Skip to main content

competitive/data_structure/
wavelet_matrix.rs

1use super::{
2    AbelianGroup, BinaryIndexedTree, BitVector, Compressor, OrderedCompressor,
3    RankSelectDictionaries, VecCompress,
4};
5use std::{
6    mem::{self, MaybeUninit},
7    ops::Range,
8};
9
10const ACCESS: u8 = 0;
11const RANK: u8 = 1;
12const QUANTILE: u8 = 2;
13const RANK_LESSTHAN: u8 = 3;
14
15#[derive(Debug, Clone)]
16#[repr(C)]
17struct WaveletMatrixQuadBlock {
18    lo: u64,
19    hi: u64,
20    rank: [u32; 4],
21}
22
23#[derive(Debug, Clone)]
24struct WaveletMatrixQuadVector {
25    blocks: Vec<WaveletMatrixQuadBlock>,
26    starts: [usize; 4],
27    select_samples: [Vec<usize>; 4],
28}
29
30impl WaveletMatrixQuadVector {
31    fn from_words(low: &[u64], high: Option<&[u64]>, len: usize) -> Self {
32        let mut blocks = Vec::with_capacity(len / 64 + 1);
33        let mut rank = [0; 4];
34        for (word, &lo) in low.iter().enumerate() {
35            let count = (len - word * 64).min(64);
36            let hi = high.map_or(0, |high| high[word]);
37            blocks.push(WaveletMatrixQuadBlock { lo, hi, rank });
38            let low = lo.count_ones();
39            let high = hi.count_ones();
40            let both = (lo & hi).count_ones();
41            rank[0] += count as u32 - low - (high - both);
42            rank[1] += low - both;
43            rank[2] += high - both;
44            rank[3] += both;
45        }
46        if len.is_multiple_of(64) {
47            blocks.push(WaveletMatrixQuadBlock { lo: 0, hi: 0, rank });
48        }
49        // Two stable binary partitions order the children as 0, 2, 1, 3.
50        let starts = [
51            0,
52            (rank[0] as usize + rank[2] as usize),
53            rank[0] as usize,
54            rank[0] as usize + rank[2] as usize + rank[1] as usize,
55        ];
56        let mut select_samples: [Vec<usize>; 4] = std::array::from_fn(|_| Vec::new());
57        for (i, block) in blocks.iter().enumerate().take(len.div_ceil(64)) {
58            for digit in 0..4 {
59                let start = block.rank[digit] as usize;
60                let end = blocks
61                    .get(i + 1)
62                    .map_or(rank[digit], |next| next.rank[digit])
63                    as usize;
64                if start.div_ceil(128) != end.div_ceil(128) {
65                    select_samples[digit].push(i);
66                }
67            }
68        }
69        Self {
70            blocks,
71            starts,
72            select_samples,
73        }
74    }
75
76    #[inline]
77    fn ranks(&self, position: usize) -> [usize; 4] {
78        let block = &self.blocks[position / 64];
79        let mask = !(u64::MAX << (position % 64));
80        let low = (block.lo & mask).count_ones() as usize;
81        let high = (block.hi & mask).count_ones() as usize;
82        let both = (block.lo & block.hi & mask).count_ones() as usize;
83        [
84            block.rank[0] as usize + position % 64 - low - (high - both),
85            block.rank[1] as usize + low - both,
86            block.rank[2] as usize + high - both,
87            block.rank[3] as usize + both,
88        ]
89    }
90
91    #[inline]
92    fn rank(&self, digit: usize, position: usize) -> usize {
93        let block = &self.blocks[position / 64];
94        let low = if digit & 1 != 0 { block.lo } else { !block.lo };
95        let high = if digit & 2 != 0 { block.hi } else { !block.hi };
96        block.rank[digit] as usize
97            + (low & high & !(u64::MAX << (position % 64))).count_ones() as usize
98    }
99
100    fn select(&self, digit: usize, k: usize) -> usize {
101        let sample = k / 128;
102        let start = self.select_samples[digit][sample];
103        let end = self.select_samples[digit]
104            .get(sample + 1)
105            .map_or(self.blocks.len(), |&word| word + 1);
106        let word = start
107            + self.blocks[start..end].partition_point(|block| block.rank[digit] as usize <= k)
108            - 1;
109        let block = &self.blocks[word];
110        let lo = if digit & 1 != 0 { block.lo } else { !block.lo };
111        let hi = if digit & 2 != 0 { block.hi } else { !block.hi };
112        word * 64 + BitVector::select_word(lo & hi, k - block.rank[digit] as usize)
113    }
114
115    #[inline]
116    fn access_rank(&self, position: usize) -> (usize, usize) {
117        let block = &self.blocks[position / 64];
118        let offset = position % 64;
119        let digit =
120            ((block.lo >> offset) & 1) as usize | (((block.hi >> offset) & 1) as usize) << 1;
121        (digit, self.rank(digit, position))
122    }
123}
124
125#[cfg(target_arch = "x86_64")]
126mod simd {
127    #![allow(unsafe_op_in_unsafe_fn)] // All entry points check the required CPU features.
128    use super::BitVector;
129    use std::arch::x86_64::*;
130
131    #[target_feature(enable = "avx2")]
132    pub unsafe fn pack_words(indices: &[u32], bit: usize) -> Vec<u64> {
133        let mut result = Vec::with_capacity(indices.len().div_ceil(64));
134        let shift = _mm_cvtsi32_si128((31 - bit) as i32);
135        for chunk in indices.chunks(64) {
136            let mut word = 0u64;
137            let mut i = 0;
138            while i + 8 <= chunk.len() {
139                let v = _mm256_loadu_si256(chunk.as_ptr().add(i).cast());
140                word |= (_mm256_movemask_ps(_mm256_castsi256_ps(_mm256_sll_epi32(v, shift)))
141                    as u64)
142                    << i;
143                i += 8;
144            }
145            for (j, &value) in chunk[i..].iter().enumerate() {
146                word |= (((value >> bit) & 1) as u64) << (i + j);
147            }
148            result.push(word);
149        }
150        result
151    }
152
153    #[target_feature(enable = "avx2")]
154    pub unsafe fn partition_avx2(indices: &[u32], words: &[u64], mut one: usize, next: &mut [u32]) {
155        const PERM: [[i32; 8]; 256] = {
156            let mut table = [[0; 8]; 256];
157            let mut mask = 0;
158            while mask < 256 {
159                let mut pos = 0;
160                let mut bit = 0;
161                while bit < 2 {
162                    let mut i = 0;
163                    while i < 8 {
164                        if (mask >> i) & 1 == bit {
165                            table[mask][pos] = i;
166                            pos += 1;
167                        }
168                        i += 1;
169                    }
170                    bit += 1;
171                }
172                mask += 1;
173            }
174            table
175        };
176        let order = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
177        let mut zero = 0;
178        for (chunk, &word) in indices.chunks(64).zip(words) {
179            if word == 0 || word == u64::MAX {
180                let at = if word == 0 { &mut zero } else { &mut one };
181                next[*at..*at + chunk.len()].copy_from_slice(chunk);
182                *at += chunk.len();
183                continue;
184            }
185            let mut i = 0;
186            while i + 8 <= chunk.len() {
187                let value = _mm256_loadu_si256(chunk.as_ptr().add(i).cast());
188                let mask = ((word >> i) & 255) as usize;
189                let ones = mask.count_ones() as usize;
190                let zeros = 8 - ones;
191                let packed = _mm256_permutevar8x32_epi32(
192                    value,
193                    _mm256_loadu_si256(PERM[mask].as_ptr().cast()),
194                );
195                let rotated = _mm256_permutevar8x32_epi32(
196                    packed,
197                    _mm256_add_epi32(order, _mm256_set1_epi32(zeros as i32)),
198                );
199                _mm256_maskstore_epi32(
200                    next.as_mut_ptr().add(zero).cast(),
201                    _mm256_cmpgt_epi32(_mm256_set1_epi32(zeros as i32), order),
202                    packed,
203                );
204                _mm256_maskstore_epi32(
205                    next.as_mut_ptr().add(one).cast(),
206                    _mm256_cmpgt_epi32(_mm256_set1_epi32(ones as i32), order),
207                    rotated,
208                );
209                zero += zeros;
210                one += ones;
211                i += 8;
212            }
213            for (j, &value) in chunk[i..].iter().enumerate() {
214                let bit = (word >> (i + j)) & 1 != 0;
215                next[if bit { one } else { zero }] = value;
216                zero += !bit as usize;
217                one += bit as usize;
218            }
219        }
220    }
221
222    #[target_feature(enable = "avx2")]
223    #[inline]
224    unsafe fn popcount_avx2(value: __m256i) -> __m256i {
225        let lookup = _mm256_setr_epi8(
226            0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4, 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2,
227            3, 3, 4,
228        );
229        let mask = _mm256_set1_epi8(15);
230        let low = _mm256_shuffle_epi8(lookup, _mm256_and_si256(value, mask));
231        let high = _mm256_shuffle_epi8(
232            lookup,
233            _mm256_and_si256(_mm256_srli_epi16::<4>(value), mask),
234        );
235        _mm256_sad_epu8(_mm256_add_epi8(low, high), _mm256_setzero_si256())
236    }
237
238    #[target_feature(enable = "avx512f")]
239    #[inline]
240    unsafe fn greater_avx512(a: __m512i, b: __m512i) -> __m512i {
241        _mm512_maskz_set1_epi64(_mm512_cmpgt_epi64_mask(a, b), -1)
242    }
243
244    #[target_feature(enable = "avx512f")]
245    #[inline]
246    unsafe fn gather_avx512<const SCALE: i32>(base: *const i64, index: __m512i) -> __m512i {
247        _mm512_i64gather_epi64::<SCALE>(index, base)
248    }
249
250    macro_rules! rank_lessthan {
251        ($name:ident, $feature:literal, $lanes:literal,
252         $load:ident, $store:ident, $set:ident, $gather:ident,
253         $add:ident, $sub:ident, $and:ident, $andnot:ident, $or:ident,
254         $left:ident, $right:ident, $gt:ident, $popcount:ident) => {
255            #[target_feature(enable = $feature)]
256            pub unsafe fn $name(vectors: &[BitVector], zeros: &[usize], states: &mut [[usize; 4]]) {
257                for chunk in states.chunks_exact_mut($lanes) {
258                    let mut starts = [0u64; $lanes];
259                    let mut ends = [0u64; $lanes];
260                    let mut keys = [0u64; $lanes];
261                    let mut values = [0u64; $lanes];
262                    for (i, state) in chunk.iter().enumerate() {
263                        starts[i] = state[0] as u64;
264                        ends[i] = state[1] as u64;
265                        keys[i] = state[2] as u64;
266                    }
267                    let mut start = $load(starts.as_ptr().cast());
268                    let mut end = $load(ends.as_ptr().cast());
269                    let key = $load(keys.as_ptr().cast());
270                    let mut value = $set(0);
271                    let one = $set(1);
272                    for (level, vector) in vectors.iter().enumerate() {
273                        let d = vectors.len() - 1 - level;
274                        let base = vector.blocks().as_ptr().cast::<i64>();
275                        let rank1 = |position| {
276                            let offset = $and(position, $set(63));
277                            let index = $left($right(position, $set(6)), one);
278                            let bits = $gather::<8>(base, index);
279                            let prefix = $gather::<8>(base, $add(index, one));
280                            $add(prefix, $popcount($and(bits, $sub($left(one, offset), one))))
281                        };
282                        let rank = rank1(start);
283                        let end_rank = rank1(end);
284                        let start0 = $sub(start, rank);
285                        let end0 = $sub(end, end_rank);
286                        let count = $sub(end0, start0);
287                        let mask = $gt($and($right(key, $set(d as i64)), one), $set(0));
288                        let zero = $set(zeros[level] as i64);
289                        start = $or($andnot(mask, start0), $and(mask, $add(zero, rank)));
290                        end = $or($andnot(mask, end0), $and(mask, $add(zero, end_rank)));
291                        value = $add(value, $and(mask, count));
292                    }
293                    $store(values.as_mut_ptr().cast(), value);
294                    for (i, state) in chunk.iter_mut().enumerate() {
295                        state[3] = values[i] as usize;
296                    }
297                }
298            }
299        };
300    }
301
302    rank_lessthan!(
303        rank_lessthan_avx2,
304        "avx2",
305        4,
306        _mm256_loadu_si256,
307        _mm256_storeu_si256,
308        _mm256_set1_epi64x,
309        _mm256_i64gather_epi64,
310        _mm256_add_epi64,
311        _mm256_sub_epi64,
312        _mm256_and_si256,
313        _mm256_andnot_si256,
314        _mm256_or_si256,
315        _mm256_sllv_epi64,
316        _mm256_srlv_epi64,
317        _mm256_cmpgt_epi64,
318        popcount_avx2
319    );
320    rank_lessthan!(
321        rank_lessthan_avx512,
322        "avx512f,avx512vpopcntdq",
323        8,
324        _mm512_loadu_si512,
325        _mm512_storeu_si512,
326        _mm512_set1_epi64,
327        gather_avx512,
328        _mm512_add_epi64,
329        _mm512_sub_epi64,
330        _mm512_and_si512,
331        _mm512_andnot_si512,
332        _mm512_or_si512,
333        _mm512_sllv_epi64,
334        _mm512_srlv_epi64,
335        greater_avx512,
336        _mm512_popcnt_epi64
337    );
338
339    #[target_feature(enable = "avx512f,avx512vpopcntdq")]
340    pub unsafe fn quad_avx512(
341        layers: &[super::WaveletMatrixQuadVector],
342        states: &mut [[usize; 4]],
343    ) {
344        for chunk in states.as_chunks_mut::<8>().0 {
345            let mut starts = [0u64; 8];
346            let mut ends = [0u64; 8];
347            let mut keys = [0u64; 8];
348            let mut result = [0u64; 8];
349            for (i, s) in chunk.iter().enumerate() {
350                starts[i] = s[0] as u64;
351                ends[i] = s[1] as u64;
352                keys[i] = s[2] as u64;
353            }
354            let mut start = _mm512_loadu_si512(starts.as_ptr().cast());
355            let mut end = _mm512_loadu_si512(ends.as_ptr().cast());
356            let mut key = _mm512_loadu_si512(keys.as_ptr().cast());
357            let mut code = _mm512_set1_epi64(0);
358            let one = _mm512_set1_epi64(1);
359            macro_rules! blend {
360                ($m:expr,$x:expr,$y:expr) => {
361                    _mm512_or_si512(_mm512_andnot_si512($m, $x), _mm512_and_si512($m, $y))
362                };
363            }
364            for layer in layers {
365                let base = layer.blocks.as_ptr().cast::<i64>();
366                let ranks = |pos| {
367                    let offset = _mm512_and_si512(pos, _mm512_set1_epi64(63));
368                    let mask = _mm512_sub_epi64(_mm512_sllv_epi64(one, offset), one);
369                    let i = _mm512_sllv_epi64(
370                        _mm512_srlv_epi64(pos, _mm512_set1_epi64(6)),
371                        _mm512_set1_epi64(2),
372                    );
373                    let lo = _mm512_and_si512(gather_avx512::<8>(base, i), mask);
374                    let hi =
375                        _mm512_and_si512(gather_avx512::<8>(base, _mm512_add_epi64(i, one)), mask);
376                    let a = _mm512_popcnt_epi64(lo);
377                    let b = _mm512_popcnt_epi64(hi);
378                    let c = _mm512_popcnt_epi64(_mm512_and_si512(lo, hi));
379                    let p01 = gather_avx512::<8>(base, _mm512_add_epi64(i, _mm512_set1_epi64(2)));
380                    let p23 = gather_avx512::<8>(base, _mm512_add_epi64(i, _mm512_set1_epi64(3)));
381                    let r0 = _mm512_add_epi64(
382                        _mm512_and_si512(p01, _mm512_set1_epi64(u32::MAX as i64)),
383                        _mm512_add_epi64(_mm512_sub_epi64(_mm512_sub_epi64(offset, a), b), c),
384                    );
385                    let r1 = _mm512_add_epi64(
386                        _mm512_srlv_epi64(p01, _mm512_set1_epi64(32)),
387                        _mm512_sub_epi64(a, c),
388                    );
389                    let r2 = _mm512_add_epi64(
390                        _mm512_and_si512(p23, _mm512_set1_epi64(u32::MAX as i64)),
391                        _mm512_sub_epi64(b, c),
392                    );
393                    let r3 = _mm512_sub_epi64(_mm512_sub_epi64(_mm512_sub_epi64(pos, r0), r1), r2);
394                    [r0, r1, r2, r3]
395                };
396                let l = ranks(start);
397                let r = ranks(end);
398                let low_count =
399                    _mm512_sub_epi64(_mm512_add_epi64(r[0], r[1]), _mm512_add_epi64(l[0], l[1]));
400                let high = greater_avx512(key, _mm512_sub_epi64(low_count, one));
401                key = _mm512_sub_epi64(key, _mm512_and_si512(high, low_count));
402                let l0 = blend!(high, l[0], l[2]);
403                let r0 = blend!(high, r[0], r[2]);
404                let l1 = blend!(high, l[1], l[3]);
405                let r1 = blend!(high, r[1], r[3]);
406                let count0 = _mm512_sub_epi64(r0, l0);
407                let low = greater_avx512(key, _mm512_sub_epi64(count0, one));
408                key = _mm512_sub_epi64(key, _mm512_and_si512(low, count0));
409                let base0 = blend!(
410                    high,
411                    _mm512_set1_epi64(layer.starts[0] as i64),
412                    _mm512_set1_epi64(layer.starts[2] as i64)
413                );
414                let base1 = blend!(
415                    high,
416                    _mm512_set1_epi64(layer.starts[1] as i64),
417                    _mm512_set1_epi64(layer.starts[3] as i64)
418                );
419                let offset = blend!(low, base0, base1);
420                start = _mm512_add_epi64(offset, blend!(low, l0, l1));
421                end = _mm512_add_epi64(offset, blend!(low, r0, r1));
422                code = _mm512_or_si512(
423                    _mm512_sllv_epi64(code, _mm512_set1_epi64(2)),
424                    _mm512_or_si512(
425                        _mm512_and_si512(high, _mm512_set1_epi64(2)),
426                        _mm512_and_si512(low, one),
427                    ),
428                );
429            }
430            _mm512_storeu_si512(result.as_mut_ptr().cast(), code);
431            for (i, s) in chunk.iter_mut().enumerate() {
432                s[3] = result[i] as usize;
433            }
434        }
435    }
436}
437
438#[derive(Debug, Clone)]
439pub struct WaveletMatrix<T> {
440    len: usize,
441    bit_length: usize,
442    zeros: Vec<usize>,
443    bit_vectors: Vec<BitVector>,
444    // Binary layers preserve callback coordinates and support weighted queries.
445    quad_vectors: Vec<WaveletMatrixQuadVector>,
446    compress: VecCompress<T>,
447    #[cfg(target_arch = "x86_64")]
448    backend: super::SimdBackend,
449}
450
451impl<T> WaveletMatrix<T>
452where
453    T: Ord + Clone,
454{
455    pub fn new(v: Vec<T>) -> Self {
456        if v.len() <= u32::MAX as usize {
457            #[cfg(target_arch = "x86_64")]
458            let backend = super::simd_backend();
459            Self::from_values(
460                v,
461                |i| i as u32,
462                |i| i as usize,
463                |indices, d| {
464                    #[cfg(target_arch = "x86_64")]
465                    if backend != super::SimdBackend::Scalar && is_x86_feature_detected!("avx2") {
466                        // SAFETY: AVX2 is available.
467                        return unsafe { simd::pack_words(indices, d) };
468                    }
469                    Self::pack_words(indices, |i| (i >> d) & 1 != 0)
470                },
471                |indices, words, zeros, next| {
472                    #[cfg(target_arch = "x86_64")]
473                    if backend != super::SimdBackend::Scalar && is_x86_feature_detected!("avx2") {
474                        // SAFETY: AVX2 is available, and the partition buffers have equal length.
475                        unsafe { simd::partition_avx2(indices, words, zeros, next) };
476                        return;
477                    }
478                    Self::partition(indices, words, zeros, next);
479                },
480            )
481        } else {
482            Self::from_values(
483                v,
484                |i| i,
485                |i| i,
486                |indices, d| Self::pack_words(indices, |i| (i >> d) & 1 != 0),
487                Self::partition,
488            )
489        }
490    }
491
492    fn pack_words<I: Copy>(indices: &[I], bit: impl Fn(I) -> bool) -> Vec<u64> {
493        indices
494            .chunks(64)
495            .map(|chunk| {
496                chunk
497                    .iter()
498                    .enumerate()
499                    .fold(0, |word, (i, &index)| word | ((bit(index) as u64) << i))
500            })
501            .collect()
502    }
503
504    fn partition<I: Copy>(indices: &[I], words: &[u64], mut one: usize, next: &mut [I]) {
505        let mut zero = 0;
506        for (chunk, &word) in indices.chunks(64).zip(words) {
507            if word == 0 {
508                next[zero..zero + chunk.len()].copy_from_slice(chunk);
509                zero += chunk.len();
510            } else if word == u64::MAX {
511                next[one..one + chunk.len()].copy_from_slice(chunk);
512                one += chunk.len();
513            } else {
514                for (i, &index) in chunk.iter().enumerate() {
515                    let bit = (word >> i) & 1 != 0;
516                    next[if bit { one } else { zero }] = index;
517                    zero += !bit as usize;
518                    one += bit as usize;
519                }
520            }
521        }
522    }
523
524    fn from_values<I: Copy>(
525        v: Vec<T>,
526        code: impl Fn(usize) -> I,
527        index: impl Fn(I) -> usize,
528        pack: impl Fn(&[I], usize) -> Vec<u64>,
529        partition: impl Fn(&[I], &[u64], usize, &mut [I]),
530    ) -> Self {
531        let len = v.len();
532        let mut sorted: Vec<_> = v
533            .into_iter()
534            .enumerate()
535            .map(|(i, value)| (value, code(i)))
536            .collect();
537        sorted.sort_unstable_by(|a, b| a.0.cmp(&b.0));
538        let mut values = Vec::with_capacity(len);
539        let mut indices = vec![code(0); len];
540        for (value, i) in sorted {
541            if values.last().is_none_or(|last| last != &value) {
542                values.push(value);
543            }
544            indices[index(i)] = code(values.len() - 1);
545        }
546        let compress = VecCompress::from_sorted_unique(values);
547        let bit_length = usize::BITS as usize - compress.size().leading_zeros() as usize;
548        let mut bit_vectors = Vec::with_capacity(bit_length);
549        let mut zeros = Vec::with_capacity(bit_length);
550        let quad_bits =
551            usize::BITS as usize - compress.size().saturating_sub(1).leading_zeros() as usize;
552        let mut quad_vectors = Vec::with_capacity(quad_bits.div_ceil(2));
553        let mut next = indices.clone();
554        for d in (0..bit_length).rev() {
555            let words = pack(&indices, d);
556            if len <= u32::MAX as usize && d < quad_bits && (d % 2 == 1 || d + 1 == quad_bits) {
557                if d % 2 == 1 {
558                    let low = pack(&indices, d - 1);
559                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&low, Some(&words), len));
560                } else {
561                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&words, None, len));
562                }
563            }
564            let bits = BitVector::from_words(&words, len);
565            let zero_count = bits.rank0(len);
566            if d == 0 {
567                zeros.push(zero_count);
568                bit_vectors.push(bits);
569                break;
570            }
571            partition(&indices, &words, zero_count, &mut next);
572            zeros.push(zero_count);
573            bit_vectors.push(bits);
574            mem::swap(&mut indices, &mut next);
575        }
576        Self {
577            len,
578            bit_length,
579            zeros,
580            bit_vectors,
581            quad_vectors,
582            compress,
583            #[cfg(target_arch = "x86_64")]
584            backend: match super::simd_backend() {
585                super::SimdBackend::Avx512 if !is_x86_feature_detected!("avx512vpopcntdq") => {
586                    super::SimdBackend::Avx2
587                }
588                backend => backend,
589            },
590        }
591    }
592
593    pub fn new_with_init<F>(v: Vec<T>, mut f: F) -> Self
594    where
595        F: FnMut(usize, usize, T),
596    {
597        let this = Self::new(v.clone());
598        if !this.quad_vectors.is_empty() {
599            let bits = usize::BITS as usize
600                - this.compress.size().saturating_sub(1).leading_zeros() as usize;
601            for (mut k, value) in v.into_iter().enumerate() {
602                if this.bit_length > bits {
603                    f(this.bit_length - 1, k, value.clone());
604                }
605                for (level, vector) in this.quad_vectors.iter().enumerate() {
606                    let d = (this.quad_vectors.len() - level - 1) * 2;
607                    let (digit, rank) = vector.access_rank(k);
608                    if d + 1 < bits {
609                        let block = &vector.blocks[k / 64];
610                        let high = block.rank[2] as usize
611                            + block.rank[3] as usize
612                            + (block.hi & !(u64::MAX << (k % 64))).count_ones() as usize;
613                        let middle = if digit & 2 == 0 {
614                            k - high
615                        } else {
616                            this.zeros[this.level(d + 1)] + high
617                        };
618                        f(d + 1, middle, value.clone());
619                    }
620                    k = vector.starts[digit] + rank;
621                    f(d, k, value.clone());
622                }
623            }
624            return this;
625        }
626        for (mut k, value) in v.into_iter().enumerate() {
627            for d in (0..this.bit_length).rev() {
628                let level = this.level(d);
629                let (bit, rank1) = this.bit_vectors[level].access_rank1(k);
630                k = if bit {
631                    this.zeros[level] + rank1
632                } else {
633                    k - rank1
634                };
635                f(d, k, value.clone());
636            }
637        }
638        this
639    }
640
641    fn level(&self, d: usize) -> usize {
642        self.bit_length - 1 - d
643    }
644
645    fn rank1(&self, level: usize, k: usize) -> usize {
646        self.bit_vectors[level].rank1(k)
647    }
648
649    fn reorder<U>(&self, level: usize, current: Vec<U>) -> Vec<U> {
650        assert_eq!(current.len(), self.len);
651        let mut next = Vec::with_capacity(self.len);
652        next.resize_with(self.len, MaybeUninit::uninit);
653        let mut zero = 0;
654        let mut one = self.zeros[level];
655        let mut current = current.into_iter();
656        for block in self.bit_vectors[level].blocks() {
657            let count = current.len().min(64);
658            if block.bits == 0 || block.bits == u64::MAX {
659                let offset = if block.bits == 0 { &mut zero } else { &mut one };
660                for (slot, value) in next[*offset..*offset + count]
661                    .iter_mut()
662                    .zip(current.by_ref().take(count))
663                {
664                    slot.write(value);
665                }
666                *offset += count;
667            } else {
668                for (i, value) in current.by_ref().take(count).enumerate() {
669                    let bit = (block.bits >> i) & 1 != 0;
670                    next[if bit { one } else { zero }].write(value);
671                    zero += !bit as usize;
672                    one += bit as usize;
673                }
674            }
675        }
676        // SAFETY: the partition counts fill every slot once, and `MaybeUninit<U>` has `U`'s layout.
677        unsafe {
678            let mut next = mem::ManuallyDrop::new(next);
679            Vec::from_raw_parts(next.as_mut_ptr().cast(), next.len(), next.capacity())
680        }
681    }
682
683    fn range_by_index(&self, idx: usize, mut range: Range<usize>) -> Range<usize> {
684        if !self.quad_vectors.is_empty() {
685            for (level, vector) in self.quad_vectors.iter().enumerate() {
686                if range.is_empty() {
687                    break;
688                }
689                let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
690                range = vector.starts[digit] + vector.rank(digit, range.start)
691                    ..vector.starts[digit] + vector.rank(digit, range.end);
692            }
693            return range;
694        }
695        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
696            if range.is_empty() {
697                break;
698            }
699            let level = self.level(d);
700            let start1 = self.rank1(level, range.start);
701            let end1 = self.rank1(level, range.end);
702            if ((idx >> d) & 1) != 0 {
703                range.start = self.zeros[level] + start1;
704                range.end = self.zeros[level] + end1;
705            } else {
706                range.start -= start1;
707                range.end -= end1;
708            }
709        }
710        range
711    }
712
713    fn batch<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
714        #[cfg(target_arch = "x86_64")]
715        if !cfg!(target_feature = "popcnt") && is_x86_feature_detected!("popcnt") {
716            // SAFETY: POPCNT is checked above.
717            unsafe {
718                self.batch_popcnt::<OP>(states, count);
719            }
720            return;
721        }
722        self.batch_inner::<OP>(states, count);
723    }
724
725    #[cfg(target_arch = "x86_64")]
726    #[target_feature(enable = "popcnt")]
727    unsafe fn batch_popcnt<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
728        self.batch_inner::<OP>(states, count);
729    }
730
731    #[inline(always)]
732    fn batch_inner<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
733        if OP == RANK_LESSTHAN && self.compress.size() <= 1 {
734            for state in &mut states[..count] {
735                state[3] = if state[2] == 0 {
736                    0
737                } else {
738                    state[1] - state[0]
739                };
740            }
741            return;
742        }
743        if !self.quad_vectors.is_empty() && matches!(OP, ACCESS | RANK | QUANTILE) {
744            #[cfg(target_arch = "x86_64")]
745            if OP == QUANTILE && count >= 8 && self.backend == super::SimdBackend::Avx512 {
746                // SAFETY: checked batch ranges stay within the quad blocks after each
747                // stable partition. Construction restricts quad counters to u32, and
748                // the cached backend includes AVX512F and VPOPCNTDQ.
749                unsafe {
750                    simd::quad_avx512(&self.quad_vectors, &mut states[..count.next_multiple_of(8)]);
751                }
752                return;
753            }
754
755            for (level, vector) in self.quad_vectors.iter().enumerate() {
756                let d = (self.quad_vectors.len() - level - 1) * 2;
757                #[cfg(target_arch = "x86_64")]
758                let next = self
759                    .quad_vectors
760                    .get(level + 1)
761                    .filter(|_| self.len >= 1048576 || (OP == ACCESS && self.len >= 262144));
762                for state in &mut states[..count] {
763                    if OP == ACCESS {
764                        let (digit, rank) = vector.access_rank(state[0]);
765                        state[0] = vector.starts[digit] + rank;
766                        state[3] = state[3] * 4 + digit;
767                    } else if OP == RANK {
768                        let digit = (state[2] >> d) & 3;
769                        state[0] = vector.starts[digit] + vector.rank(digit, state[0]);
770                        state[1] = vector.starts[digit] + vector.rank(digit, state[1]);
771                    } else {
772                        let start = vector.ranks(state[0]);
773                        let end = vector.ranks(state[1]);
774                        let prefix = [
775                            0,
776                            end[0] - start[0],
777                            end[0] + end[1] - start[0] - start[1],
778                            state[1] - state[0] - (end[3] - start[3]),
779                        ];
780                        let digit = (state[2] >= prefix[1]) as usize
781                            + (state[2] >= prefix[2]) as usize
782                            + (state[2] >= prefix[3]) as usize;
783                        state[2] -= prefix[digit];
784                        state[0] = vector.starts[digit] + start[digit];
785                        state[1] = vector.starts[digit] + end[digit];
786                        state[3] = state[3] * 4 + digit;
787                    }
788                    #[cfg(target_arch = "x86_64")]
789                    if let Some(next) = next {
790                        // SAFETY: stable partitions keep both endpoints within the next vector.
791                        unsafe {
792                            std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
793                                next.blocks.as_ptr().add(state[0] / 64).cast(),
794                            );
795                            if OP != ACCESS {
796                                std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
797                                    next.blocks.as_ptr().add(state[1] / 64).cast(),
798                                );
799                            }
800                        }
801                    }
802                }
803            }
804            return;
805        }
806        let first = (matches!(OP, ACCESS | RANK | QUANTILE)
807            && self.compress.size().is_power_of_two()) as usize;
808        #[cfg(target_arch = "x86_64")]
809        if OP == RANK_LESSTHAN && count >= 8 && self.len <= u32::MAX as usize {
810            // SAFETY: the caller checked each initial position. Rank transitions remain in
811            // 0..=len, including the sentinel block. Padding uses position zero. The
812            // backend was detected at construction, and x86-64 blocks contain two u64s.
813            unsafe {
814                match self.backend {
815                    super::SimdBackend::Avx512 => {
816                        simd::rank_lessthan_avx512(
817                            &self.bit_vectors[first..],
818                            &self.zeros[first..],
819                            &mut states[..count.next_multiple_of(8)],
820                        );
821                        return;
822                    }
823                    super::SimdBackend::Avx2 if self.len < 1048576 => {
824                        simd::rank_lessthan_avx2(
825                            &self.bit_vectors[first..],
826                            &self.zeros[first..],
827                            &mut states[..count.next_multiple_of(4)],
828                        );
829                        return;
830                    }
831                    _ => {}
832                }
833            }
834        }
835        for d in (0..self.bit_length - first).rev() {
836            let level = self.level(d);
837            for state in &mut states[..count] {
838                let (bit, start1) = self.bit_vectors[level].access_rank1(state[0]);
839                let start0 = state[0] - start1;
840                let end1 = if OP == ACCESS {
841                    0
842                } else {
843                    self.rank1(level, state[1])
844                };
845                let end0 = if OP == ACCESS { 0 } else { state[1] - end1 };
846                let count0 = if OP == ACCESS { 0 } else { end0 - start0 };
847                let bit = match OP {
848                    ACCESS => bit,
849                    RANK | RANK_LESSTHAN => (state[2] >> d) & 1 != 0,
850                    _ => state[2] >= count0,
851                };
852                state[0] = if bit {
853                    self.zeros[level] + start1
854                } else {
855                    start0
856                };
857                if OP != ACCESS {
858                    state[1] = if bit { self.zeros[level] + end1 } else { end0 };
859                }
860                if OP == ACCESS || OP == QUANTILE {
861                    state[3] |= (bit as usize) << d;
862                }
863                if OP == QUANTILE {
864                    state[2] -= if bit { count0 } else { 0 };
865                }
866                if OP == RANK_LESSTHAN {
867                    state[3] += if bit { count0 } else { 0 };
868                }
869            }
870        }
871    }
872
873    /// get k-th value
874    pub fn access(&self, mut k: usize) -> T {
875        if !self.quad_vectors.is_empty() {
876            let mut index = 0;
877            for vector in &self.quad_vectors {
878                let (digit, rank) = vector.access_rank(k);
879                index = index * 4 + digit;
880                k = vector.starts[digit] + rank;
881            }
882            return self.compress.values()[index].clone();
883        }
884        let mut idx = 0;
885        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
886            let level = self.level(d);
887            let (bit, rank1) = self.bit_vectors[level].access_rank1(k);
888            idx |= (bit as usize) << d;
889            k = if bit {
890                self.zeros[level] + rank1
891            } else {
892                k - rank1
893            };
894        }
895        self.compress.values()[idx].clone()
896    }
897
898    /// Returns the values at `indices` in input order.
899    pub fn access_batch(&self, indices: impl IntoIterator<Item = usize>) -> Vec<T> {
900        let indices: Vec<_> = indices.into_iter().collect();
901        let mut result = Vec::with_capacity(indices.len());
902        for indices in indices.chunks(16) {
903            let mut states = [[0; 4]; 16];
904            for (state, &index) in states.iter_mut().zip(indices) {
905                assert!(index < self.len);
906                state[0] = index;
907            }
908            self.batch::<ACCESS>(&mut states, indices.len());
909            result.extend(
910                states[..indices.len()]
911                    .iter()
912                    .map(|state| self.compress.values()[state[3]].clone()),
913            );
914        }
915        result
916    }
917
918    /// the number of val in range
919    pub fn rank(&self, val: T, range: Range<usize>) -> usize {
920        match self.compress.index_exact(&val) {
921            Some(idx) => self.range_by_index(idx, range).len(),
922            None => 0,
923        }
924    }
925
926    /// Returns the number of exact matches for each `(value, range)` query.
927    pub fn rank_batch(&self, queries: impl IntoIterator<Item = (T, Range<usize>)>) -> Vec<usize> {
928        let queries: Vec<_> = queries.into_iter().collect();
929        let mut result = Vec::with_capacity(queries.len());
930        for queries in queries.chunks(16) {
931            let mut states = [[0; 4]; 16];
932            for (state, (value, range)) in states.iter_mut().zip(queries) {
933                assert!(range.start <= range.end && range.end <= self.len);
934                if let Some(index) = self.compress.index_exact(value) {
935                    *state = [range.start, range.end, index, 0];
936                }
937            }
938            self.batch::<RANK>(&mut states, queries.len());
939            result.extend(
940                states[..queries.len()]
941                    .iter()
942                    .map(|state| state[1] - state[0]),
943            );
944        }
945        result
946    }
947
948    /// index of k-th val
949    pub fn select(&self, val: T, k: usize) -> Option<usize> {
950        let idx = self.compress.index_exact(&val)?;
951        let range = self.range_by_index(idx, 0..self.len);
952        if range.len() <= k {
953            return None;
954        }
955        let mut i = range.start + k;
956        if !self.quad_vectors.is_empty() {
957            for (level, vector) in self.quad_vectors.iter().enumerate().rev() {
958                let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
959                i = vector.select(digit, i - vector.starts[digit]);
960            }
961            return Some(i);
962        }
963        for level in (self.compress.size().is_power_of_two() as usize..self.bit_length).rev() {
964            if i >= self.zeros[level] {
965                i = self.bit_vectors[level]
966                    .select1(i - self.zeros[level])
967                    .unwrap();
968            } else {
969                i = self.bit_vectors[level].select0(i).unwrap();
970            }
971        }
972        Some(i)
973    }
974
975    /// get k-th smallest value in range
976    pub fn quantile(&self, mut range: Range<usize>, mut k: usize) -> T {
977        if !self.quad_vectors.is_empty() {
978            return self.quad_quantile(range, k, 0, 0);
979        }
980        let mut idx = 0;
981        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
982            let level = self.level(d);
983            let start1 = self.rank1(level, range.start);
984            let end1 = self.rank1(level, range.end);
985            let start0 = range.start - start1;
986            let end0 = range.end - end1;
987            let z = end0 - start0;
988            let bit = z <= k;
989            k -= if bit { z } else { 0 };
990            idx |= (bit as usize) << d;
991            range.start = if bit {
992                self.zeros[level] + start1
993            } else {
994                start0
995            };
996            range.end = if bit { self.zeros[level] + end1 } else { end0 };
997        }
998        self.compress.values()[idx].clone()
999    }
1000
1001    #[inline(always)]
1002    fn quad_quantile(
1003        &self,
1004        mut range: Range<usize>,
1005        mut k: usize,
1006        level: usize,
1007        mut index: usize,
1008    ) -> T {
1009        for vector in &self.quad_vectors[level..] {
1010            let start = vector.ranks(range.start);
1011            let end = vector.ranks(range.end);
1012            let mut digit = 0;
1013            while digit < 3 && k >= end[digit] - start[digit] {
1014                k -= end[digit] - start[digit];
1015                digit += 1;
1016            }
1017            index = index * 4 + digit;
1018            range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1019        }
1020        self.compress.values()[index].clone()
1021    }
1022
1023    pub fn quantile_batch(
1024        &self,
1025        queries: impl IntoIterator<Item = (Range<usize>, usize)>,
1026    ) -> Vec<T> {
1027        let queries: Vec<_> = queries.into_iter().collect();
1028        let mut result = Vec::with_capacity(queries.len());
1029        for queries in queries.chunks(16) {
1030            let mut states = [[0; 4]; 16];
1031            for (state, (range, k)) in states.iter_mut().zip(queries) {
1032                assert!(range.start <= range.end && range.end <= self.len && *k < range.len());
1033                *state = [range.start, range.end, *k, 0];
1034            }
1035            if let [(range, k)] = queries {
1036                result.push(self.quantile(range.clone(), *k));
1037                continue;
1038            }
1039            self.batch::<QUANTILE>(&mut states, queries.len());
1040            result.extend(
1041                states[..queries.len()]
1042                    .iter()
1043                    .map(|state| self.compress.values()[state[3]].clone()),
1044            );
1045        }
1046        result
1047    }
1048
1049    /// Returns the requested order statistics of one range. `ranks` must be sorted
1050    /// in nondecreasing order and every rank must be less than the range length.
1051    pub fn quantiles_sorted(&self, range: Range<usize>, ranks: &[usize]) -> Vec<T> {
1052        assert!(range.start <= range.end && range.end <= self.len);
1053        assert!(ranks.windows(2).all(|pair| pair[0] <= pair[1]));
1054        assert!(ranks.last().is_none_or(|&k| k < range.len()));
1055        if ranks.is_empty() {
1056            return Vec::new();
1057        }
1058        if ranks[0] == ranks[ranks.len() - 1] {
1059            return vec![self.quantile(range, ranks[0]); ranks.len()];
1060        }
1061        let use_batch = ranks.len() < 1024;
1062        #[cfg(target_arch = "x86_64")]
1063        let use_batch = use_batch || self.backend == super::SimdBackend::Avx512;
1064        // Sparse requests amortize less of the shared traversal's branching work.
1065        if use_batch && ranks.len() < self.compress.size().div_ceil(16) {
1066            return self.quantile_batch(ranks.iter().map(|&k| (range.clone(), k)));
1067        }
1068        let mut result = Vec::with_capacity(ranks.len());
1069        if !self.quad_vectors.is_empty() {
1070            let mut stack = vec![(range, 0, 0, 0, ranks)];
1071            while let Some((range, base, index, level, ranks)) = stack.pop() {
1072                if level == self.quad_vectors.len() {
1073                    result.resize(
1074                        result.len() + ranks.len(),
1075                        self.compress.values()[index].clone(),
1076                    );
1077                    continue;
1078                }
1079                if ranks.len() == 1 {
1080                    result.push(self.quad_quantile(range, ranks[0] - base, level, index));
1081                    continue;
1082                }
1083                let vector = &self.quad_vectors[level];
1084                let start = vector.ranks(range.start);
1085                let end = vector.ranks(range.end);
1086                let mut split = base + range.len();
1087                let mut requested = ranks.len();
1088                for digit in (0..4).rev() {
1089                    split -= end[digit] - start[digit];
1090                    let boundary = ranks[..requested].partition_point(|&k| k < split);
1091                    if boundary < requested {
1092                        stack.push((
1093                            vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit],
1094                            split,
1095                            index * 4 + digit,
1096                            level + 1,
1097                            &ranks[boundary..requested],
1098                        ));
1099                    }
1100                    requested = boundary;
1101                }
1102            }
1103            return result;
1104        }
1105        let mut stack = vec![(range, 0, 0, self.bit_length, ranks)];
1106        while let Some((range, base, index, bits, ranks)) = stack.pop() {
1107            if bits == 0 {
1108                result.resize(
1109                    result.len() + ranks.len(),
1110                    self.compress.values()[index].clone(),
1111                );
1112                continue;
1113            }
1114            let d = bits - 1;
1115            let level = self.level(d);
1116            let start1 = self.rank1(level, range.start);
1117            let end1 = self.rank1(level, range.end);
1118            let start0 = range.start - start1;
1119            let end0 = range.end - end1;
1120            let split = base + end0 - start0;
1121            let boundary = ranks.partition_point(|&k| k < split);
1122            if boundary < ranks.len() {
1123                stack.push((
1124                    self.zeros[level] + start1..self.zeros[level] + end1,
1125                    split,
1126                    index | (1 << d),
1127                    d,
1128                    &ranks[boundary..],
1129                ));
1130            }
1131            if boundary != 0 {
1132                stack.push((start0..end0, base, index, d, &ranks[..boundary]));
1133            }
1134        }
1135        result
1136    }
1137
1138    /// get k-th smallest value out of range
1139    pub fn quantile_outer(&self, mut range: Range<usize>, mut k: usize) -> T {
1140        if !self.quad_vectors.is_empty() {
1141            let mut outer = 0..self.len;
1142            let mut index = 0;
1143            for vector in &self.quad_vectors {
1144                let start = vector.ranks(range.start);
1145                let end = vector.ranks(range.end);
1146                let outer_start = vector.ranks(outer.start);
1147                let outer_end = vector.ranks(outer.end);
1148                let mut digit = 0;
1149                while digit < 3 {
1150                    let count = outer_end[digit] - outer_start[digit] - (end[digit] - start[digit]);
1151                    if k < count {
1152                        break;
1153                    }
1154                    k -= count;
1155                    digit += 1;
1156                }
1157                index = index * 4 + digit;
1158                range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1159                outer = vector.starts[digit] + outer_start[digit]
1160                    ..vector.starts[digit] + outer_end[digit];
1161            }
1162            return self.compress.values()[index].clone();
1163        }
1164        let mut idx = 0;
1165        let mut orange = 0..self.len;
1166        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
1167            let level = self.level(d);
1168            let range_start1 = self.rank1(level, range.start);
1169            let range_end1 = self.rank1(level, range.end);
1170            let outer_start1 = self.rank1(level, orange.start);
1171            let outer_end1 = self.rank1(level, orange.end);
1172            let range_start0 = range.start - range_start1;
1173            let range_end0 = range.end - range_end1;
1174            let outer_start0 = orange.start - outer_start1;
1175            let outer_end0 = orange.end - outer_end1;
1176            let z = (outer_end0 - outer_start0) - (range_end0 - range_start0);
1177            if z <= k {
1178                k -= z;
1179                idx |= 1 << d;
1180                range.start = self.zeros[level] + range_start1;
1181                range.end = self.zeros[level] + range_end1;
1182                orange.start = self.zeros[level] + outer_start1;
1183                orange.end = self.zeros[level] + outer_end1;
1184            } else {
1185                range.start = range_start0;
1186                range.end = range_end0;
1187                orange.start = outer_start0;
1188                orange.end = outer_end0;
1189            }
1190        }
1191        self.compress.values()[idx].clone()
1192    }
1193
1194    /// the number of value less than val in range
1195    pub fn rank_lessthan(&self, val: T, range: Range<usize>) -> usize {
1196        let index = self.compress.index_lower_bound(&val);
1197        if index == self.compress.size() {
1198            return range.len();
1199        }
1200        if self.len >= 262144 && !self.quad_vectors.is_empty() {
1201            self.quad_rank_lessthan_index(index, range, 0)
1202        } else {
1203            let bits = usize::BITS as usize
1204                - self.compress.size().saturating_sub(1).leading_zeros() as usize;
1205            self.rank_lessthan_index(index, range, bits)
1206        }
1207    }
1208
1209    /// Returns the number of values below each query's threshold.
1210    pub fn rank_lessthan_batch(
1211        &self,
1212        queries: impl IntoIterator<Item = (T, Range<usize>)>,
1213    ) -> Vec<usize> {
1214        let queries: Vec<_> = queries.into_iter().collect();
1215        let mut result = Vec::with_capacity(queries.len());
1216        for queries in queries.chunks(16) {
1217            let mut states = [[0; 4]; 16];
1218            for (state, (value, range)) in states.iter_mut().zip(queries) {
1219                assert!(range.start <= range.end && range.end <= self.len);
1220                *state = [
1221                    range.start,
1222                    range.end,
1223                    self.compress.index_lower_bound(value),
1224                    0,
1225                ];
1226            }
1227            self.batch::<RANK_LESSTHAN>(&mut states, queries.len());
1228            result.extend(states[..queries.len()].iter().map(|state| state[3]));
1229        }
1230        result
1231    }
1232
1233    fn quad_rank_lessthan_index(
1234        &self,
1235        index: usize,
1236        mut range: Range<usize>,
1237        start_level: usize,
1238    ) -> usize {
1239        let mut result = 0;
1240        for (level, vector) in self.quad_vectors.iter().enumerate().skip(start_level) {
1241            let d = (self.quad_vectors.len() - level - 1) * 2;
1242            if d + 2 <= index.trailing_zeros() as usize {
1243                break;
1244            }
1245            let digit = (index >> d) & 3;
1246            let ranks = |position: usize| {
1247                let block = &vector.blocks[position / 64];
1248                let mask = !(u64::MAX << (position % 64));
1249                let low = if digit & 1 == 0 { !block.lo } else { block.lo };
1250                let high = if digit & 2 == 0 { !block.hi } else { block.hi };
1251                let exact = block.rank[digit] as usize + (low & high & mask).count_ones() as usize;
1252                let less = match digit {
1253                    0 => 0,
1254                    1 => {
1255                        block.rank[0] as usize
1256                            + (!(block.lo | block.hi) & mask).count_ones() as usize
1257                    }
1258                    2 => {
1259                        block.rank[0] as usize
1260                            + block.rank[1] as usize
1261                            + (!block.hi & mask).count_ones() as usize
1262                    }
1263                    _ => position - exact,
1264                };
1265                (exact, less)
1266            };
1267            let (start, start_less) = ranks(range.start);
1268            let (end, end_less) = ranks(range.end);
1269            result += end_less - start_less;
1270            range = vector.starts[digit] + start..vector.starts[digit] + end;
1271        }
1272        result
1273    }
1274
1275    fn rank_lessthan_index(&self, idx: usize, mut range: Range<usize>, bits: usize) -> usize {
1276        let mut res = 0;
1277        for d in (idx.trailing_zeros() as usize..bits).rev() {
1278            let level = self.level(d);
1279            let start1 = self.rank1(level, range.start);
1280            let end1 = self.rank1(level, range.end);
1281            let bit = (idx >> d) & 1 != 0;
1282            let start0 = range.start - start1;
1283            let end0 = range.end - end1;
1284            res += if bit { end0 - start0 } else { 0 };
1285            range.start = if bit {
1286                self.zeros[level] + start1
1287            } else {
1288                start0
1289            };
1290            range.end = if bit { self.zeros[level] + end1 } else { end0 };
1291        }
1292        res
1293    }
1294
1295    /// the number of valrange in range
1296    pub fn rank_range(&self, valrange: Range<T>, mut range: Range<usize>) -> usize {
1297        let lower = self.compress.index_lower_bound(&valrange.start);
1298        let upper = self.compress.index_lower_bound(&valrange.end);
1299        if lower >= upper {
1300            return 0;
1301        }
1302        if !self.quad_vectors.is_empty() {
1303            if upper == self.compress.size() {
1304                return range.len() - self.quad_rank_lessthan_index(lower, range, 0);
1305            }
1306            for (level, vector) in self.quad_vectors.iter().enumerate() {
1307                let d = (self.quad_vectors.len() - level - 1) * 2;
1308                let low = (lower >> d) & 3;
1309                let high = (upper >> d) & 3;
1310                if low == high {
1311                    range = vector.starts[low] + vector.rank(low, range.start)
1312                        ..vector.starts[low] + vector.rank(low, range.end);
1313                    continue;
1314                }
1315                let start = vector.ranks(range.start);
1316                let end = vector.ranks(range.end);
1317                let count: usize = (low..high).map(|digit| end[digit] - start[digit]).sum();
1318                return count
1319                    - self.quad_rank_lessthan_index(
1320                        lower,
1321                        vector.starts[low] + start[low]..vector.starts[low] + end[low],
1322                        level + 1,
1323                    )
1324                    + self.quad_rank_lessthan_index(
1325                        upper,
1326                        vector.starts[high] + start[high]..vector.starts[high] + end[high],
1327                        level + 1,
1328                    );
1329            }
1330            return 0;
1331        }
1332        for d in (0..self.bit_length).rev() {
1333            let level = self.level(d);
1334            let start1 = self.rank1(level, range.start);
1335            let end1 = self.rank1(level, range.end);
1336            let zero_range = range.start - start1..range.end - end1;
1337            let one_range = self.zeros[level] + start1..self.zeros[level] + end1;
1338            if ((lower ^ upper) >> d) & 1 != 0 {
1339                return zero_range.len() - self.rank_lessthan_index(lower, zero_range, d)
1340                    + self.rank_lessthan_index(upper, one_range, d);
1341            }
1342            range = if (lower >> d) & 1 != 0 {
1343                one_range
1344            } else {
1345                zero_range
1346            };
1347        }
1348        0
1349    }
1350
1351    /// Returns each query's count in a half-open value range.
1352    pub fn rank_range_batch(
1353        &self,
1354        queries: impl IntoIterator<Item = (Range<T>, Range<usize>)>,
1355    ) -> Vec<usize> {
1356        let queries: Vec<_> = queries.into_iter().collect();
1357        let mut result = Vec::with_capacity(queries.len());
1358        for queries in queries.chunks(8) {
1359            let mut states = [[0; 4]; 16];
1360            for (i, (values, range)) in queries.iter().enumerate() {
1361                assert!(range.start <= range.end && range.end <= self.len);
1362                let lower = self.compress.index_lower_bound(&values.start);
1363                let upper = self.compress.index_lower_bound(&values.end);
1364                if lower < upper {
1365                    states[i * 2] = [range.start, range.end, lower, 0];
1366                    states[i * 2 + 1] = [range.start, range.end, upper, 0];
1367                }
1368            }
1369            self.batch::<RANK_LESSTHAN>(&mut states, queries.len() * 2);
1370            result.extend(
1371                states[..queries.len() * 2]
1372                    .as_chunks::<2>()
1373                    .0
1374                    .iter()
1375                    .map(|pair| pair[1][3] - pair[0][3]),
1376            );
1377        }
1378        result
1379    }
1380
1381    pub fn query_less_than<F>(&self, val: T, mut range: Range<usize>, mut f: F)
1382    where
1383        F: FnMut(usize, Range<usize>),
1384    {
1385        let idx = self.compress.index_lower_bound(&val);
1386        if !self.quad_vectors.is_empty() {
1387            if idx == self.compress.size() && idx.is_power_of_two() {
1388                f(self.bit_length - 1, range);
1389                return;
1390            }
1391            for (level, vector) in self.quad_vectors.iter().enumerate().take(
1392                self.quad_vectors
1393                    .len()
1394                    .saturating_sub(idx.trailing_zeros() as usize / 2),
1395            ) {
1396                let d = (self.quad_vectors.len() - level - 1) * 2;
1397                let digit = (idx >> d) & 3;
1398                let start = vector.ranks(range.start);
1399                let end = vector.ranks(range.end);
1400                if digit & 2 != 0 {
1401                    f(d + 1, start[0] + start[1]..end[0] + end[1]);
1402                }
1403                if digit & 1 != 0 {
1404                    let zero = digit & 2;
1405                    f(
1406                        d,
1407                        vector.starts[zero] + start[zero]..vector.starts[zero] + end[zero],
1408                    );
1409                }
1410                range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1411            }
1412            return;
1413        }
1414        for d in (idx.trailing_zeros() as usize..self.bit_length).rev() {
1415            let level = self.level(d);
1416            let start1 = self.rank1(level, range.start);
1417            let end1 = self.rank1(level, range.end);
1418            let start0 = range.start - start1;
1419            let end0 = range.end - end1;
1420            if ((idx >> d) & 1) != 0 {
1421                f(d, start0..end0);
1422                range.start = self.zeros[level] + start1;
1423                range.end = self.zeros[level] + end1;
1424            } else {
1425                range.start = start0;
1426                range.end = end0;
1427            }
1428        }
1429    }
1430
1431    pub fn build_fold<M>(&self, weights: &[M::T]) -> WaveletMatrixFold<'_, T, M>
1432    where
1433        M: AbelianGroup,
1434    {
1435        assert_eq!(weights.len(), self.len);
1436        let mut offsets = Vec::with_capacity(self.bit_length);
1437        let mut prefix = Vec::with_capacity(self.zeros.iter().map(|&zero| zero + 1).sum());
1438        let mut current: Vec<M::T> = weights.to_vec();
1439        for level in 0..self.bit_length {
1440            current = self.reorder(level, current);
1441            offsets.push(prefix.len());
1442            let mut acc = M::unit();
1443            prefix.push(acc.clone());
1444            for w in &current[..self.zeros[level]] {
1445                acc = M::operate(&acc, w);
1446                prefix.push(acc.clone());
1447            }
1448        }
1449        WaveletMatrixFold {
1450            wavelet_matrix: self,
1451            prefix,
1452            offsets,
1453        }
1454    }
1455
1456    pub fn build_point_add<M>(&self, weights: &[M::T]) -> WaveletMatrixPointAdd<'_, T, M>
1457    where
1458        M: AbelianGroup,
1459    {
1460        assert_eq!(weights.len(), self.len);
1461        let mut current = weights.to_vec();
1462        let mut bits = Vec::with_capacity(self.bit_length);
1463        for level in 0..self.bit_length {
1464            current = self.reorder(level, current);
1465            bits.push(BinaryIndexedTree::from_slice(&current[..self.zeros[level]]));
1466        }
1467        WaveletMatrixPointAdd {
1468            wavelet_matrix: self,
1469            bits,
1470        }
1471    }
1472}
1473
1474pub struct WaveletMatrixPointAdd<'a, T, M>
1475where
1476    T: Ord + Clone,
1477    M: AbelianGroup,
1478{
1479    wavelet_matrix: &'a WaveletMatrix<T>,
1480    bits: Vec<BinaryIndexedTree<M>>,
1481}
1482
1483impl<'a, T, M> WaveletMatrixPointAdd<'a, T, M>
1484where
1485    T: Ord + Clone,
1486    M: AbelianGroup,
1487{
1488    pub fn update(&mut self, mut index: usize, value: M::T) {
1489        debug_assert!(index < self.wavelet_matrix.len);
1490        for d in (0..self.wavelet_matrix.bit_length).rev() {
1491            let level = self.wavelet_matrix.level(d);
1492            let (bit, rank1) = self.wavelet_matrix.bit_vectors[level].access_rank1(index);
1493            if bit {
1494                index = self.wavelet_matrix.zeros[level] + rank1;
1495            } else {
1496                index -= rank1;
1497                self.bits[level].update(index, value.clone());
1498            }
1499        }
1500    }
1501
1502    pub fn fold_lessthan(&self, value: T, range: Range<usize>) -> M::T {
1503        let mut result = M::unit();
1504        self.wavelet_matrix
1505            .query_less_than(value, range, |d, range| {
1506                M::operate_assign(
1507                    &mut result,
1508                    &self.bits[self.wavelet_matrix.level(d)].fold_abelian(range.start, range.end),
1509                );
1510            });
1511        result
1512    }
1513
1514    pub fn fold_range(&self, values: Range<T>, range: Range<usize>) -> M::T {
1515        let lower = self
1516            .wavelet_matrix
1517            .compress
1518            .index_lower_bound(&values.start);
1519        let upper = self.wavelet_matrix.compress.index_lower_bound(&values.end);
1520        if lower >= upper {
1521            return M::unit();
1522        }
1523        let mut range = range;
1524        for d in (0..self.wavelet_matrix.bit_length).rev() {
1525            let level = self.wavelet_matrix.level(d);
1526            let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1527            let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1528            let start0 = range.start - start1;
1529            let end0 = range.end - end1;
1530            if ((lower >> d) & 1) == ((upper >> d) & 1) {
1531                if ((lower >> d) & 1) == 0 {
1532                    range = start0..end0;
1533                } else {
1534                    range = self.wavelet_matrix.zeros[level] + start1
1535                        ..self.wavelet_matrix.zeros[level] + end1;
1536                }
1537                continue;
1538            }
1539            let zero_range = start0..end0;
1540            let one_range =
1541                self.wavelet_matrix.zeros[level] + start1..self.wavelet_matrix.zeros[level] + end1;
1542            let lower_sum = self.fold_lessthan_index(lower, zero_range.clone(), d);
1543            let upper_sum = self.fold_lessthan_index(upper, one_range, d);
1544            let zero_sum = self.bits[level].fold_abelian(zero_range.start, zero_range.end);
1545            let mut result = M::rinv_operate(&zero_sum, &lower_sum);
1546            M::operate_assign(&mut result, &upper_sum);
1547            return result;
1548        }
1549        M::unit()
1550    }
1551
1552    fn fold_lessthan_index(&self, idx: usize, mut range: Range<usize>, bits: usize) -> M::T {
1553        let mut result = M::unit();
1554        for d in (idx.trailing_zeros() as usize..bits).rev() {
1555            let level = self.wavelet_matrix.level(d);
1556            let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1557            let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1558            let start0 = range.start - start1;
1559            let end0 = range.end - end1;
1560            if ((idx >> d) & 1) != 0 {
1561                M::operate_assign(&mut result, &self.bits[level].fold_abelian(start0, end0));
1562                range.start = self.wavelet_matrix.zeros[level] + start1;
1563                range.end = self.wavelet_matrix.zeros[level] + end1;
1564            } else {
1565                range.start = start0;
1566                range.end = end0;
1567            }
1568        }
1569        result
1570    }
1571}
1572
1573#[derive(Debug, Clone)]
1574pub struct WaveletMatrixFold<'a, T, M>
1575where
1576    T: Ord + Clone,
1577    M: AbelianGroup,
1578{
1579    wavelet_matrix: &'a WaveletMatrix<T>,
1580    prefix: Vec<M::T>,
1581    offsets: Vec<usize>,
1582}
1583
1584impl<'a, T, M> WaveletMatrixFold<'a, T, M>
1585where
1586    T: Ord + Clone,
1587    M: AbelianGroup,
1588{
1589    pub fn fold_lessthan(&self, val: T, range: Range<usize>) -> M::T {
1590        self.fold_lessthan_with_count(val, range).1
1591    }
1592
1593    pub fn fold_lessthan_with_count(&self, val: T, range: Range<usize>) -> (usize, M::T) {
1594        debug_assert!(range.end <= self.wavelet_matrix.len);
1595        let [result] = self.fold_lessthan_indices_with_count(
1596            [self.wavelet_matrix.compress.index_lower_bound(&val)],
1597            [range],
1598            self.wavelet_matrix.bit_length,
1599        );
1600        result
1601    }
1602
1603    pub fn fold_range(&self, valrange: Range<T>, range: Range<usize>) -> M::T {
1604        self.fold_range_with_count(valrange, range).1
1605    }
1606
1607    pub fn fold_range_with_count(
1608        &self,
1609        valrange: Range<T>,
1610        mut range: Range<usize>,
1611    ) -> (usize, M::T) {
1612        debug_assert!(range.end <= self.wavelet_matrix.len);
1613        let lower = self
1614            .wavelet_matrix
1615            .compress
1616            .index_lower_bound(&valrange.start);
1617        let upper = self
1618            .wavelet_matrix
1619            .compress
1620            .index_lower_bound(&valrange.end);
1621        if lower >= upper {
1622            return (0, M::unit());
1623        }
1624        for d in (0..self.wavelet_matrix.bit_length).rev() {
1625            let level = self.wavelet_matrix.level(d);
1626            let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1627            let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1628            let start0 = range.start - start1;
1629            let end0 = range.end - end1;
1630            if ((lower >> d) & 1) == ((upper >> d) & 1) {
1631                if ((lower >> d) & 1) == 0 {
1632                    range = start0..end0;
1633                } else {
1634                    range = self.wavelet_matrix.zeros[level] + start1
1635                        ..self.wavelet_matrix.zeros[level] + end1;
1636                }
1637                continue;
1638            }
1639            let zero_range = start0..end0;
1640            let one_range =
1641                self.wavelet_matrix.zeros[level] + start1..self.wavelet_matrix.zeros[level] + end1;
1642            let [(lower_count, lower_sum), (upper_count, upper_sum)] = self
1643                .fold_lessthan_indices_with_count(
1644                    [lower, upper],
1645                    [zero_range.clone(), one_range],
1646                    d,
1647                );
1648            let zero_sum = self.range_sum(level, zero_range.clone());
1649            return (
1650                zero_range.len() - lower_count + upper_count,
1651                M::operate(&M::rinv_operate(&zero_sum, &lower_sum), &upper_sum),
1652            );
1653        }
1654        (0, M::unit())
1655    }
1656
1657    #[inline]
1658    fn range_sum(&self, level: usize, range: Range<usize>) -> M::T {
1659        let offset = self.offsets[level];
1660        M::rinv_operate(
1661            &self.prefix[offset + range.end],
1662            &self.prefix[offset + range.start],
1663        )
1664    }
1665
1666    fn fold_lessthan_indices_with_count<const N: usize>(
1667        &self,
1668        indices: [usize; N],
1669        mut ranges: [Range<usize>; N],
1670        bits: usize,
1671    ) -> [(usize, M::T); N] {
1672        let mut results = std::array::from_fn(|_| (0, M::unit()));
1673        let last = indices
1674            .iter()
1675            .map(|index| index.trailing_zeros() as usize)
1676            .min()
1677            .unwrap_or(bits);
1678        for d in (last..bits).rev() {
1679            let level = self.wavelet_matrix.level(d);
1680            for ((&index, range), (count, sum)) in indices.iter().zip(&mut ranges).zip(&mut results)
1681            {
1682                let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1683                let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1684                let start0 = range.start - start1;
1685                let end0 = range.end - end1;
1686                if ((index >> d) & 1) != 0 {
1687                    *count += end0 - start0;
1688                    M::operate_assign(sum, &self.range_sum(level, start0..end0));
1689                    range.start = self.wavelet_matrix.zeros[level] + start1;
1690                    range.end = self.wavelet_matrix.zeros[level] + end1;
1691                } else {
1692                    range.start = start0;
1693                    range.end = end0;
1694                }
1695            }
1696        }
1697        results
1698    }
1699
1700    /// Folds the weights below each query's threshold, traversing the queries together.
1701    pub fn fold_lessthan_batch(
1702        &self,
1703        queries: impl IntoIterator<Item = (T, Range<usize>)>,
1704    ) -> Vec<M::T> {
1705        let queries: Vec<_> = queries.into_iter().collect();
1706        let mut result = Vec::with_capacity(queries.len());
1707        for queries in queries.chunks(16) {
1708            let mut indices = [0; 16];
1709            let mut ranges = std::array::from_fn(|_| 0..0);
1710            for (i, (value, range)) in queries.iter().enumerate() {
1711                assert!(range.start <= range.end && range.end <= self.wavelet_matrix.len);
1712                indices[i] = self.wavelet_matrix.compress.index_lower_bound(value);
1713                ranges[i] = range.clone();
1714            }
1715            result.extend(
1716                self.fold_lessthan_indices_with_count(
1717                    indices,
1718                    ranges,
1719                    self.wavelet_matrix.bit_length,
1720                )
1721                .into_iter()
1722                .take(queries.len())
1723                .map(|(_, sum)| sum),
1724            );
1725        }
1726        result
1727    }
1728}
1729
1730#[cfg(test)]
1731mod tests {
1732    use super::*;
1733    use crate::tools::Xorshift;
1734
1735    #[test]
1736    fn test_wavelet_matrix() {
1737        let mut rng = Xorshift::default();
1738        let mut lengths: Vec<usize> = vec![
1739            0, 1, 2, 3, 7, 8, 15, 16, 17, 63, 64, 65, 127, 128, 129, 255, 256, 257, 1000, 262144,
1740            1048577,
1741        ];
1742        lengths.extend((0..32).map(|_| rng.random(0usize..1024)));
1743        for n in lengths {
1744            let sigma: i64 = if n > 1024 { 4095 } else { rng.random(1..65) };
1745            let values: Vec<_> = (0..n).map(|_| rng.random(0..sigma) * 2).collect();
1746            let ranks: Vec<_> = (0..33)
1747                .map(|i| {
1748                    let left: usize = rng.random(0..=n);
1749                    let right = rng.random(left..=(left + 257).min(n));
1750                    let range = if i == 0 { 0..n } else { left..right };
1751                    (rng.random(-1..=sigma * 2), range)
1752                })
1753                .collect();
1754            let counts: Vec<_> = ranks
1755                .iter()
1756                .map(|(v, r)| values[r.clone()].iter().filter(|&x| x == v).count())
1757                .collect();
1758            let less: Vec<_> = ranks
1759                .iter()
1760                .map(|(v, r)| values[r.clone()].iter().filter(|&x| x < v).count())
1761                .collect();
1762            let positions: Vec<Vec<_>> = ranks
1763                .iter()
1764                .map(|(v, _)| (0..n).filter(|&i| values[i] == *v).collect())
1765                .collect();
1766            let ranges: Vec<_> = ranks
1767                .iter()
1768                .map(|(v, r)| (*v..*v + rng.random(0..=sigma), r.clone()))
1769                .collect();
1770            let range_counts: Vec<_> = ranges
1771                .iter()
1772                .map(|(v, r)| values[r.clone()].iter().filter(|x| v.contains(x)).count())
1773                .collect();
1774            let indices: Vec<_> = (0..33)
1775                .filter(|_| n != 0)
1776                .map(|_| rng.random(0..n))
1777                .collect();
1778            let queries: Vec<_> = indices
1779                .iter()
1780                .enumerate()
1781                .map(|(i, &left)| {
1782                    let right = rng.random(left + 1..=(left + 257).min(n));
1783                    let range = if i == 0 { 0..n } else { left..right };
1784                    let k = rng.random(0..range.len());
1785                    (range, k)
1786                })
1787                .collect();
1788            let expected: Vec<_> = queries
1789                .iter()
1790                .map(|(r, k)| {
1791                    let mut sorted = values[r.clone()].to_vec();
1792                    sorted.sort_unstable();
1793                    sorted[*k]
1794                })
1795                .collect();
1796            let mut sorted = values.clone();
1797            sorted.sort_unstable();
1798            let mut sorted_ranks: Vec<_> = indices.iter().map(|_| rng.random(0..n)).collect();
1799            sorted_ranks.sort_unstable();
1800            let original = WaveletMatrix::new(values.clone());
1801            for binary_only in [false, true] {
1802                let mut wm = original.clone();
1803                if binary_only {
1804                    wm.quad_vectors.clear();
1805                }
1806                for_each_backend(&mut wm, |wm| {
1807                    for (i, ((value, range), positions)) in ranks.iter().zip(&positions).enumerate()
1808                    {
1809                        assert_eq!(wm.rank(*value, range.clone()), counts[i]);
1810                        assert_eq!(wm.rank_lessthan(*value, range.clone()), less[i]);
1811                        assert_eq!(
1812                            wm.rank_range(ranges[i].0.clone(), range.clone()),
1813                            range_counts[i]
1814                        );
1815                        let k = rng.random(0..=positions.len());
1816                        assert_eq!(wm.select(*value, k), positions.get(k).copied());
1817                        assert_eq!(wm.select(*value, positions.len()), None);
1818                    }
1819                    for (&index, ((range, k), &value)) in
1820                        indices.iter().zip(queries.iter().zip(&expected))
1821                    {
1822                        assert_eq!(wm.access(index), values[index]);
1823                        assert_eq!(wm.quantile(range.clone(), *k), value);
1824                    }
1825                    for q in 0..=ranks.len() {
1826                        assert_eq!(wm.rank_batch(ranks[..q].iter().cloned()), counts[..q]);
1827                        assert_eq!(
1828                            wm.rank_lessthan_batch(ranks[..q].iter().cloned()),
1829                            less[..q]
1830                        );
1831                        assert_eq!(
1832                            wm.rank_range_batch(ranges[..q].iter().cloned()),
1833                            range_counts[..q]
1834                        );
1835                    }
1836                    for q in 0..=queries.len() {
1837                        assert_eq!(
1838                            wm.access_batch(indices[..q].iter().copied()),
1839                            indices[..q].iter().map(|&i| values[i]).collect::<Vec<_>>()
1840                        );
1841                        assert_eq!(
1842                            wm.quantile_batch(queries[..q].iter().cloned()),
1843                            expected[..q]
1844                        );
1845                        assert_eq!(
1846                            wm.quantiles_sorted(0..n, &sorted_ranks[..q]),
1847                            sorted_ranks[..q]
1848                                .iter()
1849                                .map(|&k| sorted[k])
1850                                .collect::<Vec<_>>()
1851                        );
1852                    }
1853                    if let Some((range, _)) = queries.last() {
1854                        let mut sorted = values[range.clone()].to_vec();
1855                        sorted.sort_unstable();
1856                        let mut ranks: Vec<_> =
1857                            (0..33).map(|_| rng.random(0..sorted.len())).collect();
1858                        ranks.sort_unstable();
1859                        assert_eq!(
1860                            wm.quantiles_sorted(range.clone(), &ranks),
1861                            ranks.iter().map(|&k| sorted[k]).collect::<Vec<_>>()
1862                        );
1863                        let mut outside = values.clone();
1864                        outside.drain(range.clone());
1865                        outside.sort_unstable();
1866                        if !outside.is_empty() {
1867                            let k = rng.random(0..outside.len());
1868                            assert_eq!(wm.quantile_outer(range.clone(), k), outside[k]);
1869                        }
1870                    }
1871                });
1872            }
1873        }
1874    }
1875
1876    #[test]
1877    fn test_wavelet_matrix_fold() {
1878        use crate::algebra::{Associative, Commutative, Invertible, Magma, Unital};
1879        use std::cmp::Reverse;
1880
1881        enum Sum {}
1882        impl Magma for Sum {
1883            type T = Box<i64>;
1884            fn operate(a: &Self::T, b: &Self::T) -> Self::T {
1885                Box::new(**a + **b)
1886            }
1887        }
1888        impl Unital for Sum {
1889            fn unit() -> Self::T {
1890                Box::new(0)
1891            }
1892        }
1893        impl Associative for Sum {}
1894        impl Commutative for Sum {}
1895        impl Invertible for Sum {
1896            fn inverse(a: &Self::T) -> Self::T {
1897                Box::new(-**a)
1898            }
1899        }
1900
1901        let mut rng = Xorshift::default();
1902        for n in 0..=65 {
1903            let sigma: i64 = rng.random(1..=16);
1904            let values: Vec<_> = (0..n)
1905                .map(|_| Box::new(rng.random(-sigma..=sigma)))
1906                .collect();
1907            let weights: Vec<_> = (0..n).map(|_| Box::new(rng.random(-100..=100))).collect();
1908            let mut dictionary = values.clone();
1909            dictionary.sort_unstable();
1910            dictionary.dedup();
1911            let height = usize::BITS as usize - dictionary.len().leading_zeros() as usize;
1912            let levels: Vec<Vec<_>> = (0..height)
1913                .map(|d| {
1914                    let mut order: Vec<_> = (0..n).collect();
1915                    // Each lower bit takes precedence over the previously partitioned higher bits.
1916                    order.sort_by_key(|&i| {
1917                        (dictionary.binary_search(&values[i]).unwrap() >> d).reverse_bits()
1918                    });
1919                    order
1920                })
1921                .collect();
1922            let mut expected = Vec::new();
1923            for (d, order) in levels.iter().enumerate() {
1924                for (position, &i) in order.iter().enumerate() {
1925                    expected.push((i, Reverse(d), position, values[i].clone()));
1926                }
1927            }
1928            expected.sort_by_key(|&(i, d, _, _)| (i, d));
1929            let mut callbacks = Vec::new();
1930            let original = WaveletMatrix::new_with_init(values.clone(), |d, i, value| {
1931                callbacks.push((d, i, value))
1932            });
1933            assert_eq!(
1934                callbacks,
1935                expected
1936                    .into_iter()
1937                    .map(|(_, Reverse(d), i, value)| (d, i, value))
1938                    .collect::<Vec<_>>()
1939            );
1940            let queries: Vec<_> = (0..33)
1941                .map(|i| {
1942                    if i == 0 {
1943                        return (Box::new(-sigma - 1)..Box::new(sigma + 1), 0..n);
1944                    }
1945                    let left = rng.random(0..=n);
1946                    let right = rng.random(left..=n);
1947                    let lower = rng.random(-sigma - 1..=sigma + 1);
1948                    let upper = rng.random(lower..=sigma + 1);
1949                    (Box::new(lower)..Box::new(upper), left..right)
1950                })
1951                .collect();
1952            for binary_only in [false, true] {
1953                let mut wm = original.clone();
1954                if binary_only {
1955                    wm.quad_vectors.clear();
1956                }
1957                for_each_backend(&mut wm, |wm| {
1958                    let fold = wm.build_fold::<Sum>(&weights);
1959                    let mut dynamic = wm.build_point_add::<Sum>(&weights);
1960                    let mut updated = weights.clone();
1961                    let mut less_sums = Vec::new();
1962                    for (step, (bounds, range)) in queries.iter().enumerate() {
1963                        if n != 0 {
1964                            let i = if step == 0 { n - 1 } else { rng.random(0..n) };
1965                            let delta = rng.random(-100..=100);
1966                            dynamic.update(i, Box::new(delta));
1967                            *updated[i] += delta;
1968                        }
1969                        let selected: Vec<_> = range
1970                            .clone()
1971                            .filter(|&i| bounds.contains(&values[i]))
1972                            .collect();
1973                        let sum: i64 = selected.iter().map(|&i| *weights[i]).sum();
1974                        let updated_sum: i64 = selected.iter().map(|&i| *updated[i]).sum();
1975                        assert_eq!(*fold.fold_range(bounds.clone(), range.clone()), sum);
1976                        assert_eq!(
1977                            fold.fold_range_with_count(bounds.clone(), range.clone()),
1978                            (selected.len(), Box::new(sum))
1979                        );
1980                        assert_eq!(
1981                            *dynamic.fold_range(bounds.clone(), range.clone()),
1982                            updated_sum
1983                        );
1984                        let selected: Vec<_> =
1985                            range.clone().filter(|&i| values[i] < bounds.end).collect();
1986                        let sum: i64 = selected.iter().map(|&i| *weights[i]).sum();
1987                        let updated_sum: i64 = selected.iter().map(|&i| *updated[i]).sum();
1988                        assert_eq!(*fold.fold_lessthan(bounds.end.clone(), range.clone()), sum);
1989                        assert_eq!(
1990                            fold.fold_lessthan_with_count(bounds.end.clone(), range.clone()),
1991                            (selected.len(), Box::new(sum))
1992                        );
1993                        assert_eq!(
1994                            *dynamic.fold_lessthan(bounds.end.clone(), range.clone()),
1995                            updated_sum
1996                        );
1997                        less_sums.push(Box::new(sum));
1998                        let mut actual = Vec::new();
1999                        let mut dimensions = Vec::new();
2000                        wm.query_less_than(bounds.end.clone(), range.clone(), |d, r| {
2001                            dimensions.push(d);
2002                            actual.extend_from_slice(&levels[d][r]);
2003                        });
2004                        actual.sort_unstable();
2005                        assert_eq!(actual, selected);
2006                        let bound = dictionary.partition_point(|v| v < &bounds.end);
2007                        assert_eq!(
2008                            dimensions,
2009                            (0..height)
2010                                .rev()
2011                                .filter(|&d| (bound >> d) & 1 != 0)
2012                                .collect::<Vec<_>>()
2013                        );
2014                    }
2015                    for q in 0..=queries.len() {
2016                        assert_eq!(
2017                            fold.fold_lessthan_batch(
2018                                queries[..q].iter().map(|(v, r)| (v.end.clone(), r.clone()))
2019                            ),
2020                            less_sums[..q]
2021                        );
2022                    }
2023                });
2024            }
2025        }
2026    }
2027
2028    fn for_each_backend<T>(wm: &mut WaveletMatrix<T>, mut test: impl FnMut(&WaveletMatrix<T>)) {
2029        #[cfg(target_arch = "x86_64")]
2030        for backend in [
2031            crate::tools::SimdBackend::Scalar,
2032            crate::tools::SimdBackend::Avx2,
2033            crate::tools::SimdBackend::Avx512,
2034        ] {
2035            if (backend == crate::tools::SimdBackend::Avx2 && !is_x86_feature_detected!("avx2"))
2036                || (backend == crate::tools::SimdBackend::Avx512
2037                    && !(is_x86_feature_detected!("avx512f")
2038                        && is_x86_feature_detected!("avx512vpopcntdq")))
2039            {
2040                continue;
2041            }
2042            wm.backend = backend;
2043            test(wm);
2044        }
2045        #[cfg(not(target_arch = "x86_64"))]
2046        test(wm);
2047    }
2048}