Skip to main content

WaveletMatrixQuadVector

Struct WaveletMatrixQuadVector 

Source
struct WaveletMatrixQuadVector {
    blocks: Vec<WaveletMatrixQuadBlock>,
    starts: [usize; 4],
    select_samples: [Vec<usize>; 4],
}

Fields§

§blocks: Vec<WaveletMatrixQuadBlock>§starts: [usize; 4]§select_samples: [Vec<usize>; 4]

Implementations§

Source§

impl WaveletMatrixQuadVector

Source

fn from_words(low: &[u64], high: Option<&[u64]>, len: usize) -> Self

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 559)
524    fn from_values<I: Copy>(
525        v: Vec<T>,
526        code: impl Fn(usize) -> I,
527        index: impl Fn(I) -> usize,
528        pack: impl Fn(&[I], usize) -> Vec<u64>,
529        partition: impl Fn(&[I], &[u64], usize, &mut [I]),
530    ) -> Self {
531        let len = v.len();
532        let mut sorted: Vec<_> = v
533            .into_iter()
534            .enumerate()
535            .map(|(i, value)| (value, code(i)))
536            .collect();
537        sorted.sort_unstable_by(|a, b| a.0.cmp(&b.0));
538        let mut values = Vec::with_capacity(len);
539        let mut indices = vec![code(0); len];
540        for (value, i) in sorted {
541            if values.last().is_none_or(|last| last != &value) {
542                values.push(value);
543            }
544            indices[index(i)] = code(values.len() - 1);
545        }
546        let compress = VecCompress::from_sorted_unique(values);
547        let bit_length = usize::BITS as usize - compress.size().leading_zeros() as usize;
548        let mut bit_vectors = Vec::with_capacity(bit_length);
549        let mut zeros = Vec::with_capacity(bit_length);
550        let quad_bits =
551            usize::BITS as usize - compress.size().saturating_sub(1).leading_zeros() as usize;
552        let mut quad_vectors = Vec::with_capacity(quad_bits.div_ceil(2));
553        let mut next = indices.clone();
554        for d in (0..bit_length).rev() {
555            let words = pack(&indices, d);
556            if len <= u32::MAX as usize && d < quad_bits && (d % 2 == 1 || d + 1 == quad_bits) {
557                if d % 2 == 1 {
558                    let low = pack(&indices, d - 1);
559                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&low, Some(&words), len));
560                } else {
561                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&words, None, len));
562                }
563            }
564            let bits = BitVector::from_words(&words, len);
565            let zero_count = bits.rank0(len);
566            if d == 0 {
567                zeros.push(zero_count);
568                bit_vectors.push(bits);
569                break;
570            }
571            partition(&indices, &words, zero_count, &mut next);
572            zeros.push(zero_count);
573            bit_vectors.push(bits);
574            mem::swap(&mut indices, &mut next);
575        }
576        Self {
577            len,
578            bit_length,
579            zeros,
580            bit_vectors,
581            quad_vectors,
582            compress,
583            #[cfg(target_arch = "x86_64")]
584            backend: match super::simd_backend() {
585                super::SimdBackend::Avx512 if !is_x86_feature_detected!("avx512vpopcntdq") => {
586                    super::SimdBackend::Avx2
587                }
588                backend => backend,
589            },
590        }
591    }
Source

