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 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)] 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 ¤t[..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(¤t[..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 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 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}