Skip to main content

WaveletMatrix

Struct WaveletMatrix 

Source
pub struct WaveletMatrix<T> {
    len: usize,
    bit_length: usize,
    zeros: Vec<usize>,
    bit_vectors: Vec<BitVector>,
    quad_vectors: Vec<WaveletMatrixQuadVector>,
    compress: VecCompress<T>,
    backend: SimdBackend,
}

Fields§

§len: usize§bit_length: usize§zeros: Vec<usize>§bit_vectors: Vec<BitVector>§quad_vectors: Vec<WaveletMatrixQuadVector>§compress: VecCompress<T>§backend: SimdBackend

Implementations§

Source§

impl<T> WaveletMatrix<T>
where T: Ord + Clone,

Source

pub fn new(v: Vec<T>) -> Self

Examples found in repository?
crates/library_checker/src/data_structure/static_range_frequency.rs (line 20)
17pub fn static_range_frequency_wavelet_matrix(reader: impl Read, writer: impl Write) {
18    prepare_io!(reader, writer);
19    sc!(n, q, a: [usize; n]);
20    let wm = WaveletMatrix::new(a);
21    for _ in 0..q {
22        sc!(l, r, x: usize);
23        let ans = wm.rank(x, l..r);
24        pp!(ans);
25    }
26}
More examples
Hide additional examples
crates/library_checker/src/data_structure/range_kth_smallest.rs (line 8)
5pub fn range_kth_smallest(reader: impl Read, writer: impl Write) {
6    prepare_io!(reader, writer);
7    sc!(n, q, a: [usize; n], queries: [(usize, usize, usize); iter q]);
8    let wm = WaveletMatrix::new(a);
9    let results = wm.quantile_batch(queries.map(|(l, r, k)| (l..r, k)));
10    pp!(@lf @it results);
11}
crates/library_checker/src/data_structure/static_range_sum_with_upper_bound.rs (line 26)
22pub fn static_range_sum_with_upper_bound_wavelet_matrix(reader: impl Read, writer: impl Write) {
23    prepare_io!(reader, writer);
24    sc!(n, q, a: [i64; n]);
25    let weights = a.clone();
26    let wm = WaveletMatrix::new(a);
27    let fold = wm.build_fold::<AdditiveOperation<i64>>(&weights);
28    for _ in 0..q {
29        sc!(l, r, x: i64);
30        let (count, sum) = fold.fold_lessthan_with_count(x + 1, l..r);
31        pp!(count, sum);
32    }
33}
crates/library_checker/src/data_structure/rectangle_sum.rs (line 20)
9pub fn rectangle_sum(reader: impl Read, writer: impl Write) {
10    prepare_io!(buffered; reader, writer);
11    sc!(n, q, mut xyw: [(u32, u32, i64); n], queries: [(u32, u32, u32, u32); q]);
12    xyw.radix_sort_by_key(|&(x, ..)| x);
13    let xs: Vec<_> = xyw.iter().map(|&(x, ..)| x).collect();
14    let search = StaticSearch::from_sorted(&xs);
15    let endpoints: Vec<_> = queries.iter().flat_map(|&(l, _, r, _)| [l, r]).collect();
16    let mut positions = vec![0; endpoints.len()];
17    search.lower_bound_batch(&endpoints, &mut positions);
18    let ys = xyw.iter().map(|&(_, y, _)| y).collect();
19    let weights: Vec<_> = xyw.iter().map(|&(_, _, w)| w).collect();
20    let wm = WaveletMatrix::new(ys);
21    let fold = wm.build_fold::<AdditiveOperation<i64>>(&weights);
22    let result = fold.fold_lessthan_batch(
23        queries
24            .into_iter()
25            .zip(positions.as_chunks::<2>().0)
26            .flat_map(|((_, d, _, u), &[l, r])| [(d, l..r), (u, l..r)]),
27    );
28    for &[lower, upper] in result.as_chunks::<2>().0 {
29        pp!(upper - lower);
30    }
31}
crates/library_checker/src/data_structure/point_add_rectangle_sum.rs (line 39)
17pub fn point_add_rectangle_sum(reader: impl Read, writer: impl Write) {
18    prepare_io!(reader, writer);
19    sc!(n, q, xyw: [(u32, u32, u64); iter n]);
20    let mut points: Vec<_> = xyw.map(|(x, y, w)| (x, y, w as i64)).collect();
21    sc!(queries: [Query; q]);
22    points.extend(queries.iter().filter_map(|&query| match query {
23        Query::Add { x, y, .. } => Some((x, y, 0)),
24        Query::Sum { .. } => None,
25    }));
26    let mut order: Vec<_> = (0..points.len()).collect();
27    order.radix_sort_by_key(|&i| points[i].0);
28    let mut positions = vec![0; points.len()];
29    let mut xs = Vec::with_capacity(points.len());
30    let mut ys = Vec::with_capacity(points.len());
31    let mut weights = Vec::with_capacity(points.len());
32    for (i, &point) in order.iter().enumerate() {
33        positions[point] = i;
34        let (x, y, w) = points[point];
35        xs.push(x);
36        ys.push(y);
37        weights.push(w);
38    }
39    let wm = WaveletMatrix::new(ys);
40    let mut fold: WaveletMatrixPointAdd<_, AdditiveOperation<i64>> = wm.build_point_add(&weights);
41
42    let mut point = n;
43    for query in queries {
44        match query {
45            Query::Add { w, .. } => {
46                fold.update(positions[point], w as i64);
47                point += 1;
48            }
49            Query::Sum { l, d, r, u } => {
50                let l = xs.partition_point(|&x| x < l);
51                let r = xs.partition_point(|&x| x < r);
52                pp!(fold.fold_range(d..u, l..r));
53            }
54        }
55    }
56}
crates/competitive/src/data_structure/wavelet_matrix.rs (line 597)
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    }
Source