fn ranks(&self, position: usize) -> [usize; 4]

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 772)
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 rank(&self, digit: usize, position: usize) -> usize

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 121)
116    fn access_rank(&self, position: usize) -> (usize, usize) {
117        let block = &self.blocks[position / 64];
118        let offset = position % 64;
119        let digit =
120            ((block.lo >> offset) & 1) as usize | (((block.hi >> offset) & 1) as usize) << 1;
121        (digit, self.rank(digit, position))
122    }
123}
124
125#[cfg(target_arch = "x86_64")]
126mod simd {
127    #![allow(unsafe_op_in_unsafe_fn)] // All entry points check the required CPU features.
128    use super::BitVector;
129    use std::arch::x86_64::*;
130
131    #[target_feature(enable = "avx2")]
132    pub unsafe fn pack_words(indices: &[u32], bit: usize) -> Vec<u64> {
133        let mut result = Vec::with_capacity(indices.len().div_ceil(64));
134        let shift = _mm_cvtsi32_si128((31 - bit) as i32);
135        for chunk in indices.chunks(64) {
136            let mut word = 0u64;
137            let mut i = 0;
138            while i + 8 <= chunk.len() {
139                let v = _mm256_loadu_si256(chunk.as_ptr().add(i).cast());
140                word |= (_mm256_movemask_ps(_mm256_castsi256_ps(_mm256_sll_epi32(v, shift)))
141                    as u64)
142                    << i;
143                i += 8;
144            }
145            for (j, &value) in chunk[i..].iter().enumerate() {
146                word |= (((value >> bit) & 1) as u64) << (i + j);
147            }
148            result.push(word);
149        }
150        result
151    }
152
153    #[target_feature(enable = "avx2")]
154    pub unsafe fn partition_avx2(indices: &[u32], words: &[u64], mut one: usize, next: &mut [u32]) {
155        const PERM: [[i32; 8]; 256] = {
156            let mut table = [[0; 8]; 256];
157            let mut mask = 0;
158            while mask < 256 {
159                let mut pos = 0;
160                let mut bit = 0;
161                while bit < 2 {
162                    let mut i = 0;
163                    while i < 8 {
164                        if (mask >> i) & 1 == bit {
165                            table[mask][pos] = i;
166                            pos += 1;
167                        }
168                        i += 1;
169                    }
170                    bit += 1;
171                }
172                mask += 1;
173            }
174            table
175        };
176        let order = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
177        let mut zero = 0;
178        for (chunk, &word) in indices.chunks(64).zip(words) {
179            if word == 0 || word == u64::MAX {
180                let at = if word == 0 { &mut zero } else { &mut one };
181                next[*at..*at + chunk.len()].copy_from_slice(chunk);
182                *at += chunk.len();
183                continue;
184            }
185            let mut i = 0;
186            while i + 8 <= chunk.len() {
187                let value = _mm256_loadu_si256(chunk.as_ptr().add(i).cast());
188                let mask = ((word >> i) & 255) as usize;
189                let ones = mask.count_ones() as usize;
190                let zeros = 8 - ones;
191                let packed = _mm256_permutevar8x32_epi32(
192                    value,
193                    _mm256_loadu_si256(PERM[mask].as_ptr().cast()),
194                );
195                let rotated = _mm256_permutevar8x32_epi32(
196                    packed,
197                    _mm256_add_epi32(order, _mm256_set1_epi32(zeros as i32)),
198                );
199                _mm256_maskstore_epi32(
200                    next.as_mut_ptr().add(zero).cast(),
201                    _mm256_cmpgt_epi32(_mm256_set1_epi32(zeros as i32), order),
202                    packed,
203                );
204                _mm256_maskstore_epi32(
205                    next.as_mut_ptr().add(one).cast(),
206                    _mm256_cmpgt_epi32(_mm256_set1_epi32(ones as i32), order),
207                    rotated,
208                );
209                zero += zeros;
210                one += ones;
211                i += 8;
212            }
213            for (j, &value) in chunk[i..].iter().enumerate() {
214                let bit = (word >> (i + j)) & 1 != 0;
215                next[if bit { one } else { zero }] = value;
216                zero += !bit as usize;
217                one += bit as usize;
218            }
219        }
220    }
221
222    #[target_feature(enable = "avx2")]
223    #[inline]
224    unsafe fn popcount_avx2(value: __m256i) -> __m256i {
225        let lookup = _mm256_setr_epi8(
226            0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2, 3, 3, 4, 0, 1, 1, 2, 1, 2, 2, 3, 1, 2, 2, 3, 2,
227            3, 3, 4,
228        );
229        let mask = _mm256_set1_epi8(15);
230        let low = _mm256_shuffle_epi8(lookup, _mm256_and_si256(value, mask));
231        let high = _mm256_shuffle_epi8(
232            lookup,
233            _mm256_and_si256(_mm256_srli_epi16::<4>(value), mask),
234        );
235        _mm256_sad_epu8(_mm256_add_epi8(low, high), _mm256_setzero_si256())
236    }
237
238    #[target_feature(enable = "avx512f")]
239    #[inline]
240    unsafe fn greater_avx512(a: __m512i, b: __m512i) -> __m512i {
241        _mm512_maskz_set1_epi64(_mm512_cmpgt_epi64_mask(a, b), -1)
242    }
243
244    #[target_feature(enable = "avx512f")]
245    #[inline]
246    unsafe fn gather_avx512<const SCALE: i32>(base: *const i64, index: __m512i) -> __m512i {
247        _mm512_i64gather_epi64::<SCALE>(index, base)
248    }
249
250    macro_rules! rank_lessthan {
251        ($name:ident, $feature:literal, $lanes:literal,
252         $load:ident, $store:ident, $set:ident, $gather:ident,
253         $add:ident, $sub:ident, $and:ident, $andnot:ident, $or:ident,
254         $left:ident, $right:ident, $gt:ident, $popcount:ident) => {
255            #[target_feature(enable = $feature)]
256            pub unsafe fn $name(vectors: &[BitVector], zeros: &[usize], states: &mut [[usize; 4]]) {
257                for chunk in states.chunks_exact_mut($lanes) {
258                    let mut starts = [0u64; $lanes];
259                    let mut ends = [0u64; $lanes];
260                    let mut keys = [0u64; $lanes];
261                    let mut values = [0u64; $lanes];
262                    for (i, state) in chunk.iter().enumerate() {
263                        starts[i] = state[0] as u64;
264                        ends[i] = state[1] as u64;
265                        keys[i] = state[2] as u64;
266                    }
267                    let mut start = $load(starts.as_ptr().cast());
268                    let mut end = $load(ends.as_ptr().cast());
269                    let key = $load(keys.as_ptr().cast());
270                    let mut value = $set(0);
271                    let one = $set(1);
272                    for (level, vector) in vectors.iter().enumerate() {
273                        let d = vectors.len() - 1 - level;
274                        let base = vector.blocks().as_ptr().cast::<i64>();
275                        let rank1 = |position| {
276                            let offset = $and(position, $set(63));
277                            let index = $left($right(position, $set(6)), one);
278                            let bits = $gather::<8>(base, index);
279                            let prefix = $gather::<8>(base, $add(index, one));
280                            $add(prefix, $popcount($and(bits, $sub($left(one, offset), one))))
281                        };
282                        let rank = rank1(start);
283                        let end_rank = rank1(end);
284                        let start0 = $sub(start, rank);
285                        let end0 = $sub(end, end_rank);
286                        let count = $sub(end0, start0);
287                        let mask = $gt($and($right(key, $set(d as i64)), one), $set(0));
288                        let zero = $set(zeros[level] as i64);
289                        start = $or($andnot(mask, start0), $and(mask, $add(zero, rank)));
290                        end = $or($andnot(mask, end0), $and(mask, $add(zero, end_rank)));
291                        value = $add(value, $and(mask, count));
292                    }
293                    $store(values.as_mut_ptr().cast(), value);
294                    for (i, state) in chunk.iter_mut().enumerate() {
295                        state[3] = values[i] as usize;
296                    }
297                }
298            }
299        };
300    }
301
302    rank_lessthan!(
303        rank_lessthan_avx2,
304        "avx2",
305        4,
306        _mm256_loadu_si256,
307        _mm256_storeu_si256,
308        _mm256_set1_epi64x,
309        _mm256_i64gather_epi64,
310        _mm256_add_epi64,
311        _mm256_sub_epi64,
312        _mm256_and_si256,
313        _mm256_andnot_si256,
314        _mm256_or_si256,
315        _mm256_sllv_epi64,
316        _mm256_srlv_epi64,
317        _mm256_cmpgt_epi64,
318        popcount_avx2
319    );
320    rank_lessthan!(
321        rank_lessthan_avx512,
322        "avx512f,avx512vpopcntdq",
323        8,
324        _mm512_loadu_si512,
325        _mm512_storeu_si512,
326        _mm512_set1_epi64,
327        gather_avx512,
328        _mm512_add_epi64,
329        _mm512_sub_epi64,
330        _mm512_and_si512,
331        _mm512_andnot_si512,
332        _mm512_or_si512,
333        _mm512_sllv_epi64,
334        _mm512_srlv_epi64,
335        greater_avx512,
336        _mm512_popcnt_epi64
337    );
338
339    #[target_feature(enable = "avx512f,avx512vpopcntdq")]
340    pub unsafe fn quad_avx512(
341        layers: &[super::WaveletMatrixQuadVector],
342        states: &mut [[usize; 4]],
343    ) {
344        for chunk in states.as_chunks_mut::<8>().0 {
345            let mut starts = [0u64; 8];
346            let mut ends = [0u64; 8];
347            let mut keys = [0u64; 8];
348            let mut result = [0u64; 8];
349            for (i, s) in chunk.iter().enumerate() {
350                starts[i] = s[0] as u64;
351                ends[i] = s[1] as u64;
352                keys[i] = s[2] as u64;
353            }
354            let mut start = _mm512_loadu_si512(starts.as_ptr().cast());
355            let mut end = _mm512_loadu_si512(ends.as_ptr().cast());
356            let mut key = _mm512_loadu_si512(keys.as_ptr().cast());
357            let mut code = _mm512_set1_epi64(0);
358            let one = _mm512_set1_epi64(1);
359            macro_rules! blend {
360                ($m:expr,$x:expr,$y:expr) => {
361                    _mm512_or_si512(_mm512_andnot_si512($m, $x), _mm512_and_si512($m, $y))
362                };
363            }
364            for layer in layers {
365                let base = layer.blocks.as_ptr().cast::<i64>();
366                let ranks = |pos| {
367                    let offset = _mm512_and_si512(pos, _mm512_set1_epi64(63));
368                    let mask = _mm512_sub_epi64(_mm512_sllv_epi64(one, offset), one);
369                    let i = _mm512_sllv_epi64(
370                        _mm512_srlv_epi64(pos, _mm512_set1_epi64(6)),
371                        _mm512_set1_epi64(2),
372                    );
373                    let lo = _mm512_and_si512(gather_avx512::<8>(base, i), mask);
374                    let hi =
375                        _mm512_and_si512(gather_avx512::<8>(base, _mm512_add_epi64(i, one)), mask);
376                    let a = _mm512_popcnt_epi64(lo);
377                    let b = _mm512_popcnt_epi64(hi);
378                    let c = _mm512_popcnt_epi64(_mm512_and_si512(lo, hi));
379                    let p01 = gather_avx512::<8>(base, _mm512_add_epi64(i, _mm512_set1_epi64(2)));
380                    let p23 = gather_avx512::<8>(base, _mm512_add_epi64(i, _mm512_set1_epi64(3)));
381                    let r0 = _mm512_add_epi64(
382                        _mm512_and_si512(p01, _mm512_set1_epi64(u32::MAX as i64)),
383                        _mm512_add_epi64(_mm512_sub_epi64(_mm512_sub_epi64(offset, a), b), c),
384                    );
385                    let r1 = _mm512_add_epi64(
386                        _mm512_srlv_epi64(p01, _mm512_set1_epi64(32)),
387                        _mm512_sub_epi64(a, c),
388                    );
389                    let r2 = _mm512_add_epi64(
390                        _mm512_and_si512(p23, _mm512_set1_epi64(u32::MAX as i64)),
391                        _mm512_sub_epi64(b, c),
392                    );
393                    let r3 = _mm512_sub_epi64(_mm512_sub_epi64(_mm512_sub_epi64(pos, r0), r1), r2);
394                    [r0, r1, r2, r3]
395                };
396                let l = ranks(start);
397                let r = ranks(end);
398                let low_count =
399                    _mm512_sub_epi64(_mm512_add_epi64(r[0], r[1]), _mm512_add_epi64(l[0], l[1]));
400                let high = greater_avx512(key, _mm512_sub_epi64(low_count, one));
401                key = _mm512_sub_epi64(key, _mm512_and_si512(high, low_count));
402                let l0 = blend!(high, l[0], l[2]);
403                let r0 = blend!(high, r[0], r[2]);
404                let l1 = blend!(high, l[1], l[3]);
405                let r1 = blend!(high, r[1], r[3]);
406                let count0 = _mm512_sub_epi64(r0, l0);
407                let low = greater_avx512(key, _mm512_sub_epi64(count0, one));
408                key = _mm512_sub_epi64(key, _mm512_and_si512(low, count0));
409                let base0 = blend!(
410                    high,
411                    _mm512_set1_epi64(layer.starts[0] as i64),
412                    _mm512_set1_epi64(layer.starts[2] as i64)
413                );
414                let base1 = blend!(
415                    high,
416                    _mm512_set1_epi64(layer.starts[1] as i64),
417                    _mm512_set1_epi64(layer.starts[3] as i64)
418                );
419                let offset = blend!(low, base0, base1);
420                start = _mm512_add_epi64(offset, blend!(low, l0, l1));
421                end = _mm512_add_epi64(offset, blend!(low, r0, r1));
422                code = _mm512_or_si512(
423                    _mm512_sllv_epi64(code, _mm512_set1_epi64(2)),
424                    _mm512_or_si512(
425                        _mm512_and_si512(high, _mm512_set1_epi64(2)),
426                        _mm512_and_si512(low, one),
427                    ),
428                );
429            }
430            _mm512_storeu_si512(result.as_mut_ptr().cast(), code);
431            for (i, s) in chunk.iter_mut().enumerate() {
432                s[3] = result[i] as usize;
433            }
434        }
435    }
436}
437
438#[derive(Debug, Clone)]
439pub struct WaveletMatrix<T> {
440    len: usize,
441    bit_length: usize,
442    zeros: Vec<usize>,
443    bit_vectors: Vec<BitVector>,
444    // Binary layers preserve callback coordinates and support weighted queries.
445    quad_vectors: Vec<WaveletMatrixQuadVector>,
446    compress: VecCompress<T>,
447    #[cfg(target_arch = "x86_64")]
448    backend: super::SimdBackend,
449}
450
451impl<T> WaveletMatrix<T>
452where
453    T: Ord + Clone,
454{
455    pub fn new(v: Vec<T>) -> Self {
456        if v.len() <= u32::MAX as usize {
457            #[cfg(target_arch = "x86_64")]
458            let backend = super::simd_backend();
459            Self::from_values(
460                v,
461                |i| i as u32,
462                |i| i as usize,
463                |indices, d| {
464                    #[cfg(target_arch = "x86_64")]
465                    if backend != super::SimdBackend::Scalar && is_x86_feature_detected!("avx2") {
466                        // SAFETY: AVX2 is available.
467                        return unsafe { simd::pack_words(indices, d) };
468                    }
469                    Self::pack_words(indices, |i| (i >> d) & 1 != 0)
470                },
471                |indices, words, zeros, next| {
472                    #[cfg(target_arch = "x86_64")]
473                    if backend != super::SimdBackend::Scalar && is_x86_feature_detected!("avx2") {
474                        // SAFETY: AVX2 is available, and the partition buffers have equal length.
475                        unsafe { simd::partition_avx2(indices, words, zeros, next) };
476                        return;
477                    }
478                    Self::partition(indices, words, zeros, next);
479                },
480            )
481        } else {
482            Self::from_values(
483                v,
484                |i| i,
485                |i| i,
486                |indices, d| Self::pack_words(indices, |i| (i >> d) & 1 != 0),
487                Self::partition,
488            )
489        }
490    }
491
492    fn pack_words<I: Copy>(indices: &[I], bit: impl Fn(I) -> bool) -> Vec<u64> {
493        indices
494            .chunks(64)
495            .map(|chunk| {
496                chunk
497                    .iter()
498                    .enumerate()
499                    .fold(0, |word, (i, &index)| word | ((bit(index) as u64) << i))
500            })
501            .collect()
502    }
503
504    fn partition<I: Copy>(indices: &[I], words: &[u64], mut one: usize, next: &mut [I]) {
505        let mut zero = 0;
506        for (chunk, &word) in indices.chunks(64).zip(words) {
507            if word == 0 {
508                next[zero..zero + chunk.len()].copy_from_slice(chunk);
509                zero += chunk.len();
510            } else if word == u64::MAX {
511                next[one..one + chunk.len()].copy_from_slice(chunk);
512                one += chunk.len();
513            } else {
514                for (i, &index) in chunk.iter().enumerate() {
515                    let bit = (word >> i) & 1 != 0;
516                    next[if bit { one } else { zero }] = index;
517                    zero += !bit as usize;
518                    one += bit as usize;
519                }
520            }
521        }
522    }
523
524    fn from_values<I: Copy>(
525        v: Vec<T>,
526        code: impl Fn(usize) -> I,
527        index: impl Fn(I) -> usize,
528        pack: impl Fn(&[I], usize) -> Vec<u64>,
529        partition: impl Fn(&[I], &[u64], usize, &mut [I]),
530    ) -> Self {
531        let len = v.len();
532        let mut sorted: Vec<_> = v
533            .into_iter()
534            .enumerate()
535            .map(|(i, value)| (value, code(i)))
536            .collect();
537        sorted.sort_unstable_by(|a, b| a.0.cmp(&b.0));
538        let mut values = Vec::with_capacity(len);
539        let mut indices = vec![code(0); len];
540        for (value, i) in sorted {
541            if values.last().is_none_or(|last| last != &value) {
542                values.push(value);
543            }
544            indices[index(i)] = code(values.len() - 1);
545        }
546        let compress = VecCompress::from_sorted_unique(values);
547        let bit_length = usize::BITS as usize - compress.size().leading_zeros() as usize;
548        let mut bit_vectors = Vec::with_capacity(bit_length);
549        let mut zeros = Vec::with_capacity(bit_length);
550        let quad_bits =
551            usize::BITS as usize - compress.size().saturating_sub(1).leading_zeros() as usize;
552        let mut quad_vectors = Vec::with_capacity(quad_bits.div_ceil(2));
553        let mut next = indices.clone();
554        for d in (0..bit_length).rev() {
555            let words = pack(&indices, d);
556            if len <= u32::MAX as usize && d < quad_bits && (d % 2 == 1 || d + 1 == quad_bits) {
557                if d % 2 == 1 {
558                    let low = pack(&indices, d - 1);
559                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&low, Some(&words), len));
560                } else {
561                    quad_vectors.push(WaveletMatrixQuadVector::from_words(&words, None, len));
562                }
563            }
564            let bits = BitVector::from_words(&words, len);
565            let zero_count = bits.rank0(len);
566            if d == 0 {
567                zeros.push(zero_count);
568                bit_vectors.push(bits);
569                break;
570            }
571            partition(&indices, &words, zero_count, &mut next);
572            zeros.push(zero_count);
573            bit_vectors.push(bits);
574            mem::swap(&mut indices, &mut next);
575        }
576        Self {
577            len,
578            bit_length,
579            zeros,
580            bit_vectors,
581            quad_vectors,
582            compress,
583            #[cfg(target_arch = "x86_64")]
584            backend: match super::simd_backend() {
585                super::SimdBackend::Avx512 if !is_x86_feature_detected!("avx512vpopcntdq") => {
586                    super::SimdBackend::Avx2
587                }
588                backend => backend,
589            },
590        }
591    }
592
593    pub fn new_with_init<F>(v: Vec<T>, mut f: F) -> Self
594    where
595        F: FnMut(usize, usize, T),
596    {
597        let this = Self::new(v.clone());
598        if !this.quad_vectors.is_empty() {
599            let bits = usize::BITS as usize
600                - this.compress.size().saturating_sub(1).leading_zeros() as usize;
601            for (mut k, value) in v.into_iter().enumerate() {
602                if this.bit_length > bits {
603                    f(this.bit_length - 1, k, value.clone());
604                }
605                for (level, vector) in this.quad_vectors.iter().enumerate() {
606                    let d = (this.quad_vectors.len() - level - 1) * 2;
607                    let (digit, rank) = vector.access_rank(k);
608                    if d + 1 < bits {
609                        let block = &vector.blocks[k / 64];
610                        let high = block.rank[2] as usize
611                            + block.rank[3] as usize
612                            + (block.hi & !(u64::MAX << (k % 64))).count_ones() as usize;
613                        let middle = if digit & 2 == 0 {
614                            k - high
615                        } else {
616                            this.zeros[this.level(d + 1)] + high
617                        };
618                        f(d + 1, middle, value.clone());
619                    }
620                    k = vector.starts[digit] + rank;
621                    f(d, k, value.clone());
622                }
623            }
624            return this;
625        }
626        for (mut k, value) in v.into_iter().enumerate() {
627            for d in (0..this.bit_length).rev() {
628                let level = this.level(d);
629                let (bit, rank1) = this.bit_vectors[level].access_rank1(k);
630                k = if bit {
631                    this.zeros[level] + rank1
632                } else {
633                    k - rank1
634                };
635                f(d, k, value.clone());
636            }
637        }
638        this
639    }
640
641    fn level(&self, d: usize) -> usize {
642        self.bit_length - 1 - d
643    }
644
645    fn rank1(&self, level: usize, k: usize) -> usize {
646        self.bit_vectors[level].rank1(k)
647    }
648
649    fn reorder<U>(&self, level: usize, current: Vec<U>) -> Vec<U> {
650        assert_eq!(current.len(), self.len);
651        let mut next = Vec::with_capacity(self.len);
652        next.resize_with(self.len, MaybeUninit::uninit);
653        let mut zero = 0;
654        let mut one = self.zeros[level];
655        let mut current = current.into_iter();
656        for block in self.bit_vectors[level].blocks() {
657            let count = current.len().min(64);
658            if block.bits == 0 || block.bits == u64::MAX {
659                let offset = if block.bits == 0 { &mut zero } else { &mut one };
660                for (slot, value) in next[*offset..*offset + count]
661                    .iter_mut()
662                    .zip(current.by_ref().take(count))
663                {
664                    slot.write(value);
665                }
666                *offset += count;
667            } else {
668                for (i, value) in current.by_ref().take(count).enumerate() {
669                    let bit = (block.bits >> i) & 1 != 0;
670                    next[if bit { one } else { zero }].write(value);
671                    zero += !bit as usize;
672                    one += bit as usize;
673                }
674            }
675        }
676        // SAFETY: the partition counts fill every slot once, and `MaybeUninit<U>` has `U`'s layout.
677        unsafe {
678            let mut next = mem::ManuallyDrop::new(next);
679            Vec::from_raw_parts(next.as_mut_ptr().cast(), next.len(), next.capacity())
680        }
681    }
682
683    fn range_by_index(&self, idx: usize, mut range: Range<usize>) -> Range<usize> {
684        if !self.quad_vectors.is_empty() {
685            for (level, vector) in self.quad_vectors.iter().enumerate() {
686                if range.is_empty() {
687                    break;
688                }
689                let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
690                range = vector.starts[digit] + vector.rank(digit, range.start)
691                    ..vector.starts[digit] + vector.rank(digit, range.end);
692            }
693            return range;
694        }
695        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
696            if range.is_empty() {
697                break;
698            }
699            let level = self.level(d);
700            let start1 = self.rank1(level, range.start);
701            let end1 = self.rank1(level, range.end);
702            if ((idx >> d) & 1) != 0 {
703                range.start = self.zeros[level] + start1;
704                range.end = self.zeros[level] + end1;
705            } else {
706                range.start -= start1;
707                range.end -= end1;
708            }
709        }
710        range
711    }
712
713    fn batch<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
714        #[cfg(target_arch = "x86_64")]
715        if !cfg!(target_feature = "popcnt") && is_x86_feature_detected!("popcnt") {
716            // SAFETY: POPCNT is checked above.
717            unsafe {
718                self.batch_popcnt::<OP>(states, count);
719            }
720            return;
721        }
722        self.batch_inner::<OP>(states, count);
723    }
724
725    #[cfg(target_arch = "x86_64")]
726    #[target_feature(enable = "popcnt")]
727    unsafe fn batch_popcnt<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
728        self.batch_inner::<OP>(states, count);
729    }
730
731    #[inline(always)]
732    fn batch_inner<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
733        if OP == RANK_LESSTHAN && self.compress.size() <= 1 {
734            for state in &mut states[..count] {
735                state[3] = if state[2] == 0 {
736                    0
737                } else {
738                    state[1] - state[0]
739                };
740            }
741            return;
742        }
743        if !self.quad_vectors.is_empty() && matches!(OP, ACCESS | RANK | QUANTILE) {
744            #[cfg(target_arch = "x86_64")]
745            if OP == QUANTILE && count >= 8 && self.backend == super::SimdBackend::Avx512 {
746                // SAFETY: checked batch ranges stay within the quad blocks after each
747                // stable partition. Construction restricts quad counters to u32, and
748                // the cached backend includes AVX512F and VPOPCNTDQ.
749                unsafe {
750                    simd::quad_avx512(&self.quad_vectors, &mut states[..count.next_multiple_of(8)]);
751                }
752                return;
753            }
754
755            for (level, vector) in self.quad_vectors.iter().enumerate() {
756                let d = (self.quad_vectors.len() - level - 1) * 2;
757                #[cfg(target_arch = "x86_64")]
758                let next = self
759                    .quad_vectors
760                    .get(level + 1)
761                    .filter(|_| self.len >= 1048576 || (OP == ACCESS && self.len >= 262144));
762                for state in &mut states[..count] {
763                    if OP == ACCESS {
764                        let (digit, rank) = vector.access_rank(state[0]);
765                        state[0] = vector.starts[digit] + rank;
766                        state[3] = state[3] * 4 + digit;
767                    } else if OP == RANK {
768                        let digit = (state[2] >> d) & 3;
769                        state[0] = vector.starts[digit] + vector.rank(digit, state[0]);
770                        state[1] = vector.starts[digit] + vector.rank(digit, state[1]);
771                    } else {
772                        let start = vector.ranks(state[0]);
773                        let end = vector.ranks(state[1]);
774                        let prefix = [
775                            0,
776                            end[0] - start[0],
777                            end[0] + end[1] - start[0] - start[1],
778                            state[1] - state[0] - (end[3] - start[3]),
779                        ];
780                        let digit = (state[2] >= prefix[1]) as usize
781                            + (state[2] >= prefix[2]) as usize
782                            + (state[2] >= prefix[3]) as usize;
783                        state[2] -= prefix[digit];
784                        state[0] = vector.starts[digit] + start[digit];
785                        state[1] = vector.starts[digit] + end[digit];
786                        state[3] = state[3] * 4 + digit;
787                    }
788                    #[cfg(target_arch = "x86_64")]
789                    if let Some(next) = next {
790                        // SAFETY: stable partitions keep both endpoints within the next vector.
791                        unsafe {
792                            std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
793                                next.blocks.as_ptr().add(state[0] / 64).cast(),
794                            );
795                            if OP != ACCESS {
796                                std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
797                                    next.blocks.as_ptr().add(state[1] / 64).cast(),
798                                );
799                            }
800                        }
801                    }
802                }
803            }
804            return;
805        }
806        let first = (matches!(OP, ACCESS | RANK | QUANTILE)
807            && self.compress.size().is_power_of_two()) as usize;
808        #[cfg(target_arch = "x86_64")]
809        if OP == RANK_LESSTHAN && count >= 8 && self.len <= u32::MAX as usize {
810            // SAFETY: the caller checked each initial position. Rank transitions remain in
811            // 0..=len, including the sentinel block. Padding uses position zero. The
812            // backend was detected at construction, and x86-64 blocks contain two u64s.
813            unsafe {
814                match self.backend {
815                    super::SimdBackend::Avx512 => {
816                        simd::rank_lessthan_avx512(
817                            &self.bit_vectors[first..],
818                            &self.zeros[first..],
819                            &mut states[..count.next_multiple_of(8)],
820                        );
821                        return;
822                    }
823                    super::SimdBackend::Avx2 if self.len < 1048576 => {
824                        simd::rank_lessthan_avx2(
825                            &self.bit_vectors[first..],
826                            &self.zeros[first..],
827                            &mut states[..count.next_multiple_of(4)],
828                        );
829                        return;
830                    }
831                    _ => {}
832                }
833            }
834        }
835        for d in (0..self.bit_length - first).rev() {
836            let level = self.level(d);
837            for state in &mut states[..count] {
838                let (bit, start1) = self.bit_vectors[level].access_rank1(state[0]);
839                let start0 = state[0] - start1;
840                let end1 = if OP == ACCESS {
841                    0
842                } else {
843                    self.rank1(level, state[1])
844                };
845                let end0 = if OP == ACCESS { 0 } else { state[1] - end1 };
846                let count0 = if OP == ACCESS { 0 } else { end0 - start0 };
847                let bit = match OP {
848                    ACCESS => bit,
849                    RANK | RANK_LESSTHAN => (state[2] >> d) & 1 != 0,
850                    _ => state[2] >= count0,
851                };
852                state[0] = if bit {
853                    self.zeros[level] + start1
854                } else {
855                    start0
856                };
857                if OP != ACCESS {
858                    state[1] = if bit { self.zeros[level] + end1 } else { end0 };
859                }
860                if OP == ACCESS || OP == QUANTILE {
861                    state[3] |= (bit as usize) << d;
862                }
863                if OP == QUANTILE {
864                    state[2] -= if bit { count0 } else { 0 };
865                }
866                if OP == RANK_LESSTHAN {
867                    state[3] += if bit { count0 } else { 0 };
868                }
869            }
870        }
871    }
872
873    /// get k-th value
874    pub fn access(&self, mut k: usize) -> T {
875        if !self.quad_vectors.is_empty() {
876            let mut index = 0;
877            for vector in &self.quad_vectors {
878                let (digit, rank) = vector.access_rank(k);
879                index = index * 4 + digit;
880                k = vector.starts[digit] + rank;
881            }
882            return self.compress.values()[index].clone();
883        }
884        let mut idx = 0;
885        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
886            let level = self.level(d);
887            let (bit, rank1) = self.bit_vectors[level].access_rank1(k);
888            idx |= (bit as usize) << d;
889            k = if bit {
890                self.zeros[level] + rank1
891            } else {
892                k - rank1
893            };
894        }
895        self.compress.values()[idx].clone()
896    }
897
898    /// Returns the values at `indices` in input order.
899    pub fn access_batch(&self, indices: impl IntoIterator<Item = usize>) -> Vec<T> {
900        let indices: Vec<_> = indices.into_iter().collect();
901        let mut result = Vec::with_capacity(indices.len());
902        for indices in indices.chunks(16) {
903            let mut states = [[0; 4]; 16];
904            for (state, &index) in states.iter_mut().zip(indices) {
905                assert!(index < self.len);
906                state[0] = index;
907            }
908            self.batch::<ACCESS>(&mut states, indices.len());
909            result.extend(
910                states[..indices.len()]
911                    .iter()
912                    .map(|state| self.compress.values()[state[3]].clone()),
913            );
914        }
915        result
916    }
917
918    /// the number of val in range
919    pub fn rank(&self, val: T, range: Range<usize>) -> usize {
920        match self.compress.index_exact(&val) {
921            Some(idx) => self.range_by_index(idx, range).len(),
922            None => 0,
923        }
924    }
925
926    /// Returns the number of exact matches for each `(value, range)` query.
927    pub fn rank_batch(&self, queries: impl IntoIterator<Item = (T, Range<usize>)>) -> Vec<usize> {
928        let queries: Vec<_> = queries.into_iter().collect();
929        let mut result = Vec::with_capacity(queries.len());
930        for queries in queries.chunks(16) {
931            let mut states = [[0; 4]; 16];
932            for (state, (value, range)) in states.iter_mut().zip(queries) {
933                assert!(range.start <= range.end && range.end <= self.len);
934                if let Some(index) = self.compress.index_exact(value) {
935                    *state = [range.start, range.end, index, 0];
936                }
937            }
938            self.batch::<RANK>(&mut states, queries.len());
939            result.extend(
940                states[..queries.len()]
941                    .iter()
942                    .map(|state| state[1] - state[0]),
943            );
944        }
945        result
946    }
947
948    /// index of k-th val
949    pub fn select(&self, val: T, k: usize) -> Option<usize> {
950        let idx = self.compress.index_exact(&val)?;
951        let range = self.range_by_index(idx, 0..self.len);
952        if range.len() <= k {
953            return None;
954        }
955        let mut i = range.start + k;
956        if !self.quad_vectors.is_empty() {
957            for (level, vector) in self.quad_vectors.iter().enumerate().rev() {
958                let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
959                i = vector.select(digit, i - vector.starts[digit]);
960            }
961            return Some(i);
962        }
963        for level in (self.compress.size().is_power_of_two() as usize..self.bit_length).rev() {
964            if i >= self.zeros[level] {
965                i = self.bit_vectors[level]
966                    .select1(i - self.zeros[level])
967                    .unwrap();
968            } else {
969                i = self.bit_vectors[level].select0(i).unwrap();
970            }
971        }
972        Some(i)
973    }
974
975    /// get k-th smallest value in range
976    pub fn quantile(&self, mut range: Range<usize>, mut k: usize) -> T {
977        if !self.quad_vectors.is_empty() {
978            return self.quad_quantile(range, k, 0, 0);
979        }
980        let mut idx = 0;
981        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
982            let level = self.level(d);
983            let start1 = self.rank1(level, range.start);
984            let end1 = self.rank1(level, range.end);
985            let start0 = range.start - start1;
986            let end0 = range.end - end1;
987            let z = end0 - start0;
988            let bit = z <= k;
989            k -= if bit { z } else { 0 };
990            idx |= (bit as usize) << d;
991            range.start = if bit {
992                self.zeros[level] + start1
993            } else {
994                start0
995            };
996            range.end = if bit { self.zeros[level] + end1 } else { end0 };
997        }
998        self.compress.values()[idx].clone()
999    }
1000
1001    #[inline(always)]
1002    fn quad_quantile(
1003        &self,
1004        mut range: Range<usize>,
1005        mut k: usize,
1006        level: usize,
1007        mut index: usize,
1008    ) -> T {
1009        for vector in &self.quad_vectors[level..] {
1010            let start = vector.ranks(range.start);
1011            let end = vector.ranks(range.end);
1012            let mut digit = 0;
1013            while digit < 3 && k >= end[digit] - start[digit] {
1014                k -= end[digit] - start[digit];
1015                digit += 1;
1016            }
1017            index = index * 4 + digit;
1018            range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1019        }
1020        self.compress.values()[index].clone()
1021    }
1022
1023    pub fn quantile_batch(
1024        &self,
1025        queries: impl IntoIterator<Item = (Range<usize>, usize)>,
1026    ) -> Vec<T> {
1027        let queries: Vec<_> = queries.into_iter().collect();
1028        let mut result = Vec::with_capacity(queries.len());
1029        for queries in queries.chunks(16) {
1030            let mut states = [[0; 4]; 16];
1031            for (state, (range, k)) in states.iter_mut().zip(queries) {
1032                assert!(range.start <= range.end && range.end <= self.len && *k < range.len());
1033                *state = [range.start, range.end, *k, 0];
1034            }
1035            if let [(range, k)] = queries {
1036                result.push(self.quantile(range.clone(), *k));
1037                continue;
1038            }
1039            self.batch::<QUANTILE>(&mut states, queries.len());
1040            result.extend(
1041                states[..queries.len()]
1042                    .iter()
1043                    .map(|state| self.compress.values()[state[3]].clone()),
1044            );
1045        }
1046        result
1047    }
1048
1049    /// Returns the requested order statistics of one range. `ranks` must be sorted
1050    /// in nondecreasing order and every rank must be less than the range length.
1051    pub fn quantiles_sorted(&self, range: Range<usize>, ranks: &[usize]) -> Vec<T> {
1052        assert!(range.start <= range.end && range.end <= self.len);
1053        assert!(ranks.windows(2).all(|pair| pair[0] <= pair[1]));
1054        assert!(ranks.last().is_none_or(|&k| k < range.len()));
1055        if ranks.is_empty() {
1056            return Vec::new();
1057        }
1058        if ranks[0] == ranks[ranks.len() - 1] {
1059            return vec![self.quantile(range, ranks[0]); ranks.len()];
1060        }
1061        let use_batch = ranks.len() < 1024;
1062        #[cfg(target_arch = "x86_64")]
1063        let use_batch = use_batch || self.backend == super::SimdBackend::Avx512;
1064        // Sparse requests amortize less of the shared traversal's branching work.
1065        if use_batch && ranks.len() < self.compress.size().div_ceil(16) {
1066            return self.quantile_batch(ranks.iter().map(|&k| (range.clone(), k)));
1067        }
1068        let mut result = Vec::with_capacity(ranks.len());
1069        if !self.quad_vectors.is_empty() {
1070            let mut stack = vec![(range, 0, 0, 0, ranks)];
1071            while let Some((range, base, index, level, ranks)) = stack.pop() {
1072                if level == self.quad_vectors.len() {
1073                    result.resize(
1074                        result.len() + ranks.len(),
1075                        self.compress.values()[index].clone(),
1076                    );
1077                    continue;
1078                }
1079                if ranks.len() == 1 {
1080                    result.push(self.quad_quantile(range, ranks[0] - base, level, index));
1081                    continue;
1082                }
1083                let vector = &self.quad_vectors[level];
1084                let start = vector.ranks(range.start);
1085                let end = vector.ranks(range.end);
1086                let mut split = base + range.len();
1087                let mut requested = ranks.len();
1088                for digit in (0..4).rev() {
1089                    split -= end[digit] - start[digit];
1090                    let boundary = ranks[..requested].partition_point(|&k| k < split);
1091                    if boundary < requested {
1092                        stack.push((
1093                            vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit],
1094                            split,
1095                            index * 4 + digit,
1096                            level + 1,
1097                            &ranks[boundary..requested],
1098                        ));
1099                    }
1100                    requested = boundary;
1101                }
1102            }
1103            return result;
1104        }
1105        let mut stack = vec![(range, 0, 0, self.bit_length, ranks)];
1106        while let Some((range, base, index, bits, ranks)) = stack.pop() {
1107            if bits == 0 {
1108                result.resize(
1109                    result.len() + ranks.len(),
1110                    self.compress.values()[index].clone(),
1111                );
1112                continue;
1113            }
1114            let d = bits - 1;
1115            let level = self.level(d);
1116            let start1 = self.rank1(level, range.start);
1117            let end1 = self.rank1(level, range.end);
1118            let start0 = range.start - start1;
1119            let end0 = range.end - end1;
1120            let split = base + end0 - start0;
1121            let boundary = ranks.partition_point(|&k| k < split);
1122            if boundary < ranks.len() {
1123                stack.push((
1124                    self.zeros[level] + start1..self.zeros[level] + end1,
1125                    split,
1126                    index | (1 << d),
1127                    d,
1128                    &ranks[boundary..],
1129                ));
1130            }
1131            if boundary != 0 {
1132                stack.push((start0..end0, base, index, d, &ranks[..boundary]));
1133            }
1134        }
1135        result
1136    }
1137
1138    /// get k-th smallest value out of range
1139    pub fn quantile_outer(&self, mut range: Range<usize>, mut k: usize) -> T {
1140        if !self.quad_vectors.is_empty() {
1141            let mut outer = 0..self.len;
1142            let mut index = 0;
1143            for vector in &self.quad_vectors {
1144                let start = vector.ranks(range.start);
1145                let end = vector.ranks(range.end);
1146                let outer_start = vector.ranks(outer.start);
1147                let outer_end = vector.ranks(outer.end);
1148                let mut digit = 0;
1149                while digit < 3 {
1150                    let count = outer_end[digit] - outer_start[digit] - (end[digit] - start[digit]);
1151                    if k < count {
1152                        break;
1153                    }
1154                    k -= count;
1155                    digit += 1;
1156                }
1157                index = index * 4 + digit;
1158                range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1159                outer = vector.starts[digit] + outer_start[digit]
1160                    ..vector.starts[digit] + outer_end[digit];
1161            }
1162            return self.compress.values()[index].clone();
1163        }
1164        let mut idx = 0;
1165        let mut orange = 0..self.len;
1166        for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
1167            let level = self.level(d);
1168            let range_start1 = self.rank1(level, range.start);
1169            let range_end1 = self.rank1(level, range.end);
1170            let outer_start1 = self.rank1(level, orange.start);
1171            let outer_end1 = self.rank1(level, orange.end);
1172            let range_start0 = range.start - range_start1;
1173            let range_end0 = range.end - range_end1;
1174            let outer_start0 = orange.start - outer_start1;
1175            let outer_end0 = orange.end - outer_end1;
1176            let z = (outer_end0 - outer_start0) - (range_end0 - range_start0);
1177            if z <= k {
1178                k -= z;
1179                idx |= 1 << d;
1180                range.start = self.zeros[level] + range_start1;
1181                range.end = self.zeros[level] + range_end1;
1182                orange.start = self.zeros[level] + outer_start1;
1183                orange.end = self.zeros[level] + outer_end1;
1184            } else {
1185                range.start = range_start0;
1186                range.end = range_end0;
1187                orange.start = outer_start0;
1188                orange.end = outer_end0;
1189            }
1190        }
1191        self.compress.values()[idx].clone()
1192    }
1193
1194    /// the number of value less than val in range
1195    pub fn rank_lessthan(&self, val: T, range: Range<usize>) -> usize {
1196        let index = self.compress.index_lower_bound(&val);
1197        if index == self.compress.size() {
1198            return range.len();
1199        }
1200        if self.len >= 262144 && !self.quad_vectors.is_empty() {
1201            self.quad_rank_lessthan_index(index, range, 0)
1202        } else {
1203            let bits = usize::BITS as usize
1204                - self.compress.size().saturating_sub(1).leading_zeros() as usize;
1205            self.rank_lessthan_index(index, range, bits)
1206        }
1207    }
1208
1209    /// Returns the number of values below each query's threshold.
1210    pub fn rank_lessthan_batch(
1211        &self,
1212        queries: impl IntoIterator<Item = (T, Range<usize>)>,
1213    ) -> Vec<usize> {
1214        let queries: Vec<_> = queries.into_iter().collect();
1215        let mut result = Vec::with_capacity(queries.len());
1216        for queries in queries.chunks(16) {
1217            let mut states = [[0; 4]; 16];
1218            for (state, (value, range)) in states.iter_mut().zip(queries) {
1219                assert!(range.start <= range.end && range.end <= self.len);
1220                *state = [
1221                    range.start,
1222                    range.end,
1223                    self.compress.index_lower_bound(value),
1224                    0,
1225                ];
1226            }
1227            self.batch::<RANK_LESSTHAN>(&mut states, queries.len());
1228            result.extend(states[..queries.len()].iter().map(|state| state[3]));
1229        }
1230        result
1231    }
1232
1233    fn quad_rank_lessthan_index(
1234        &self,
1235        index: usize,
1236        mut range: Range<usize>,
1237        start_level: usize,
1238    ) -> usize {
1239        let mut result = 0;
1240        for (level, vector) in self.quad_vectors.iter().enumerate().skip(start_level) {
1241            let d = (self.quad_vectors.len() - level - 1) * 2;
1242            if d + 2 <= index.trailing_zeros() as usize {
1243                break;
1244            }
1245            let digit = (index >> d) & 3;
1246            let ranks = |position: usize| {
1247                let block = &vector.blocks[position / 64];
1248                let mask = !(u64::MAX << (position % 64));
1249                let low = if digit & 1 == 0 { !block.lo } else { block.lo };
1250                let high = if digit & 2 == 0 { !block.hi } else { block.hi };
1251                let exact = block.rank[digit] as usize + (low & high & mask).count_ones() as usize;
1252                let less = match digit {
1253                    0 => 0,
1254                    1 => {
1255                        block.rank[0] as usize
1256                            + (!(block.lo | block.hi) & mask).count_ones() as usize
1257                    }
1258                    2 => {
1259                        block.rank[0] as usize
1260                            + block.rank[1] as usize
1261                            + (!block.hi & mask).count_ones() as usize
1262                    }
1263                    _ => position - exact,
1264                };
1265                (exact, less)
1266            };
1267            let (start, start_less) = ranks(range.start);
1268            let (end, end_less) = ranks(range.end);
1269            result += end_less - start_less;
1270            range = vector.starts[digit] + start..vector.starts[digit] + end;
1271        }
1272        result
1273    }
1274
1275    fn rank_lessthan_index(&self, idx: usize, mut range: Range<usize>, bits: usize) -> usize {
1276        let mut res = 0;
1277        for d in (idx.trailing_zeros() as usize..bits).rev() {
1278            let level = self.level(d);
1279            let start1 = self.rank1(level, range.start);
1280            let end1 = self.rank1(level, range.end);
1281            let bit = (idx >> d) & 1 != 0;
1282            let start0 = range.start - start1;
1283            let end0 = range.end - end1;
1284            res += if bit { end0 - start0 } else { 0 };
1285            range.start = if bit {
1286                self.zeros[level] + start1
1287            } else {
1288                start0
1289            };
1290            range.end = if bit { self.zeros[level] + end1 } else { end0 };
1291        }
1292        res
1293    }
1294
1295    /// the number of valrange in range
1296    pub fn rank_range(&self, valrange: Range<T>, mut range: Range<usize>) -> usize {
1297        let lower = self.compress.index_lower_bound(&valrange.start);
1298        let upper = self.compress.index_lower_bound(&valrange.end);
1299        if lower >= upper {
1300            return 0;
1301        }
1302        if !self.quad_vectors.is_empty() {
1303            if upper == self.compress.size() {
1304                return range.len() - self.quad_rank_lessthan_index(lower, range, 0);
1305            }
1306            for (level, vector) in self.quad_vectors.iter().enumerate() {
1307                let d = (self.quad_vectors.len() - level - 1) * 2;
1308                let low = (lower >> d) & 3;
1309                let high = (upper >> d) & 3;
1310                if low == high {
1311                    range = vector.starts[low] + vector.rank(low, range.start)
1312                        ..vector.starts[low] + vector.rank(low, range.end);
1313                    continue;
1314                }
1315                let start = vector.ranks(range.start);
1316                let end = vector.ranks(range.end);
1317                let count: usize = (low..high).map(|digit| end[digit] - start[digit]).sum();
1318                return count
1319                    - self.quad_rank_lessthan_index(
1320                        lower,
1321                        vector.starts[low] + start[low]..vector.starts[low] + end[low],
1322                        level + 1,
1323                    )
1324                    + self.quad_rank_lessthan_index(
1325                        upper,
1326                        vector.starts[high] + start[high]..vector.starts[high] + end[high],
1327                        level + 1,
1328                    );
1329            }
1330            return 0;
1331        }
1332        for d in (0..self.bit_length).rev() {
1333            let level = self.level(d);
1334            let start1 = self.rank1(level, range.start);
1335            let end1 = self.rank1(level, range.end);
1336            let zero_range = range.start - start1..range.end - end1;
1337            let one_range = self.zeros[level] + start1..self.zeros[level] + end1;
1338            if ((lower ^ upper) >> d) & 1 != 0 {
1339                return zero_range.len() - self.rank_lessthan_index(lower, zero_range, d)
1340                    + self.rank_lessthan_index(upper, one_range, d);
1341            }
1342            range = if (lower >> d) & 1 != 0 {
1343                one_range
1344            } else {
1345                zero_range
1346            };
1347        }
1348        0
1349    }
Source

fn select(&self, digit: usize, k: usize) -> usize

Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 959)
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 access_rank(&self, position: usize) -> (usize, usize)

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

Trait Implementations§

Source§

impl Clone for WaveletMatrixQuadVector

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 Debug for WaveletMatrixQuadVector

Source§

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

Formats the value using the given formatter. Read more

Auto Trait Implementations§

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.