1use super::SimdBackend;
2#[cfg(target_arch = "x86_64")]
3use super::{avx512_enabled, simd};
4use std::{marker::PhantomData, ops::Range};
5
6#[inline]
7fn static_search_backend(bits: u32) -> SimdBackend {
8 #[cfg(target_arch = "x86_64")]
9 {
10 if avx512_enabled()
11 && is_x86_feature_detected!("avx512f")
12 && (bits != 16 || is_x86_feature_detected!("avx512bw"))
13 {
14 return SimdBackend::Avx512;
15 }
16 if is_x86_feature_detected!("avx2") {
17 return SimdBackend::Avx2;
18 }
19 }
20 let _ = bits;
21 SimdBackend::Scalar
22}
23
24pub trait SimdKey: Copy + Ord {
29 const BITS: u32;
30
31 fn encode(self) -> u128;
32}
33
34macro_rules! impl_unsigned_simd_key {
35 ($($value:ty),* $(,)?) => {
36 $(
37 impl SimdKey for $value {
38 const BITS: u32 = <$value>::BITS;
39
40 #[inline(always)]
41 fn encode(self) -> u128 {
42 self as u128
43 }
44 }
45 )*
46 };
47}
48
49macro_rules! impl_signed_simd_key {
50 ($(($signed:ty, $unsigned:ty)),* $(,)?) => {
51 $(
52 impl SimdKey for $signed {
53 const BITS: u32 = <$signed>::BITS;
54
55 #[inline(always)]
56 fn encode(self) -> u128 {
57 ((self as $unsigned) ^ ((1 as $unsigned) << (<$signed>::BITS - 1))) as u128
58 }
59 }
60 )*
61 };
62}
63
64impl_unsigned_simd_key!(u8, u16, u32, u64, u128, usize);
65impl_signed_simd_key!(
66 (i8, u8),
67 (i16, u16),
68 (i32, u32),
69 (i64, u64),
70 (i128, u128),
71 (isize, usize),
72);
73
74#[derive(Clone, Debug)]
75enum DirectStaticSearch {
76 U16(Vec<u16>),
77 U32(Vec<u32>),
78}
79
80impl DirectStaticSearch {
81 fn build<K: SimdKey>(values: &[K], bits: u32) -> Self {
82 assert!(values.len() <= u32::MAX as usize);
83 if values.len() <= u16::MAX as usize {
84 Self::U16(build_direct_positions(values, bits, |position| {
85 position as u16
86 }))
87 } else {
88 Self::U32(build_direct_positions(values, bits, |position| {
89 position as u32
90 }))
91 }
92 }
93
94 #[inline(always)]
95 fn lower_bound(&self, value: u128) -> usize {
96 let value = usize::try_from(value).expect("SimdKey::encode exceeds usize");
97 match self {
98 Self::U16(positions) => positions[value] as usize,
99 Self::U32(positions) => positions[value] as usize,
100 }
101 }
102
103 #[inline(always)]
104 fn upper_bound(&self, value: u128) -> usize {
105 let value = usize::try_from(value)
106 .ok()
107 .and_then(|value| value.checked_add(1))
108 .expect("SimdKey::encode exceeds usize");
109 match self {
110 Self::U16(positions) => positions[value] as usize,
111 Self::U32(positions) => positions[value] as usize,
112 }
113 }
114
115 #[inline(always)]
116 fn contains(&self, value: u128) -> bool {
117 let value = usize::try_from(value)
118 .ok()
119 .and_then(|value| value.checked_add(1).map(|next| (value, next)))
120 .expect("SimdKey::encode exceeds usize");
121 match self {
122 Self::U16(positions) => positions[value.0] != positions[value.1],
123 Self::U32(positions) => positions[value.0] != positions[value.1],
124 }
125 }
126}
127
128fn build_direct_positions<K, P>(values: &[K], bits: u32, position: impl Fn(usize) -> P) -> Vec<P>
129where
130 K: SimdKey,
131 P: Copy,
132{
133 let len = (1 << bits) + 1;
134 let mut positions = Vec::with_capacity(len);
135 let maximum = (1u128 << bits) - 1;
136 let mut previous: Option<(K, u128)> = None;
137 for (index, &value) in values.iter().enumerate() {
138 let encoded = value.encode();
139 assert!(
140 encoded <= maximum,
141 "SimdKey::encode exceeds its declared width"
142 );
143 if let Some((previous_value, previous_encoded)) = previous {
144 assert_eq!(
145 previous_value.cmp(&value),
146 previous_encoded.cmp(&encoded),
147 "SimdKey::encode does not preserve order"
148 );
149 }
150 if previous.is_none_or(|(_, previous)| previous != encoded) {
151 positions.resize(encoded as usize + 1, position(index));
152 }
153 previous = Some((value, encoded));
154 }
155 positions.resize(len, position(values.len()));
156 positions
157}
158
159macro_rules! bound_batch {
161 (lower, $this:ident, $values:ident, $search:expr) => {{
162 if $this.len == 0 {
163 return [0; 16];
164 }
165 let mut $values = *$values;
166 let mut beyond = [false; 16];
167 for index in 0..16 {
168 beyond[index] = $values[index] > $this.maximum;
169 $values[index] = $values[index].min($this.maximum);
170 }
171 let mut result = $search;
172 for index in 0..16 {
173 if beyond[index] {
174 result[index] = $this.len;
175 }
176 }
177 result
178 }};
179 (upper, $this:ident, $values:ident, $search:expr) => {{
180 if $this.len == 0 {
181 return [0; 16];
182 }
183 let mut $values = *$values;
184 let mut beyond = [false; 16];
185 for index in 0..16 {
186 beyond[index] = $values[index] >= $this.maximum;
187 }
188 let Some(value) = $values.iter().copied().find(|&value| value < $this.maximum) else {
189 return [$this.len; 16];
190 };
191 for index in 0..16 {
192 if beyond[index] {
193 $values[index] = value;
194 }
195 }
196 let mut result = $search;
197 for index in 0..16 {
198 if beyond[index] {
199 result[index] = $this.len;
200 }
201 }
202 result
203 }};
204}
205
206#[repr(C, align(64))]
207#[derive(Clone, Debug)]
208struct SearchBlock<T, const B: usize>([T; B]);
209
210#[derive(Clone, Debug)]
211struct StaticSearchTree<T, const B: usize> {
212 values: Vec<SearchBlock<T, B>>,
213 len: usize,
214 maximum: T,
215 levels: Vec<Vec<SearchBlock<T, B>>>,
216 #[cfg(target_arch = "x86_64")]
217 backend: SimdBackend,
218}
219
220impl<T: Copy + Ord, const B: usize> StaticSearchTree<T, B> {
221 fn build<K>(
222 values: &[K],
223 sentinel: T,
224 maximum_encoded: u128,
225 convert: impl Fn(u128) -> T,
226 backend: SimdBackend,
227 ) -> Self
228 where
229 K: SimdKey,
230 {
231 let _ = &backend;
232 let len = values.len();
233 let mut previous: Option<(K, T)> = None;
234 let mut separators = Vec::with_capacity(values.len().div_ceil(B));
235 let mut blocks = Vec::with_capacity(separators.capacity());
236 for chunk in values.chunks(B) {
237 let mut block = [sentinel; B];
238 for (index, &value) in chunk.iter().enumerate() {
239 let encoded = value.encode();
240 assert!(
241 encoded <= maximum_encoded,
242 "SimdKey::encode exceeds its declared width"
243 );
244 let encoded = convert(encoded);
245 if let Some((previous_value, previous_encoded)) = previous {
246 assert!(
247 previous_value.cmp(&value) == previous_encoded.cmp(&encoded),
248 "SimdKey::encode does not preserve order"
249 );
250 }
251 previous = Some((value, encoded));
252 block[index] = encoded;
253 }
254 separators.push(block[chunk.len() - 1]);
255 blocks.push(SearchBlock(block));
256 }
257 let maximum = separators.last().copied().unwrap_or(sentinel);
258 let mut levels = Vec::new();
259 while separators.len() > 1 {
260 let mut blocks = Vec::with_capacity(separators.len().div_ceil(B));
261 let mut next = Vec::with_capacity(blocks.capacity());
262 for chunk in separators.chunks(B) {
263 let mut block = [sentinel; B];
264 block[..chunk.len()].copy_from_slice(chunk);
265 blocks.push(SearchBlock(block));
266 next.push(chunk[chunk.len() - 1]);
267 }
268 levels.push(blocks);
269 separators = next;
270 }
271 Self {
272 values: blocks,
273 len,
274 maximum,
275 levels,
276 #[cfg(target_arch = "x86_64")]
277 backend,
278 }
279 }
280
281 #[inline(always)]
282 fn descend<F>(&self, value: T, mut position: F) -> usize
283 where
284 F: FnMut(&[T; B], T) -> usize,
285 {
286 let mut block = 0;
287 for level in self.levels.iter().rev() {
288 let values = &unsafe { level.get_unchecked(block) }.0;
292 block = block * B + position(values, value);
293 }
294 let values = &unsafe { self.values.get_unchecked(block) }.0;
295 (block * B + position(values, value)).min(self.len)
296 }
297
298 #[inline(always)]
299 fn get(&self, index: usize) -> T {
300 unsafe {
301 *self
302 .values
303 .get_unchecked(index / B)
304 .0
305 .get_unchecked(index % B)
306 }
307 }
308
309 #[inline(always)]
310 fn descend_batch<F>(&self, values: &[T; 16], mut position: F) -> [usize; 16]
311 where
312 F: FnMut(&[T; B], T) -> usize,
313 {
314 let mut blocks = [0; 16];
315 for level in self.levels.iter().rev() {
316 for index in 0..16 {
317 let block_values = &unsafe { level.get_unchecked(blocks[index]) }.0;
320 blocks[index] = blocks[index] * B + position(block_values, values[index]);
321 }
322 }
323 for index in 0..16 {
324 let block = blocks[index];
325 let block_values = &unsafe { self.values.get_unchecked(block) }.0;
326 blocks[index] = (block * B + position(block_values, values[index])).min(self.len);
327 }
328 blocks
329 }
330
331 #[inline(always)]
332 fn lower_bound_scalar(&self, value: T) -> usize {
333 self.descend(value, |values, value| {
334 values.partition_point(|¤t| current < value)
335 })
336 }
337
338 #[inline(always)]
339 fn upper_bound_scalar(&self, value: T) -> usize {
340 self.descend(value, |values, value| {
341 values.partition_point(|¤t| current <= value)
342 })
343 }
344
345 #[inline(always)]
346 fn lower_bound_batch_scalar(&self, values: &[T; 16]) -> [usize; 16] {
347 self.descend_batch(values, |values, value| {
348 values.partition_point(|¤t| current < value)
349 })
350 }
351
352 #[inline(always)]
353 fn upper_bound_batch_scalar(&self, values: &[T; 16]) -> [usize; 16] {
354 self.descend_batch(values, |values, value| {
355 values.partition_point(|¤t| current <= value)
356 })
357 }
358}
359
360macro_rules! impl_static_search_tree {
361 (
362 $value:ty,
363 $branch:expr,
364 $first_ge_avx2:ident,
365 $first_gt_avx2:ident,
366 $first_ge_avx512:ident,
367 $first_gt_avx512:ident,
368 $avx512_features:literal
369 ) => {
370 impl StaticSearchTree<$value, $branch> {
371 #[inline]
372 fn lower_bound(&self, value: $value) -> usize {
373 if self.len == 0 || value > self.maximum {
374 return self.len;
375 }
376 #[cfg(target_arch = "x86_64")]
377 return match self.backend {
378 SimdBackend::Scalar => self.lower_bound_scalar(value),
379 SimdBackend::Avx2 => unsafe { self.lower_bound_avx2(value) },
382 SimdBackend::Avx512 => unsafe { self.lower_bound_avx512(value) },
384 };
385 #[cfg(not(target_arch = "x86_64"))]
386 self.lower_bound_scalar(value)
387 }
388
389 #[inline]
390 fn upper_bound(&self, value: $value) -> usize {
391 if self.len == 0 {
392 return 0;
393 }
394 if value >= self.maximum {
395 return self.len;
396 }
397 #[cfg(target_arch = "x86_64")]
398 return match self.backend {
399 SimdBackend::Scalar => self.upper_bound_scalar(value),
400 SimdBackend::Avx2 => unsafe { self.upper_bound_avx2(value) },
403 SimdBackend::Avx512 => unsafe { self.upper_bound_avx512(value) },
405 };
406 #[cfg(not(target_arch = "x86_64"))]
407 self.upper_bound_scalar(value)
408 }
409
410 #[inline]
411 fn contains(&self, value: $value) -> bool {
412 let index = self.lower_bound(value);
413 index < self.len && self.get(index) == value
414 }
415
416 #[inline]
417 fn lower_bound_batch(&self, values: &[$value; 16]) -> [usize; 16] {
418 #[cfg(target_arch = "x86_64")]
419 if self.backend == SimdBackend::Avx512 {
420 return unsafe { self.lower_bound_batch_avx512(values) };
422 }
423 bound_batch!(lower, self, values, {
424 #[cfg(target_arch = "x86_64")]
425 let result = if self.backend == SimdBackend::Avx2 {
426 unsafe { self.lower_bound_batch_avx2(&values) }
428 } else {
429 self.lower_bound_batch_scalar(&values)
430 };
431 #[cfg(not(target_arch = "x86_64"))]
432 let result = self.lower_bound_batch_scalar(&values);
433 result
434 })
435 }
436
437 #[inline]
438 fn upper_bound_batch(&self, values: &[$value; 16]) -> [usize; 16] {
439 #[cfg(target_arch = "x86_64")]
440 if self.backend == SimdBackend::Avx512 {
441 return unsafe { self.upper_bound_batch_avx512(values) };
443 }
444 bound_batch!(upper, self, values, {
445 #[cfg(target_arch = "x86_64")]
446 let result = if self.backend == SimdBackend::Avx2 {
447 unsafe { self.upper_bound_batch_avx2(&values) }
449 } else {
450 self.upper_bound_batch_scalar(&values)
451 };
452 #[cfg(not(target_arch = "x86_64"))]
453 let result = self.upper_bound_batch_scalar(&values);
454 result
455 })
456 }
457
458 #[cfg(target_arch = "x86_64")]
459 #[target_feature(enable = "avx2")]
460 unsafe fn lower_bound_avx2(&self, value: $value) -> usize {
461 self.descend(value, |values, value| unsafe {
462 simd::$first_ge_avx2(values, value)
463 })
464 }
465
466 #[cfg(target_arch = "x86_64")]
467 #[target_feature(enable = "avx2")]
468 unsafe fn upper_bound_avx2(&self, value: $value) -> usize {
469 self.descend(value, |values, value| unsafe {
470 simd::$first_gt_avx2(values, value)
471 })
472 }
473
474 #[cfg(target_arch = "x86_64")]
475 #[target_feature(enable = "avx2")]
476 unsafe fn lower_bound_batch_avx2(&self, values: &[$value; 16]) -> [usize; 16] {
477 self.descend_batch(values, |values, value| unsafe {
478 simd::$first_ge_avx2(values, value)
479 })
480 }
481
482 #[cfg(target_arch = "x86_64")]
483 #[target_feature(enable = "avx2")]
484 unsafe fn upper_bound_batch_avx2(&self, values: &[$value; 16]) -> [usize; 16] {
485 self.descend_batch(values, |values, value| unsafe {
486 simd::$first_gt_avx2(values, value)
487 })
488 }
489
490 #[cfg(target_arch = "x86_64")]
491 #[target_feature(enable = $avx512_features)]
492 unsafe fn lower_bound_avx512(&self, value: $value) -> usize {
493 self.descend(value, |values, value| unsafe {
494 simd::$first_ge_avx512(values, value)
495 })
496 }
497
498 #[cfg(target_arch = "x86_64")]
499 #[target_feature(enable = $avx512_features)]
500 unsafe fn upper_bound_avx512(&self, value: $value) -> usize {
501 self.descend(value, |values, value| unsafe {
502 simd::$first_gt_avx512(values, value)
503 })
504 }
505
506 #[cfg(target_arch = "x86_64")]
507 #[target_feature(enable = $avx512_features)]
508 unsafe fn lower_bound_batch_avx512(&self, values: &[$value; 16]) -> [usize; 16] {
509 bound_batch!(
510 lower,
511 self,
512 values,
513 self.descend_batch(&values, |values, value| unsafe {
514 simd::$first_ge_avx512(values, value)
515 })
516 )
517 }
518
519 #[cfg(target_arch = "x86_64")]
520 #[target_feature(enable = $avx512_features)]
521 unsafe fn upper_bound_batch_avx512(&self, values: &[$value; 16]) -> [usize; 16] {
522 bound_batch!(
523 upper,
524 self,
525 values,
526 self.descend_batch(&values, |values, value| unsafe {
527 simd::$first_gt_avx512(values, value)
528 })
529 )
530 }
531 }
532 };
533}
534
535impl_static_search_tree!(
536 u16,
537 32,
538 first_ge_u16x32_avx2,
539 first_gt_u16x32_avx2,
540 first_ge_u16x32_avx512,
541 first_gt_u16x32_avx512,
542 "avx512f,avx512bw"
543);
544impl_static_search_tree!(
545 u32,
546 16,
547 first_ge_u32x16_avx2,
548 first_gt_u32x16_avx2,
549 first_ge_u32x16_avx512,
550 first_gt_u32x16_avx512,
551 "avx512f"
552);
553impl_static_search_tree!(
554 u64,
555 8,
556 first_ge_u64x8_avx2,
557 first_gt_u64x8_avx2,
558 first_ge_u64x8_avx512,
559 first_gt_u64x8_avx512,
560 "avx512f"
561);
562
563impl StaticSearchTree<u128, 4> {
564 #[inline]
565 fn lower_bound(&self, value: u128) -> usize {
566 if self.len == 0 || value > self.maximum {
567 self.len
568 } else {
569 self.descend(value, |values, value| {
570 (values[0] < value) as usize
571 + (values[1] < value) as usize
572 + (values[2] < value) as usize
573 + (values[3] < value) as usize
574 })
575 }
576 }
577
578 #[inline]
579 fn upper_bound(&self, value: u128) -> usize {
580 if self.len == 0 || value >= self.maximum {
581 self.len
582 } else {
583 self.descend(value, |values, value| {
584 (values[0] <= value) as usize
585 + (values[1] <= value) as usize
586 + (values[2] <= value) as usize
587 + (values[3] <= value) as usize
588 })
589 }
590 }
591
592 #[inline]
593 fn contains(&self, value: u128) -> bool {
594 let index = self.lower_bound(value);
595 index < self.len && self.get(index) == value
596 }
597
598 #[inline]
599 fn lower_bound_batch(&self, values: &[u128; 16]) -> [usize; 16] {
600 bound_batch!(
601 lower,
602 self,
603 values,
604 self.descend_batch(&values, |values, value| {
605 (values[0] < value) as usize
606 + (values[1] < value) as usize
607 + (values[2] < value) as usize
608 + (values[3] < value) as usize
609 })
610 )
611 }
612
613 #[inline]
614 fn upper_bound_batch(&self, values: &[u128; 16]) -> [usize; 16] {
615 bound_batch!(
616 upper,
617 self,
618 values,
619 self.descend_batch(&values, |values, value| {
620 (values[0] <= value) as usize
621 + (values[1] <= value) as usize
622 + (values[2] <= value) as usize
623 + (values[3] <= value) as usize
624 })
625 )
626 }
627}
628
629fn search_batch<K, T>(
630 values: &[K],
631 output: &mut [usize],
632 convert: impl Fn(u128) -> T,
633 single: impl Fn(T) -> usize,
634 batch: impl Fn(&[T; 16]) -> [usize; 16],
635) where
636 K: SimdKey,
637 T: Copy,
638{
639 let mut offset = 0;
640 while offset + 16 <= values.len() {
641 let values = std::array::from_fn(|index| convert(values[offset + index].encode()));
642 output[offset..offset + 16].copy_from_slice(&batch(&values));
643 offset += 16;
644 }
645 let remaining = values.len() - offset;
646 if remaining >= 8 {
647 let mut encoded = [convert(values[offset].encode()); 16];
648 for index in 1..remaining {
649 encoded[index] = convert(values[offset + index].encode());
650 }
651 let positions = batch(&encoded);
652 output[offset..].copy_from_slice(&positions[..remaining]);
653 } else {
654 for (&value, position) in values[offset..].iter().zip(&mut output[offset..]) {
655 *position = single(convert(value.encode()));
656 }
657 }
658}
659
660#[derive(Clone, Debug)]
661enum StaticSearchStorage {
662 Direct(DirectStaticSearch),
663 U16(StaticSearchTree<u16, 32>),
664 U32(StaticSearchTree<u32, 16>),
665 U64(StaticSearchTree<u64, 8>),
666 U128(StaticSearchTree<u128, 4>),
667}
668
669impl StaticSearchStorage {
670 #[inline(always)]
671 fn lower_bound(&self, value: u128) -> usize {
672 match self {
673 Self::Direct(search) => search.lower_bound(value),
674 Self::U16(search) => search.lower_bound(
675 u16::try_from(value).expect("SimdKey::encode exceeds its declared width"),
676 ),
677 Self::U32(search) => search.lower_bound(
678 u32::try_from(value).expect("SimdKey::encode exceeds its declared width"),
679 ),
680 Self::U64(search) => search.lower_bound(
681 u64::try_from(value).expect("SimdKey::encode exceeds its declared width"),
682 ),
683 Self::U128(search) => search.lower_bound(value),
684 }
685 }
686
687 #[inline(always)]
688 fn upper_bound(&self, value: u128) -> usize {
689 match self {
690 Self::Direct(search) => search.upper_bound(value),
691 Self::U16(search) => search.upper_bound(
692 u16::try_from(value).expect("SimdKey::encode exceeds its declared width"),
693 ),
694 Self::U32(search) => search.upper_bound(
695 u32::try_from(value).expect("SimdKey::encode exceeds its declared width"),
696 ),
697 Self::U64(search) => search.upper_bound(
698 u64::try_from(value).expect("SimdKey::encode exceeds its declared width"),
699 ),
700 Self::U128(search) => search.upper_bound(value),
701 }
702 }
703
704 #[inline(always)]
705 fn contains(&self, value: u128) -> bool {
706 match self {
707 Self::Direct(search) => search.contains(value),
708 Self::U16(search) => search.contains(
709 u16::try_from(value).expect("SimdKey::encode exceeds its declared width"),
710 ),
711 Self::U32(search) => search.contains(
712 u32::try_from(value).expect("SimdKey::encode exceeds its declared width"),
713 ),
714 Self::U64(search) => search.contains(
715 u64::try_from(value).expect("SimdKey::encode exceeds its declared width"),
716 ),
717 Self::U128(search) => search.contains(value),
718 }
719 }
720
721 fn lower_bound_batch<K: SimdKey>(&self, values: &[K], output: &mut [usize]) {
722 match self {
723 Self::Direct(search) => {
724 for (&value, position) in values.iter().zip(output) {
725 *position = search.lower_bound(value.encode());
726 }
727 }
728 Self::U16(search) => search_batch(
729 values,
730 output,
731 |value| u16::try_from(value).expect("SimdKey::encode exceeds its declared width"),
732 |value| search.lower_bound(value),
733 |values| search.lower_bound_batch(values),
734 ),
735 Self::U32(search) => search_batch(
736 values,
737 output,
738 |value| u32::try_from(value).expect("SimdKey::encode exceeds its declared width"),
739 |value| search.lower_bound(value),
740 |values| search.lower_bound_batch(values),
741 ),
742 Self::U64(search) => search_batch(
743 values,
744 output,
745 |value| u64::try_from(value).expect("SimdKey::encode exceeds its declared width"),
746 |value| search.lower_bound(value),
747 |values| search.lower_bound_batch(values),
748 ),
749 Self::U128(search) => search_batch(
750 values,
751 output,
752 |value| value,
753 |value| search.lower_bound(value),
754 |values| search.lower_bound_batch(values),
755 ),
756 }
757 }
758
759 fn upper_bound_batch<K: SimdKey>(&self, values: &[K], output: &mut [usize]) {
760 match self {
761 Self::Direct(search) => {
762 for (&value, position) in values.iter().zip(output) {
763 *position = search.upper_bound(value.encode());
764 }
765 }
766 Self::U16(search) => search_batch(
767 values,
768 output,
769 |value| u16::try_from(value).expect("SimdKey::encode exceeds its declared width"),
770 |value| search.upper_bound(value),
771 |values| search.upper_bound_batch(values),
772 ),
773 Self::U32(search) => search_batch(
774 values,
775 output,
776 |value| u32::try_from(value).expect("SimdKey::encode exceeds its declared width"),
777 |value| search.upper_bound(value),
778 |values| search.upper_bound_batch(values),
779 ),
780 Self::U64(search) => search_batch(
781 values,
782 output,
783 |value| u64::try_from(value).expect("SimdKey::encode exceeds its declared width"),
784 |value| search.upper_bound(value),
785 |values| search.upper_bound_batch(values),
786 ),
787 Self::U128(search) => search_batch(
788 values,
789 output,
790 |value| value,
791 |value| search.upper_bound(value),
792 |values| search.upper_bound_batch(values),
793 ),
794 }
795 }
796}
797
798#[derive(Clone, Debug)]
803pub struct StaticSearch<K> {
804 storage: StaticSearchStorage,
805 len: usize,
806 marker: PhantomData<fn() -> K>,
807}
808
809impl<K: SimdKey> StaticSearch<K> {
810 pub fn from_sorted(values: &[K]) -> Self {
816 Self::build(values, static_search_backend(K::BITS), false)
817 }
818
819 pub fn from_sorted_direct(values: &[K]) -> Self {
828 assert!(matches!(K::BITS, 8 | 16));
829 Self::build(values, static_search_backend(K::BITS), true)
830 }
831
832 #[inline]
833 pub fn len(&self) -> usize {
834 self.len
835 }
836
837 #[inline]
838 pub fn is_empty(&self) -> bool {
839 self.len == 0
840 }
841
842 #[inline]
844 pub fn lower_bound(&self, value: K) -> usize {
845 self.storage.lower_bound(value.encode())
846 }
847
848 #[inline]
850 pub fn upper_bound(&self, value: K) -> usize {
851 self.storage.upper_bound(value.encode())
852 }
853
854 pub fn lower_bound_batch(&self, values: &[K], output: &mut [usize]) {
860 assert_eq!(values.len(), output.len());
861 self.storage.lower_bound_batch(values, output);
862 }
863
864 pub fn upper_bound_batch(&self, values: &[K], output: &mut [usize]) {
870 assert_eq!(values.len(), output.len());
871 self.storage.upper_bound_batch(values, output);
872 }
873
874 #[inline]
875 pub fn range(&self, value: K) -> Range<usize> {
876 let value = value.encode();
877 self.storage.lower_bound(value)..self.storage.upper_bound(value)
878 }
879
880 #[inline]
881 pub fn contains(&self, value: K) -> bool {
882 self.storage.contains(value.encode())
883 }
884
885 fn build(values: &[K], backend: SimdBackend, direct: bool) -> Self {
886 assert!(matches!(K::BITS, 8 | 16 | 32 | 64 | 128));
887 assert!(values.windows(2).all(|pair| pair[0] <= pair[1]));
888 let len = values.len();
889 let storage = match K::BITS {
890 8 => StaticSearchStorage::Direct(DirectStaticSearch::build(values, K::BITS)),
891 16 => {
892 if direct {
893 StaticSearchStorage::Direct(DirectStaticSearch::build(values, K::BITS))
894 } else {
895 StaticSearchStorage::U16(StaticSearchTree::build(
896 values,
897 u16::MAX,
898 u16::MAX as u128,
899 |value| value as u16,
900 backend,
901 ))
902 }
903 }
904 32 => StaticSearchStorage::U32(StaticSearchTree::build(
905 values,
906 u32::MAX,
907 u32::MAX as u128,
908 |value| value as u32,
909 backend,
910 )),
911 64 => StaticSearchStorage::U64(StaticSearchTree::build(
912 values,
913 u64::MAX,
914 u64::MAX as u128,
915 |value| value as u64,
916 backend,
917 )),
918 128 => StaticSearchStorage::U128(StaticSearchTree::build(
919 values,
920 u128::MAX,
921 u128::MAX,
922 |value| value,
923 backend,
924 )),
925 _ => unreachable!(),
926 };
927 Self {
928 storage,
929 len,
930 marker: PhantomData,
931 }
932 }
933}
934
935#[cfg(test)]
936mod tests {
937 use super::*;
938 use crate::tools::Xorshift;
939 #[cfg(target_arch = "x86_64")]
940 use crate::tools::avx512_supported;
941 use std::fmt::Debug;
942
943 #[cfg(target_arch = "x86_64")]
944 fn backends() -> Vec<SimdBackend> {
945 let mut result = vec![SimdBackend::Scalar];
946 if is_x86_feature_detected!("avx2") {
947 result.push(SimdBackend::Avx2);
948 }
949 if avx512_supported() {
950 result.push(SimdBackend::Avx512);
951 }
952 result
953 }
954
955 #[cfg(not(target_arch = "x86_64"))]
956 fn backends() -> Vec<SimdBackend> {
957 vec![SimdBackend::Scalar]
958 }
959
960 fn check<K>(values: Vec<K>, queries: &[K])
961 where
962 K: SimdKey + Debug,
963 {
964 let verify = |search: StaticSearch<K>| {
965 assert_eq!(search.len(), values.len());
966 assert_eq!(search.is_empty(), values.is_empty());
967 for &query in queries {
968 let left = values.partition_point(|&value| value < query);
969 let right = values.partition_point(|&value| value <= query);
970 assert_eq!(search.lower_bound(query), left);
971 assert_eq!(search.upper_bound(query), right);
972 assert_eq!(search.range(query), left..right);
973 assert_eq!(search.contains(query), left != right);
974 }
975 for len in [0, 1, 7, 8, 15, 16, 17, queries.len()] {
976 let queries = &queries[..len.min(queries.len())];
977 let mut left = vec![0; queries.len()];
978 let mut right = vec![0; queries.len()];
979 search.lower_bound_batch(queries, &mut left);
980 search.upper_bound_batch(queries, &mut right);
981 assert_eq!(
982 left,
983 queries
984 .iter()
985 .map(|query| values.partition_point(|value| value < query))
986 .collect::<Vec<_>>()
987 );
988 assert_eq!(
989 right,
990 queries
991 .iter()
992 .map(|query| values.partition_point(|value| value <= query))
993 .collect::<Vec<_>>()
994 );
995 }
996 };
997 if K::BITS == 8 {
998 verify(StaticSearch::build(&values, SimdBackend::Scalar, false));
999 } else {
1000 for backend in backends() {
1001 verify(StaticSearch::build(&values, backend, false));
1002 }
1003 }
1004 if K::BITS == 16 {
1005 verify(StaticSearch::build(&values, SimdBackend::Scalar, true));
1006 }
1007 }
1008
1009 fn check_random<K>(
1010 rng: &mut Xorshift,
1011 mut random: impl FnMut(&mut Xorshift) -> K,
1012 boundaries: &[K],
1013 ) where
1014 K: SimdKey + Debug,
1015 {
1016 for len in [0, 1, 7, 8, 15, 16, 17, 31, 32, 33, 255, 256, 257, 4097] {
1017 let mut values: Vec<_> = (0..len).map(|_| random(rng)).collect();
1018 for (value, &boundary) in values.iter_mut().zip(boundaries) {
1019 *value = boundary;
1020 }
1021 if values.len() > boundaries.len() {
1022 values[boundaries.len()] = boundaries[0];
1023 }
1024 values.sort_unstable();
1025 let mut queries: Vec<_> = (0..263).map(|_| random(rng)).collect();
1026 queries.extend_from_slice(boundaries);
1027 check(values, &queries);
1028 }
1029 }
1030
1031 #[test]
1032 fn test_static_search() {
1033 let mut rng = Xorshift::default();
1034 check_random(&mut rng, |rng| rng.rand64() as u8, &[u8::MIN, 1, u8::MAX]);
1035 check_random(
1036 &mut rng,
1037 |rng| rng.rand64() as i8,
1038 &[i8::MIN, -1, 0, 1, i8::MAX],
1039 );
1040 check_random(
1041 &mut rng,
1042 |rng| rng.rand64() as u16,
1043 &[u16::MIN, 1, u16::MAX],
1044 );
1045 check_random(
1046 &mut rng,
1047 |rng| rng.rand64() as i16,
1048 &[i16::MIN, -1, 0, 1, i16::MAX],
1049 );
1050 check_random(
1051 &mut rng,
1052 |rng| rng.rand64() as u32,
1053 &[u32::MIN, 1, u32::MAX],
1054 );
1055 check_random(
1056 &mut rng,
1057 |rng| rng.rand64() as i32,
1058 &[i32::MIN, -1, 0, 1, i32::MAX],
1059 );
1060 check_random(&mut rng, |rng| rng.rand64(), &[u64::MIN, 1, u64::MAX]);
1061 check_random(
1062 &mut rng,
1063 |rng| rng.rand64() as i64,
1064 &[i64::MIN, -1, 0, 1, i64::MAX],
1065 );
1066 check_random(
1067 &mut rng,
1068 |rng| (rng.rand64() as u128) << 64 | rng.rand64() as u128,
1069 &[u128::MIN, 1, u128::MAX],
1070 );
1071 check_random(
1072 &mut rng,
1073 |rng| ((rng.rand64() as u128) << 64 | rng.rand64() as u128) as i128,
1074 &[i128::MIN, -1, 0, 1, i128::MAX],
1075 );
1076 check_random(
1077 &mut rng,
1078 |rng| rng.rand64() as usize,
1079 &[usize::MIN, 1, usize::MAX],
1080 );
1081 check_random(
1082 &mut rng,
1083 |rng| rng.rand64() as isize,
1084 &[isize::MIN, -1, 0, 1, isize::MAX],
1085 );
1086
1087 let mut values: Vec<_> = (0..=u16::MAX).map(|_| rng.rand64() as u16).collect();
1088 values.sort_unstable();
1089 let queries: Vec<_> = (0..263).map(|_| rng.rand64() as u16).collect();
1090 check(values, &queries);
1091
1092 #[derive(Clone, Copy, Debug, Eq, Ord, PartialEq, PartialOrd)]
1093 struct Pair(u32, u32);
1094
1095 impl SimdKey for Pair {
1096 const BITS: u32 = 64;
1097
1098 fn encode(self) -> u128 {
1099 ((self.0 as u128) << 32) | self.1 as u128
1100 }
1101 }
1102
1103 check_random(
1104 &mut rng,
1105 |rng| Pair(rng.rand64() as u32, rng.rand64() as u32),
1106 &[Pair(0, 0), Pair(u32::MAX, u32::MAX)],
1107 );
1108 }
1109}