fn pack_words<I: Copy>(indices: &[I], bit: impl Fn(I) -> bool) -> Vec<u64>

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 469)
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    }
Source

fn partition<I: Copy>(indices: &[I], words: &[u64], one: usize, next: &mut [I])

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 478)
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    }
Source

fn from_values<I: Copy>( v: Vec<T>, code: impl Fn(usize) -> I, index: impl Fn(I) -> usize, pack: impl Fn(&[I], usize) -> Vec<u64>, partition: impl Fn(&[I], &[u64], usize, &mut [I]), ) -> Self

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (lines 459-480)
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    }
Source

pub fn new_with_init<F>(v: Vec<T>, f: F) -> Self
where F: FnMut(usize, usize, T),

Source

fn level(&self, d: usize) -> usize

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 616)
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    }
Source

fn rank1(&self, level: usize, k: usize) -> usize

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 700)
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    }
Source

fn reorder<U>(&self, level: usize, current: Vec<U>) -> Vec<U>

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 1440)
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    }
Source

fn range_by_index(&self, idx: usize, range: Range<usize>) -> Range<usize> ⓘ

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 921)
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    }
Source

fn batch<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize)

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 908)
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    }
Source

unsafe fn batch_popcnt<const OP: u8>( &self, states: &mut [[usize; 4]; 16], count: usize, )

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 718)
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    }
Source

fn batch_inner<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize)

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 722)
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    }
Source

pub fn access(&self, k: usize) -> T

get k-th value

Source

pub fn access_batch(&self, indices: impl IntoIterator<Item = usize>) -> Vec<T>

Returns the values at indices in input order.

Source

pub fn rank(&self, val: T, range: Range<usize>) -> usize

the number of val in range

Examples found in repository?
crates/library_checker/src/data_structure/static_range_frequency.rs (line 23)
17pub fn static_range_frequency_wavelet_matrix(reader: impl Read, writer: impl Write) {
18    prepare_io!(reader, writer);
19    sc!(n, q, a: [usize; n]);
20    let wm = WaveletMatrix::new(a);
21    for _ in 0..q {
22        sc!(l, r, x: usize);
23        let ans = wm.rank(x, l..r);
24        pp!(ans);
25    }
26}
Source

pub fn rank_batch( &self, queries: impl IntoIterator<Item = (T, Range<usize>)>, ) -> Vec<usize>

Returns the number of exact matches for each (value, range) query.

Source

pub fn select(&self, val: T, k: usize) -> Option<usize>

index of k-th val

Source

pub fn quantile(&self, range: Range<usize>, k: usize) -> T

get k-th smallest value in range

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 1036)
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    }
Source

