src/bisection_tree/refine.rs

changeset 199
32f5062ee477
parent 124
6aa955ad8122
equal deleted inserted replaced
198:3868555d135c 199:32f5062ee477
1 use super::aggregator::*; 1 use super::aggregator::*;
2 use super::bt::*; 2 use super::bt::*;
3 use super::support::*; 3 use super::support::*;
4 use crate::nanleast::NaNLeast;
5 use crate::parallelism::TaskBudget; 4 use crate::parallelism::TaskBudget;
6 use crate::parallelism::{thread_pool, thread_pool_size}; 5 use crate::parallelism::{thread_pool, thread_pool_size};
7 use crate::sets::Cube; 6 use crate::sets::Cube;
8 use crate::types::*; 7 use crate::types::*;
9 use std::cmp::{max, Ord, Ordering, Ordering::*, PartialOrd}; 8 use num_traits::float::TotalOrder;
9 use std::cmp::{Ord, Ordering, Ordering::*, PartialOrd};
10 use std::collections::BinaryHeap; 10 use std::collections::BinaryHeap;
11 use std::marker::PhantomData; 11 use std::marker::PhantomData;
12 use std::sync::{Arc, Condvar, Mutex, MutexGuard}; 12 use std::sync::{Arc, Condvar, Mutex, MutexGuard};
13 13
14 /// Trait for sorting [`Aggregator`]s for [`BT`] refinement. 14 /// Trait for sorting [`Aggregator`]s for [`BT`] refinement.
17 /// with upper key less the lower key of another are discarded from the refinement process. 17 /// with upper key less the lower key of another are discarded from the refinement process.
18 /// Nodes with the highest upper sorting key are picked for refinement. 18 /// Nodes with the highest upper sorting key are picked for refinement.
19 pub trait AggregatorSorting: Sync + Send + 'static { 19 pub trait AggregatorSorting: Sync + Send + 'static {
20 // Priority 20 // Priority
21 type Agg: Aggregator; 21 type Agg: Aggregator;
22 type Sort: Ord + Copy + std::fmt::Debug + Sync + Send; 22 /// This is temporarily a Float, after removal of NanLeast, to use [
23 /// `num_traits::float::TotalOrder`] and [`num_traits::float::Float::max`].
24 /// It should be generalised by introducing a general PseudoTotalOrder trait
25 /// (because Rust's standard one is not implemented for `Float`s)
26 type Sort: Float + Copy + std::fmt::Debug + Sync + Send;
23 27
24 /// Returns lower sorting key 28 /// Returns lower sorting key
25 fn sort_lower(aggregator: &Self::Agg) -> Self::Sort; 29 fn sort_lower(aggregator: &Self::Agg) -> Self::Sort;
26 30
27 /// Returns upper sorting key 31 /// Returns upper sorting key
41 /// See [`UpperBoundSorting`] for the opposite ordering. 45 /// See [`UpperBoundSorting`] for the opposite ordering.
42 pub struct LowerBoundSorting<F: Float>(PhantomData<F>); 46 pub struct LowerBoundSorting<F: Float>(PhantomData<F>);
43 47
44 impl<F: Float> AggregatorSorting for UpperBoundSorting<F> { 48 impl<F: Float> AggregatorSorting for UpperBoundSorting<F> {
45 type Agg = Bounds<F>; 49 type Agg = Bounds<F>;
46 type Sort = NaNLeast<F>; 50 type Sort = F;
47 51
48 #[inline] 52 #[inline]
49 fn sort_lower(aggregator: &Bounds<F>) -> Self::Sort { 53 fn sort_lower(aggregator: &Bounds<F>) -> Self::Sort {
50 NaNLeast(aggregator.lower()) 54 aggregator.lower()
51 } 55 }
52 56
53 #[inline] 57 #[inline]
54 fn sort_upper(aggregator: &Bounds<F>) -> Self::Sort { 58 fn sort_upper(aggregator: &Bounds<F>) -> Self::Sort {
55 NaNLeast(aggregator.upper()) 59 aggregator.upper()
56 } 60 }
57 61
58 #[inline] 62 #[inline]
59 fn bottom() -> Self::Sort { 63 fn bottom() -> Self::Sort {
60 NaNLeast(F::NEG_INFINITY) 64 F::NEG_INFINITY
61 } 65 }
62 } 66 }
63 67
64 impl<F: Float> AggregatorSorting for LowerBoundSorting<F> { 68 impl<F: Float> AggregatorSorting for LowerBoundSorting<F> {
65 type Agg = Bounds<F>; 69 type Agg = Bounds<F>;
66 type Sort = NaNLeast<F>; 70 type Sort = F;
67 71
68 #[inline] 72 #[inline]
69 fn sort_upper(aggregator: &Bounds<F>) -> Self::Sort { 73 fn sort_upper(aggregator: &Bounds<F>) -> Self::Sort {
70 NaNLeast(-aggregator.lower()) 74 -aggregator.lower()
71 } 75 }
72 76
73 #[inline] 77 #[inline]
74 fn sort_lower(aggregator: &Bounds<F>) -> Self::Sort { 78 fn sort_lower(aggregator: &Bounds<F>) -> Self::Sort {
75 NaNLeast(-aggregator.upper()) 79 -aggregator.upper()
76 } 80 }
77 81
78 #[inline] 82 #[inline]
79 fn bottom() -> Self::Sort { 83 fn bottom() -> Self::Sort {
80 NaNLeast(F::NEG_INFINITY) 84 F::NEG_INFINITY
81 } 85 }
82 } 86 }
83 87
84 /// Return type of [`Refiner::refine`]. 88 /// Return type of [`Refiner::refine`].
85 /// 89 ///
238 S: AggregatorSorting<Agg = A>, 242 S: AggregatorSorting<Agg = A>,
239 { 243 {
240 #[inline] 244 #[inline]
241 fn cmp(&self, other: &Self) -> Ordering { 245 fn cmp(&self, other: &Self) -> Ordering {
242 self.with_aggregator(|agg1| { 246 self.with_aggregator(|agg1| {
243 other.with_aggregator(|agg2| match S::sort_upper(agg1).cmp(&S::sort_upper(agg2)) { 247 other.with_aggregator(|agg2| {
244 Equal => S::sort_lower(agg1).cmp(&S::sort_lower(agg2)), 248 match S::sort_upper(agg1).total_cmp(&S::sort_upper(agg2)) {
245 order => order, 249 Equal => S::sort_lower(agg1).total_cmp(&S::sort_lower(agg2)),
250 order => order,
251 }
246 }) 252 })
247 }) 253 })
248 } 254 }
249 } 255 }
250 256
316 ) where 322 ) where
317 S: AggregatorSorting<Agg = A>, 323 S: AggregatorSorting<Agg = A>,
318 { 324 {
319 // Insert all subnodes into the refinement heap. 325 // Insert all subnodes into the refinement heap.
320 for (node, cube) in self.nodes_and_cubes_mut(&domain) { 326 for (node, cube) in self.nodes_and_cubes_mut(&domain) {
321 container.push(RefinementInfo { 327 container.push(RefinementInfo { cube, node, refiner_info: None, sorting: PhantomData });
322 cube,
323 node,
324 refiner_info: None,
325 sorting: PhantomData,
326 });
327 } 328 }
328 } 329 }
329 } 330 }
330 331
331 impl<F: Float, D, A, const N: usize, const P: usize> Node<F, D, A, N, P> 332 impl<F: Float, D, A, const N: usize, const P: usize> Node<F, D, A, N, P>
538 } 539 }
539 540
540 // Do priority queue maintenance 541 // Do priority queue maintenance
541 if container.insert_counter > container.heap_prune_threshold { 542 if container.insert_counter > container.heap_prune_threshold {
542 // Make sure glb is good. 543 // Make sure glb is good.
543 match container.heap.iter().map(|ri| ri.sort_lower()).reduce(max) { 544 match container
545 .heap
546 .iter()
547 .map(|ri| ri.sort_lower())
548 .reduce(num_traits::Float::max)
549 {
544 Some(glb) => { 550 Some(glb) => {
545 container.glb = glb; 551 container.glb = glb;
546 // Prune 552 // Prune
547 container.heap.retain(|ri| ri.sort_upper() >= glb); 553 container.heap.retain(|ri| ri.sort_upper() >= glb);
548 } 554 }

mercurial