Skip to main content

competitive/data_structure/
static_search.rs

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
24/// Maps a key to an unsigned integer with exactly the same ordering.
25///
26/// `BITS` must be 8, 16, 32, 64, or 128. `encode` must be deterministic, fit in
27/// `BITS`, and satisfy `a.cmp(&b) == a.encode().cmp(&b.encode())`.
28pub 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
159// Expand normalization inside the AVX-512 function so LLVM can optimize the whole batch.
160macro_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            // SAFETY: each separator is the maximum of one real child group. The public entry
289            // point excludes queries beyond the global maximum, so the first matching separator
290            // selects a real group at every level.
291            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                // SAFETY: each query is capped at the global maximum before descent. As in
318                // `descend`, every selected separator therefore names a real child group.
319                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(|&current| 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(|&current| 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(|&current| 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(|&current| 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                    // SAFETY: `simd_backend` only selects supported instruction sets. Tests and
380                    // standalone benchmarks pass supported backends to the private constructor.
381                    SimdBackend::Avx2 => unsafe { self.lower_bound_avx2(value) },
382                    // SAFETY: same as above.
383                    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                    // SAFETY: `simd_backend` only selects supported instruction sets. Tests and
401                    // standalone benchmarks pass supported backends to the private constructor.
402                    SimdBackend::Avx2 => unsafe { self.upper_bound_avx2(value) },
403                    // SAFETY: same as above.
404                    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                    // SAFETY: construction selects a supported instruction set.
421                    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                        // SAFETY: same as above.
427                        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                    // SAFETY: construction selects a supported instruction set.
442                    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                        // SAFETY: same as above.
448                        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/// A static search index over sorted integer or integer-encoded keys.
799///
800/// The index adds build time and storage. For a small number of searches, search the sorted slice
801/// directly instead.
802#[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    /// Builds an index over sorted `values`.
811    ///
812    /// # Panics
813    ///
814    /// Panics if `values` is not sorted or `SimdKey` does not satisfy its contract.
815    pub fn from_sorted(values: &[K]) -> Self {
816        Self::build(values, static_search_backend(K::BITS), false)
817    }
818
819    /// Builds a direct lookup table over sorted 8-bit or 16-bit `values`.
820    ///
821    /// This layout uses a fixed table of 257 or 65,537 positions and is intended
822    /// for query-heavy workloads.
823    ///
824    /// # Panics
825    ///
826    /// Panics if `K::BITS` is neither 8 nor 16, or if `values` is not sorted.
827    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    /// Returns the first index whose value is greater than or equal to `value`.
843    #[inline]
844    pub fn lower_bound(&self, value: K) -> usize {
845        self.storage.lower_bound(value.encode())
846    }
847
848    /// Returns one past the last index whose value is less than or equal to `value`.
849    #[inline]
850    pub fn upper_bound(&self, value: K) -> usize {
851        self.storage.upper_bound(value.encode())
852    }
853
854    /// Writes the first index greater than or equal to each value into `output`.
855    ///
856    /// # Panics
857    ///
858    /// Panics if `values` and `output` have different lengths.
859    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    /// Writes one past the last index less than or equal to each value into `output`.
865    ///
866    /// # Panics
867    ///
868    /// Panics if `values` and `output` have different lengths.
869    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}