fn quad_quantile( &self, range: Range<usize>, k: usize, level: usize, index: usize, ) -> T

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 978)
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    }
Source

pub fn quantile_batch( &self, queries: impl IntoIterator<Item = (Range<usize>, usize)>, ) -> Vec<T>

Examples found in repository?
crates/library_checker/src/data_structure/range_kth_smallest.rs (line 9)
5pub fn range_kth_smallest(reader: impl Read, writer: impl Write) {
6    prepare_io!(reader, writer);
7    sc!(n, q, a: [usize; n], queries: [(usize, usize, usize); iter q]);
8    let wm = WaveletMatrix::new(a);
9    let results = wm.quantile_batch(queries.map(|(l, r, k)| (l..r, k)));
10    pp!(@lf @it results);
11}
More examples
Hide additional examples
crates/competitive/src/data_structure/wavelet_matrix.rs (line 1066)
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    }
Source

pub fn quantiles_sorted(&self, range: Range<usize>, ranks: &[usize]) -> Vec<T>

Returns the requested order statistics of one range. ranks must be sorted in nondecreasing order and every rank must be less than the range length.

Source

pub fn quantile_outer(&self, range: Range<usize>, k: usize) -> T

get k-th smallest value out of range

Source

pub fn rank_lessthan(&self, val: T, range: Range<usize>) -> usize

the number of value less than val in range

Source

pub fn rank_lessthan_batch( &self, queries: impl IntoIterator<Item = (T, Range<usize>)>, ) -> Vec<usize>

Returns the number of values below each query’s threshold.

Source

fn quad_rank_lessthan_index( &self, index: usize, range: Range<usize>, start_level: usize, ) -> usize

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 1201)
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    }
Source

fn rank_lessthan_index( &self, idx: usize, range: Range<usize>, bits: usize, ) -> usize

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 1205)
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    }
Source

pub fn rank_range(&self, valrange: Range<T>, range: Range<usize>) -> usize

the number of valrange in range

Source

pub fn rank_range_batch( &self, queries: impl IntoIterator<Item = (Range<T>, Range<usize>)>, ) -> Vec<usize>

Returns each query’s count in a half-open value range.

Source

pub fn query_less_than<F>(&self, val: T, range: Range<usize>, f: F)
where F: FnMut(usize, Range<usize>),

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (lines 1505-1510)
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    }
Source

pub fn build_fold<M>(&self, weights: &[M::T]) -> WaveletMatrixFold<'_, T, M>
where M: AbelianGroup,

Examples found in repository?
crates/library_checker/src/data_structure/static_range_sum_with_upper_bound.rs (line 27)
22pub fn static_range_sum_with_upper_bound_wavelet_matrix(reader: impl Read, writer: impl Write) {
23    prepare_io!(reader, writer);
24    sc!(n, q, a: [i64; n]);
25    let weights = a.clone();
26    let wm = WaveletMatrix::new(a);
27    let fold = wm.build_fold::<AdditiveOperation<i64>>(&weights);
28    for _ in 0..q {
29        sc!(l, r, x: i64);
30        let (count, sum) = fold.fold_lessthan_with_count(x + 1, l..r);
31        pp!(count, sum);
32    }
33}
More examples
Hide additional examples
crates/library_checker/src/data_structure/rectangle_sum.rs (line 21)
9pub fn rectangle_sum(reader: impl Read, writer: impl Write) {
10    prepare_io!(buffered; reader, writer);
11    sc!(n, q, mut xyw: [(u32, u32, i64); n], queries: [(u32, u32, u32, u32); q]);
12    xyw.radix_sort_by_key(|&(x, ..)| x);
13    let xs: Vec<_> = xyw.iter().map(|&(x, ..)| x).collect();
14    let search = StaticSearch::from_sorted(&xs);
15    let endpoints: Vec<_> = queries.iter().flat_map(|&(l, _, r, _)| [l, r]).collect();
16    let mut positions = vec![0; endpoints.len()];
17    search.lower_bound_batch(&endpoints, &mut positions);
18    let ys = xyw.iter().map(|&(_, y, _)| y).collect();
19    let weights: Vec<_> = xyw.iter().map(|&(_, _, w)| w).collect();
20    let wm = WaveletMatrix::new(ys);
21    let fold = wm.build_fold::<AdditiveOperation<i64>>(&weights);
22    let result = fold.fold_lessthan_batch(
23        queries
24            .into_iter()
25            .zip(positions.as_chunks::<2>().0)
26            .flat_map(|((_, d, _, u), &[l, r])| [(d, l..r), (u, l..r)]),
27    );
28    for &[lower, upper] in result.as_chunks::<2>().0 {
29        pp!(upper - lower);
30    }
31}
Source

pub fn build_point_add<M>( &self, weights: &[M::T], ) -> WaveletMatrixPointAdd<'_, T, M>
where M: AbelianGroup,

Examples found in repository?
crates/library_checker/src/data_structure/point_add_rectangle_sum.rs (line 40)
17pub fn point_add_rectangle_sum(reader: impl Read, writer: impl Write) {
18    prepare_io!(reader, writer);
19    sc!(n, q, xyw: [(u32, u32, u64); iter n]);
20    let mut points: Vec<_> = xyw.map(|(x, y, w)| (x, y, w as i64)).collect();
21    sc!(queries: [Query; q]);
22    points.extend(queries.iter().filter_map(|&query| match query {
23        Query::Add { x, y, .. } => Some((x, y, 0)),
24        Query::Sum { .. } => None,
25    }));
26    let mut order: Vec<_> = (0..points.len()).collect();
27    order.radix_sort_by_key(|&i| points[i].0);
28    let mut positions = vec![0; points.len()];
29    let mut xs = Vec::with_capacity(points.len());
30    let mut ys = Vec::with_capacity(points.len());
31    let mut weights = Vec::with_capacity(points.len());
32    for (i, &point) in order.iter().enumerate() {
33        positions[point] = i;
34        let (x, y, w) = points[point];
35        xs.push(x);
36        ys.push(y);
37        weights.push(w);
38    }
39    let wm = WaveletMatrix::new(ys);
40    let mut fold: WaveletMatrixPointAdd<_, AdditiveOperation<i64>> = wm.build_point_add(&weights);
41
42    let mut point = n;
43    for query in queries {
44        match query {
45            Query::Add { w, .. } => {
46                fold.update(positions[point], w as i64);
47                point += 1;
48            }
49            Query::Sum { l, d, r, u } => {
50                let l = xs.partition_point(|&x| x < l);
51                let r = xs.partition_point(|&x| x < r);
52                pp!(fold.fold_range(d..u, l..r));
53            }
54        }
55    }
56}

Trait Implementations§

Source§

impl<T: Clone> Clone for WaveletMatrix<T>

Source§

fn clone(&self) -> Self

Returns a duplicate of the value. Read more
1.0.0 (const: unstable) · Source§

fn clone_from(&mut self, source: &Self)

Performs copy-assignment from source. Read more
Source§

impl<T: Debug> Debug for WaveletMatrix<T>

Source§

fn fmt(&self, f: &mut Formatter<'_>) -> Result

Formats the value using the given formatter. Read more

Auto Trait Implementations§

§

impl<T> Freeze for WaveletMatrix<T>
where VecCompress<T>: Freeze,

§

impl<T> RefUnwindSafe for WaveletMatrix<T>

§

impl<T> Send for WaveletMatrix<T>
where VecCompress<T>: Send,

§

impl<T> Sync for WaveletMatrix<T>
where VecCompress<T>: Sync,

§

impl<T> Unpin for WaveletMatrix<T>
where VecCompress<T>: Unpin,

§

impl<T> UnsafeUnpin for WaveletMatrix<T>

§

impl<T> UnwindSafe for WaveletMatrix<T>

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
Source§

impl<T> CloneToUninit for T
where T: Clone,

Source§

unsafe fn clone_to_uninit(&self, dest: *mut u8)

🔬This is a nightly-only experimental API. (clone_to_uninit)
Performs copy-assignment from self to dest. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> ToArrayVecScalar for T

Source§

impl<T> ToOwned for T
where T: Clone,

Source§

type Owned = T

The resulting type after obtaining ownership.
Source§

fn to_owned(&self) -> T

Creates owned data from borrowed data, usually by cloning. Read more
Source§

fn clone_into(&self, target: &mut T)

Uses borrowed data to replace owned data, usually by cloning. Read more
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = !

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, !>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.