pub trait RankSelectDictionaries {
// Required methods
fn bit_length(&self) -> usize;
fn access(&self, k: usize) -> bool;
// Provided methods
fn access_rank1(&self, k: usize) -> (bool, usize) { ... }
fn rank1(&self, k: usize) -> usize { ... }
fn rank0(&self, k: usize) -> usize { ... }
fn select1(&self, k: usize) -> Option<usize> { ... }
fn select0(&self, k: usize) -> Option<usize> { ... }
}Expand description
rank_i(select_i(k)) = k rank_i(select_i(k) + 1) = k + 1
Required Methods§
Provided Methods§
Sourcefn access_rank1(&self, k: usize) -> (bool, usize)
fn access_rank1(&self, k: usize) -> (bool, usize)
Returns the k-th bit and the number of ones before it.
Examples found in repository?
More examples
crates/competitive/src/data_structure/wavelet_matrix.rs (line 629)
593 pub fn new_with_init<F>(v: Vec<T>, mut f: F) -> Self
594 where
595 F: FnMut(usize, usize, T),
596 {
597 let this = Self::new(v.clone());
598 if !this.quad_vectors.is_empty() {
599 let bits = usize::BITS as usize
600 - this.compress.size().saturating_sub(1).leading_zeros() as usize;
601 for (mut k, value) in v.into_iter().enumerate() {
602 if this.bit_length > bits {
603 f(this.bit_length - 1, k, value.clone());
604 }
605 for (level, vector) in this.quad_vectors.iter().enumerate() {
606 let d = (this.quad_vectors.len() - level - 1) * 2;
607 let (digit, rank) = vector.access_rank(k);
608 if d + 1 < bits {
609 let block = &vector.blocks[k / 64];
610 let high = block.rank[2] as usize
611 + block.rank[3] as usize
612 + (block.hi & !(u64::MAX << (k % 64))).count_ones() as usize;
613 let middle = if digit & 2 == 0 {
614 k - high
615 } else {
616 this.zeros[this.level(d + 1)] + high
617 };
618 f(d + 1, middle, value.clone());
619 }
620 k = vector.starts[digit] + rank;
621 f(d, k, value.clone());
622 }
623 }
624 return this;
625 }
626 for (mut k, value) in v.into_iter().enumerate() {
627 for d in (0..this.bit_length).rev() {
628 let level = this.level(d);
629 let (bit, rank1) = this.bit_vectors[level].access_rank1(k);
630 k = if bit {
631 this.zeros[level] + rank1
632 } else {
633 k - rank1
634 };
635 f(d, k, value.clone());
636 }
637 }
638 this
639 }
640
641 fn level(&self, d: usize) -> usize {
642 self.bit_length - 1 - d
643 }
644
645 fn rank1(&self, level: usize, k: usize) -> usize {
646 self.bit_vectors[level].rank1(k)
647 }
648
649 fn reorder<U>(&self, level: usize, current: Vec<U>) -> Vec<U> {
650 assert_eq!(current.len(), self.len);
651 let mut next = Vec::with_capacity(self.len);
652 next.resize_with(self.len, MaybeUninit::uninit);
653 let mut zero = 0;
654 let mut one = self.zeros[level];
655 let mut current = current.into_iter();
656 for block in self.bit_vectors[level].blocks() {
657 let count = current.len().min(64);
658 if block.bits == 0 || block.bits == u64::MAX {
659 let offset = if block.bits == 0 { &mut zero } else { &mut one };
660 for (slot, value) in next[*offset..*offset + count]
661 .iter_mut()
662 .zip(current.by_ref().take(count))
663 {
664 slot.write(value);
665 }
666 *offset += count;
667 } else {
668 for (i, value) in current.by_ref().take(count).enumerate() {
669 let bit = (block.bits >> i) & 1 != 0;
670 next[if bit { one } else { zero }].write(value);
671 zero += !bit as usize;
672 one += bit as usize;
673 }
674 }
675 }
676 // SAFETY: the partition counts fill every slot once, and `MaybeUninit<U>` has `U`'s layout.
677 unsafe {
678 let mut next = mem::ManuallyDrop::new(next);
679 Vec::from_raw_parts(next.as_mut_ptr().cast(), next.len(), next.capacity())
680 }
681 }
682
683 fn range_by_index(&self, idx: usize, mut range: Range<usize>) -> Range<usize> {
684 if !self.quad_vectors.is_empty() {
685 for (level, vector) in self.quad_vectors.iter().enumerate() {
686 if range.is_empty() {
687 break;
688 }
689 let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
690 range = vector.starts[digit] + vector.rank(digit, range.start)
691 ..vector.starts[digit] + vector.rank(digit, range.end);
692 }
693 return range;
694 }
695 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
696 if range.is_empty() {
697 break;
698 }
699 let level = self.level(d);
700 let start1 = self.rank1(level, range.start);
701 let end1 = self.rank1(level, range.end);
702 if ((idx >> d) & 1) != 0 {
703 range.start = self.zeros[level] + start1;
704 range.end = self.zeros[level] + end1;
705 } else {
706 range.start -= start1;
707 range.end -= end1;
708 }
709 }
710 range
711 }
712
713 fn batch<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
714 #[cfg(target_arch = "x86_64")]
715 if !cfg!(target_feature = "popcnt") && is_x86_feature_detected!("popcnt") {
716 // SAFETY: POPCNT is checked above.
717 unsafe {
718 self.batch_popcnt::<OP>(states, count);
719 }
720 return;
721 }
722 self.batch_inner::<OP>(states, count);
723 }
724
725 #[cfg(target_arch = "x86_64")]
726 #[target_feature(enable = "popcnt")]
727 unsafe fn batch_popcnt<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
728 self.batch_inner::<OP>(states, count);
729 }
730
731 #[inline(always)]
732 fn batch_inner<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
733 if OP == RANK_LESSTHAN && self.compress.size() <= 1 {
734 for state in &mut states[..count] {
735 state[3] = if state[2] == 0 {
736 0
737 } else {
738 state[1] - state[0]
739 };
740 }
741 return;
742 }
743 if !self.quad_vectors.is_empty() && matches!(OP, ACCESS | RANK | QUANTILE) {
744 #[cfg(target_arch = "x86_64")]
745 if OP == QUANTILE && count >= 8 && self.backend == super::SimdBackend::Avx512 {
746 // SAFETY: checked batch ranges stay within the quad blocks after each
747 // stable partition. Construction restricts quad counters to u32, and
748 // the cached backend includes AVX512F and VPOPCNTDQ.
749 unsafe {
750 simd::quad_avx512(&self.quad_vectors, &mut states[..count.next_multiple_of(8)]);
751 }
752 return;
753 }
754
755 for (level, vector) in self.quad_vectors.iter().enumerate() {
756 let d = (self.quad_vectors.len() - level - 1) * 2;
757 #[cfg(target_arch = "x86_64")]
758 let next = self
759 .quad_vectors
760 .get(level + 1)
761 .filter(|_| self.len >= 1048576 || (OP == ACCESS && self.len >= 262144));
762 for state in &mut states[..count] {
763 if OP == ACCESS {
764 let (digit, rank) = vector.access_rank(state[0]);
765 state[0] = vector.starts[digit] + rank;
766 state[3] = state[3] * 4 + digit;
767 } else if OP == RANK {
768 let digit = (state[2] >> d) & 3;
769 state[0] = vector.starts[digit] + vector.rank(digit, state[0]);
770 state[1] = vector.starts[digit] + vector.rank(digit, state[1]);
771 } else {
772 let start = vector.ranks(state[0]);
773 let end = vector.ranks(state[1]);
774 let prefix = [
775 0,
776 end[0] - start[0],
777 end[0] + end[1] - start[0] - start[1],
778 state[1] - state[0] - (end[3] - start[3]),
779 ];
780 let digit = (state[2] >= prefix[1]) as usize
781 + (state[2] >= prefix[2]) as usize
782 + (state[2] >= prefix[3]) as usize;
783 state[2] -= prefix[digit];
784 state[0] = vector.starts[digit] + start[digit];
785 state[1] = vector.starts[digit] + end[digit];
786 state[3] = state[3] * 4 + digit;
787 }
788 #[cfg(target_arch = "x86_64")]
789 if let Some(next) = next {
790 // SAFETY: stable partitions keep both endpoints within the next vector.
791 unsafe {
792 std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
793 next.blocks.as_ptr().add(state[0] / 64).cast(),
794 );
795 if OP != ACCESS {
796 std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
797 next.blocks.as_ptr().add(state[1] / 64).cast(),
798 );
799 }
800 }
801 }
802 }
803 }
804 return;
805 }
806 let first = (matches!(OP, ACCESS | RANK | QUANTILE)
807 && self.compress.size().is_power_of_two()) as usize;
808 #[cfg(target_arch = "x86_64")]
809 if OP == RANK_LESSTHAN && count >= 8 && self.len <= u32::MAX as usize {
810 // SAFETY: the caller checked each initial position. Rank transitions remain in
811 // 0..=len, including the sentinel block. Padding uses position zero. The
812 // backend was detected at construction, and x86-64 blocks contain two u64s.
813 unsafe {
814 match self.backend {
815 super::SimdBackend::Avx512 => {
816 simd::rank_lessthan_avx512(
817 &self.bit_vectors[first..],
818 &self.zeros[first..],
819 &mut states[..count.next_multiple_of(8)],
820 );
821 return;
822 }
823 super::SimdBackend::Avx2 if self.len < 1048576 => {
824 simd::rank_lessthan_avx2(
825 &self.bit_vectors[first..],
826 &self.zeros[first..],
827 &mut states[..count.next_multiple_of(4)],
828 );
829 return;
830 }
831 _ => {}
832 }
833 }
834 }
835 for d in (0..self.bit_length - first).rev() {
836 let level = self.level(d);
837 for state in &mut states[..count] {
838 let (bit, start1) = self.bit_vectors[level].access_rank1(state[0]);
839 let start0 = state[0] - start1;
840 let end1 = if OP == ACCESS {
841 0
842 } else {
843 self.rank1(level, state[1])
844 };
845 let end0 = if OP == ACCESS { 0 } else { state[1] - end1 };
846 let count0 = if OP == ACCESS { 0 } else { end0 - start0 };
847 let bit = match OP {
848 ACCESS => bit,
849 RANK | RANK_LESSTHAN => (state[2] >> d) & 1 != 0,
850 _ => state[2] >= count0,
851 };
852 state[0] = if bit {
853 self.zeros[level] + start1
854 } else {
855 start0
856 };
857 if OP != ACCESS {
858 state[1] = if bit { self.zeros[level] + end1 } else { end0 };
859 }
860 if OP == ACCESS || OP == QUANTILE {
861 state[3] |= (bit as usize) << d;
862 }
863 if OP == QUANTILE {
864 state[2] -= if bit { count0 } else { 0 };
865 }
866 if OP == RANK_LESSTHAN {
867 state[3] += if bit { count0 } else { 0 };
868 }
869 }
870 }
871 }
872
873 /// get k-th value
874 pub fn access(&self, mut k: usize) -> T {
875 if !self.quad_vectors.is_empty() {
876 let mut index = 0;
877 for vector in &self.quad_vectors {
878 let (digit, rank) = vector.access_rank(k);
879 index = index * 4 + digit;
880 k = vector.starts[digit] + rank;
881 }
882 return self.compress.values()[index].clone();
883 }
884 let mut idx = 0;
885 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
886 let level = self.level(d);
887 let (bit, rank1) = self.bit_vectors[level].access_rank1(k);
888 idx |= (bit as usize) << d;
889 k = if bit {
890 self.zeros[level] + rank1
891 } else {
892 k - rank1
893 };
894 }
895 self.compress.values()[idx].clone()
896 }
897
898 /// Returns the values at `indices` in input order.
899 pub fn access_batch(&self, indices: impl IntoIterator<Item = usize>) -> Vec<T> {
900 let indices: Vec<_> = indices.into_iter().collect();
901 let mut result = Vec::with_capacity(indices.len());
902 for indices in indices.chunks(16) {
903 let mut states = [[0; 4]; 16];
904 for (state, &index) in states.iter_mut().zip(indices) {
905 assert!(index < self.len);
906 state[0] = index;
907 }
908 self.batch::<ACCESS>(&mut states, indices.len());
909 result.extend(
910 states[..indices.len()]
911 .iter()
912 .map(|state| self.compress.values()[state[3]].clone()),
913 );
914 }
915 result
916 }
917
918 /// the number of val in range
919 pub fn rank(&self, val: T, range: Range<usize>) -> usize {
920 match self.compress.index_exact(&val) {
921 Some(idx) => self.range_by_index(idx, range).len(),
922 None => 0,
923 }
924 }
925
926 /// Returns the number of exact matches for each `(value, range)` query.
927 pub fn rank_batch(&self, queries: impl IntoIterator<Item = (T, Range<usize>)>) -> Vec<usize> {
928 let queries: Vec<_> = queries.into_iter().collect();
929 let mut result = Vec::with_capacity(queries.len());
930 for queries in queries.chunks(16) {
931 let mut states = [[0; 4]; 16];
932 for (state, (value, range)) in states.iter_mut().zip(queries) {
933 assert!(range.start <= range.end && range.end <= self.len);
934 if let Some(index) = self.compress.index_exact(value) {
935 *state = [range.start, range.end, index, 0];
936 }
937 }
938 self.batch::<RANK>(&mut states, queries.len());
939 result.extend(
940 states[..queries.len()]
941 .iter()
942 .map(|state| state[1] - state[0]),
943 );
944 }
945 result
946 }
947
948 /// index of k-th val
949 pub fn select(&self, val: T, k: usize) -> Option<usize> {
950 let idx = self.compress.index_exact(&val)?;
951 let range = self.range_by_index(idx, 0..self.len);
952 if range.len() <= k {
953 return None;
954 }
955 let mut i = range.start + k;
956 if !self.quad_vectors.is_empty() {
957 for (level, vector) in self.quad_vectors.iter().enumerate().rev() {
958 let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
959 i = vector.select(digit, i - vector.starts[digit]);
960 }
961 return Some(i);
962 }
963 for level in (self.compress.size().is_power_of_two() as usize..self.bit_length).rev() {
964 if i >= self.zeros[level] {
965 i = self.bit_vectors[level]
966 .select1(i - self.zeros[level])
967 .unwrap();
968 } else {
969 i = self.bit_vectors[level].select0(i).unwrap();
970 }
971 }
972 Some(i)
973 }
974
975 /// get k-th smallest value in range
976 pub fn quantile(&self, mut range: Range<usize>, mut k: usize) -> T {
977 if !self.quad_vectors.is_empty() {
978 return self.quad_quantile(range, k, 0, 0);
979 }
980 let mut idx = 0;
981 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
982 let level = self.level(d);
983 let start1 = self.rank1(level, range.start);
984 let end1 = self.rank1(level, range.end);
985 let start0 = range.start - start1;
986 let end0 = range.end - end1;
987 let z = end0 - start0;
988 let bit = z <= k;
989 k -= if bit { z } else { 0 };
990 idx |= (bit as usize) << d;
991 range.start = if bit {
992 self.zeros[level] + start1
993 } else {
994 start0
995 };
996 range.end = if bit { self.zeros[level] + end1 } else { end0 };
997 }
998 self.compress.values()[idx].clone()
999 }
1000
1001 #[inline(always)]
1002 fn quad_quantile(
1003 &self,
1004 mut range: Range<usize>,
1005 mut k: usize,
1006 level: usize,
1007 mut index: usize,
1008 ) -> T {
1009 for vector in &self.quad_vectors[level..] {
1010 let start = vector.ranks(range.start);
1011 let end = vector.ranks(range.end);
1012 let mut digit = 0;
1013 while digit < 3 && k >= end[digit] - start[digit] {
1014 k -= end[digit] - start[digit];
1015 digit += 1;
1016 }
1017 index = index * 4 + digit;
1018 range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1019 }
1020 self.compress.values()[index].clone()
1021 }
1022
1023 pub fn quantile_batch(
1024 &self,
1025 queries: impl IntoIterator<Item = (Range<usize>, usize)>,
1026 ) -> Vec<T> {
1027 let queries: Vec<_> = queries.into_iter().collect();
1028 let mut result = Vec::with_capacity(queries.len());
1029 for queries in queries.chunks(16) {
1030 let mut states = [[0; 4]; 16];
1031 for (state, (range, k)) in states.iter_mut().zip(queries) {
1032 assert!(range.start <= range.end && range.end <= self.len && *k < range.len());
1033 *state = [range.start, range.end, *k, 0];
1034 }
1035 if let [(range, k)] = queries {
1036 result.push(self.quantile(range.clone(), *k));
1037 continue;
1038 }
1039 self.batch::<QUANTILE>(&mut states, queries.len());
1040 result.extend(
1041 states[..queries.len()]
1042 .iter()
1043 .map(|state| self.compress.values()[state[3]].clone()),
1044 );
1045 }
1046 result
1047 }
1048
1049 /// Returns the requested order statistics of one range. `ranks` must be sorted
1050 /// in nondecreasing order and every rank must be less than the range length.
1051 pub fn quantiles_sorted(&self, range: Range<usize>, ranks: &[usize]) -> Vec<T> {
1052 assert!(range.start <= range.end && range.end <= self.len);
1053 assert!(ranks.windows(2).all(|pair| pair[0] <= pair[1]));
1054 assert!(ranks.last().is_none_or(|&k| k < range.len()));
1055 if ranks.is_empty() {
1056 return Vec::new();
1057 }
1058 if ranks[0] == ranks[ranks.len() - 1] {
1059 return vec![self.quantile(range, ranks[0]); ranks.len()];
1060 }
1061 let use_batch = ranks.len() < 1024;
1062 #[cfg(target_arch = "x86_64")]
1063 let use_batch = use_batch || self.backend == super::SimdBackend::Avx512;
1064 // Sparse requests amortize less of the shared traversal's branching work.
1065 if use_batch && ranks.len() < self.compress.size().div_ceil(16) {
1066 return self.quantile_batch(ranks.iter().map(|&k| (range.clone(), k)));
1067 }
1068 let mut result = Vec::with_capacity(ranks.len());
1069 if !self.quad_vectors.is_empty() {
1070 let mut stack = vec![(range, 0, 0, 0, ranks)];
1071 while let Some((range, base, index, level, ranks)) = stack.pop() {
1072 if level == self.quad_vectors.len() {
1073 result.resize(
1074 result.len() + ranks.len(),
1075 self.compress.values()[index].clone(),
1076 );
1077 continue;
1078 }
1079 if ranks.len() == 1 {
1080 result.push(self.quad_quantile(range, ranks[0] - base, level, index));
1081 continue;
1082 }
1083 let vector = &self.quad_vectors[level];
1084 let start = vector.ranks(range.start);
1085 let end = vector.ranks(range.end);
1086 let mut split = base + range.len();
1087 let mut requested = ranks.len();
1088 for digit in (0..4).rev() {
1089 split -= end[digit] - start[digit];
1090 let boundary = ranks[..requested].partition_point(|&k| k < split);
1091 if boundary < requested {
1092 stack.push((
1093 vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit],
1094 split,
1095 index * 4 + digit,
1096 level + 1,
1097 &ranks[boundary..requested],
1098 ));
1099 }
1100 requested = boundary;
1101 }
1102 }
1103 return result;
1104 }
1105 let mut stack = vec![(range, 0, 0, self.bit_length, ranks)];
1106 while let Some((range, base, index, bits, ranks)) = stack.pop() {
1107 if bits == 0 {
1108 result.resize(
1109 result.len() + ranks.len(),
1110 self.compress.values()[index].clone(),
1111 );
1112 continue;
1113 }
1114 let d = bits - 1;
1115 let level = self.level(d);
1116 let start1 = self.rank1(level, range.start);
1117 let end1 = self.rank1(level, range.end);
1118 let start0 = range.start - start1;
1119 let end0 = range.end - end1;
1120 let split = base + end0 - start0;
1121 let boundary = ranks.partition_point(|&k| k < split);
1122 if boundary < ranks.len() {
1123 stack.push((
1124 self.zeros[level] + start1..self.zeros[level] + end1,
1125 split,
1126 index | (1 << d),
1127 d,
1128 &ranks[boundary..],
1129 ));
1130 }
1131 if boundary != 0 {
1132 stack.push((start0..end0, base, index, d, &ranks[..boundary]));
1133 }
1134 }
1135 result
1136 }
1137
1138 /// get k-th smallest value out of range
1139 pub fn quantile_outer(&self, mut range: Range<usize>, mut k: usize) -> T {
1140 if !self.quad_vectors.is_empty() {
1141 let mut outer = 0..self.len;
1142 let mut index = 0;
1143 for vector in &self.quad_vectors {
1144 let start = vector.ranks(range.start);
1145 let end = vector.ranks(range.end);
1146 let outer_start = vector.ranks(outer.start);
1147 let outer_end = vector.ranks(outer.end);
1148 let mut digit = 0;
1149 while digit < 3 {
1150 let count = outer_end[digit] - outer_start[digit] - (end[digit] - start[digit]);
1151 if k < count {
1152 break;
1153 }
1154 k -= count;
1155 digit += 1;
1156 }
1157 index = index * 4 + digit;
1158 range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1159 outer = vector.starts[digit] + outer_start[digit]
1160 ..vector.starts[digit] + outer_end[digit];
1161 }
1162 return self.compress.values()[index].clone();
1163 }
1164 let mut idx = 0;
1165 let mut orange = 0..self.len;
1166 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
1167 let level = self.level(d);
1168 let range_start1 = self.rank1(level, range.start);
1169 let range_end1 = self.rank1(level, range.end);
1170 let outer_start1 = self.rank1(level, orange.start);
1171 let outer_end1 = self.rank1(level, orange.end);
1172 let range_start0 = range.start - range_start1;
1173 let range_end0 = range.end - range_end1;
1174 let outer_start0 = orange.start - outer_start1;
1175 let outer_end0 = orange.end - outer_end1;
1176 let z = (outer_end0 - outer_start0) - (range_end0 - range_start0);
1177 if z <= k {
1178 k -= z;
1179 idx |= 1 << d;
1180 range.start = self.zeros[level] + range_start1;
1181 range.end = self.zeros[level] + range_end1;
1182 orange.start = self.zeros[level] + outer_start1;
1183 orange.end = self.zeros[level] + outer_end1;
1184 } else {
1185 range.start = range_start0;
1186 range.end = range_end0;
1187 orange.start = outer_start0;
1188 orange.end = outer_end0;
1189 }
1190 }
1191 self.compress.values()[idx].clone()
1192 }
1193
1194 /// the number of value less than val in range
1195 pub fn rank_lessthan(&self, val: T, range: Range<usize>) -> usize {
1196 let index = self.compress.index_lower_bound(&val);
1197 if index == self.compress.size() {
1198 return range.len();
1199 }
1200 if self.len >= 262144 && !self.quad_vectors.is_empty() {
1201 self.quad_rank_lessthan_index(index, range, 0)
1202 } else {
1203 let bits = usize::BITS as usize
1204 - self.compress.size().saturating_sub(1).leading_zeros() as usize;
1205 self.rank_lessthan_index(index, range, bits)
1206 }
1207 }
1208
1209 /// Returns the number of values below each query's threshold.
1210 pub fn rank_lessthan_batch(
1211 &self,
1212 queries: impl IntoIterator<Item = (T, Range<usize>)>,
1213 ) -> Vec<usize> {
1214 let queries: Vec<_> = queries.into_iter().collect();
1215 let mut result = Vec::with_capacity(queries.len());
1216 for queries in queries.chunks(16) {
1217 let mut states = [[0; 4]; 16];
1218 for (state, (value, range)) in states.iter_mut().zip(queries) {
1219 assert!(range.start <= range.end && range.end <= self.len);
1220 *state = [
1221 range.start,
1222 range.end,
1223 self.compress.index_lower_bound(value),
1224 0,
1225 ];
1226 }
1227 self.batch::<RANK_LESSTHAN>(&mut states, queries.len());
1228 result.extend(states[..queries.len()].iter().map(|state| state[3]));
1229 }
1230 result
1231 }
1232
1233 fn quad_rank_lessthan_index(
1234 &self,
1235 index: usize,
1236 mut range: Range<usize>,
1237 start_level: usize,
1238 ) -> usize {
1239 let mut result = 0;
1240 for (level, vector) in self.quad_vectors.iter().enumerate().skip(start_level) {
1241 let d = (self.quad_vectors.len() - level - 1) * 2;
1242 if d + 2 <= index.trailing_zeros() as usize {
1243 break;
1244 }
1245 let digit = (index >> d) & 3;
1246 let ranks = |position: usize| {
1247 let block = &vector.blocks[position / 64];
1248 let mask = !(u64::MAX << (position % 64));
1249 let low = if digit & 1 == 0 { !block.lo } else { block.lo };
1250 let high = if digit & 2 == 0 { !block.hi } else { block.hi };
1251 let exact = block.rank[digit] as usize + (low & high & mask).count_ones() as usize;
1252 let less = match digit {
1253 0 => 0,
1254 1 => {
1255 block.rank[0] as usize
1256 + (!(block.lo | block.hi) & mask).count_ones() as usize
1257 }
1258 2 => {
1259 block.rank[0] as usize
1260 + block.rank[1] as usize
1261 + (!block.hi & mask).count_ones() as usize
1262 }
1263 _ => position - exact,
1264 };
1265 (exact, less)
1266 };
1267 let (start, start_less) = ranks(range.start);
1268 let (end, end_less) = ranks(range.end);
1269 result += end_less - start_less;
1270 range = vector.starts[digit] + start..vector.starts[digit] + end;
1271 }
1272 result
1273 }
1274
1275 fn rank_lessthan_index(&self, idx: usize, mut range: Range<usize>, bits: usize) -> usize {
1276 let mut res = 0;
1277 for d in (idx.trailing_zeros() as usize..bits).rev() {
1278 let level = self.level(d);
1279 let start1 = self.rank1(level, range.start);
1280 let end1 = self.rank1(level, range.end);
1281 let bit = (idx >> d) & 1 != 0;
1282 let start0 = range.start - start1;
1283 let end0 = range.end - end1;
1284 res += if bit { end0 - start0 } else { 0 };
1285 range.start = if bit {
1286 self.zeros[level] + start1
1287 } else {
1288 start0
1289 };
1290 range.end = if bit { self.zeros[level] + end1 } else { end0 };
1291 }
1292 res
1293 }
1294
1295 /// the number of valrange in range
1296 pub fn rank_range(&self, valrange: Range<T>, mut range: Range<usize>) -> usize {
1297 let lower = self.compress.index_lower_bound(&valrange.start);
1298 let upper = self.compress.index_lower_bound(&valrange.end);
1299 if lower >= upper {
1300 return 0;
1301 }
1302 if !self.quad_vectors.is_empty() {
1303 if upper == self.compress.size() {
1304 return range.len() - self.quad_rank_lessthan_index(lower, range, 0);
1305 }
1306 for (level, vector) in self.quad_vectors.iter().enumerate() {
1307 let d = (self.quad_vectors.len() - level - 1) * 2;
1308 let low = (lower >> d) & 3;
1309 let high = (upper >> d) & 3;
1310 if low == high {
1311 range = vector.starts[low] + vector.rank(low, range.start)
1312 ..vector.starts[low] + vector.rank(low, range.end);
1313 continue;
1314 }
1315 let start = vector.ranks(range.start);
1316 let end = vector.ranks(range.end);
1317 let count: usize = (low..high).map(|digit| end[digit] - start[digit]).sum();
1318 return count
1319 - self.quad_rank_lessthan_index(
1320 lower,
1321 vector.starts[low] + start[low]..vector.starts[low] + end[low],
1322 level + 1,
1323 )
1324 + self.quad_rank_lessthan_index(
1325 upper,
1326 vector.starts[high] + start[high]..vector.starts[high] + end[high],
1327 level + 1,
1328 );
1329 }
1330 return 0;
1331 }
1332 for d in (0..self.bit_length).rev() {
1333 let level = self.level(d);
1334 let start1 = self.rank1(level, range.start);
1335 let end1 = self.rank1(level, range.end);
1336 let zero_range = range.start - start1..range.end - end1;
1337 let one_range = self.zeros[level] + start1..self.zeros[level] + end1;
1338 if ((lower ^ upper) >> d) & 1 != 0 {
1339 return zero_range.len() - self.rank_lessthan_index(lower, zero_range, d)
1340 + self.rank_lessthan_index(upper, one_range, d);
1341 }
1342 range = if (lower >> d) & 1 != 0 {
1343 one_range
1344 } else {
1345 zero_range
1346 };
1347 }
1348 0
1349 }
1350
1351 /// Returns each query's count in a half-open value range.
1352 pub fn rank_range_batch(
1353 &self,
1354 queries: impl IntoIterator<Item = (Range<T>, Range<usize>)>,
1355 ) -> Vec<usize> {
1356 let queries: Vec<_> = queries.into_iter().collect();
1357 let mut result = Vec::with_capacity(queries.len());
1358 for queries in queries.chunks(8) {
1359 let mut states = [[0; 4]; 16];
1360 for (i, (values, range)) in queries.iter().enumerate() {
1361 assert!(range.start <= range.end && range.end <= self.len);
1362 let lower = self.compress.index_lower_bound(&values.start);
1363 let upper = self.compress.index_lower_bound(&values.end);
1364 if lower < upper {
1365 states[i * 2] = [range.start, range.end, lower, 0];
1366 states[i * 2 + 1] = [range.start, range.end, upper, 0];
1367 }
1368 }
1369 self.batch::<RANK_LESSTHAN>(&mut states, queries.len() * 2);
1370 result.extend(
1371 states[..queries.len() * 2]
1372 .as_chunks::<2>()
1373 .0
1374 .iter()
1375 .map(|pair| pair[1][3] - pair[0][3]),
1376 );
1377 }
1378 result
1379 }
1380
1381 pub fn query_less_than<F>(&self, val: T, mut range: Range<usize>, mut f: F)
1382 where
1383 F: FnMut(usize, Range<usize>),
1384 {
1385 let idx = self.compress.index_lower_bound(&val);
1386 if !self.quad_vectors.is_empty() {
1387 if idx == self.compress.size() && idx.is_power_of_two() {
1388 f(self.bit_length - 1, range);
1389 return;
1390 }
1391 for (level, vector) in self.quad_vectors.iter().enumerate().take(
1392 self.quad_vectors
1393 .len()
1394 .saturating_sub(idx.trailing_zeros() as usize / 2),
1395 ) {
1396 let d = (self.quad_vectors.len() - level - 1) * 2;
1397 let digit = (idx >> d) & 3;
1398 let start = vector.ranks(range.start);
1399 let end = vector.ranks(range.end);
1400 if digit & 2 != 0 {
1401 f(d + 1, start[0] + start[1]..end[0] + end[1]);
1402 }
1403 if digit & 1 != 0 {
1404 let zero = digit & 2;
1405 f(
1406 d,
1407 vector.starts[zero] + start[zero]..vector.starts[zero] + end[zero],
1408 );
1409 }
1410 range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1411 }
1412 return;
1413 }
1414 for d in (idx.trailing_zeros() as usize..self.bit_length).rev() {
1415 let level = self.level(d);
1416 let start1 = self.rank1(level, range.start);
1417 let end1 = self.rank1(level, range.end);
1418 let start0 = range.start - start1;
1419 let end0 = range.end - end1;
1420 if ((idx >> d) & 1) != 0 {
1421 f(d, start0..end0);
1422 range.start = self.zeros[level] + start1;
1423 range.end = self.zeros[level] + end1;
1424 } else {
1425 range.start = start0;
1426 range.end = end0;
1427 }
1428 }
1429 }
1430
1431 pub fn build_fold<M>(&self, weights: &[M::T]) -> WaveletMatrixFold<'_, T, M>
1432 where
1433 M: AbelianGroup,
1434 {
1435 assert_eq!(weights.len(), self.len);
1436 let mut offsets = Vec::with_capacity(self.bit_length);
1437 let mut prefix = Vec::with_capacity(self.zeros.iter().map(|&zero| zero + 1).sum());
1438 let mut current: Vec<M::T> = weights.to_vec();
1439 for level in 0..self.bit_length {
1440 current = self.reorder(level, current);
1441 offsets.push(prefix.len());
1442 let mut acc = M::unit();
1443 prefix.push(acc.clone());
1444 for w in ¤t[..self.zeros[level]] {
1445 acc = M::operate(&acc, w);
1446 prefix.push(acc.clone());
1447 }
1448 }
1449 WaveletMatrixFold {
1450 wavelet_matrix: self,
1451 prefix,
1452 offsets,
1453 }
1454 }
1455
1456 pub fn build_point_add<M>(&self, weights: &[M::T]) -> WaveletMatrixPointAdd<'_, T, M>
1457 where
1458 M: AbelianGroup,
1459 {
1460 assert_eq!(weights.len(), self.len);
1461 let mut current = weights.to_vec();
1462 let mut bits = Vec::with_capacity(self.bit_length);
1463 for level in 0..self.bit_length {
1464 current = self.reorder(level, current);
1465 bits.push(BinaryIndexedTree::from_slice(¤t[..self.zeros[level]]));
1466 }
1467 WaveletMatrixPointAdd {
1468 wavelet_matrix: self,
1469 bits,
1470 }
1471 }
1472}
1473
1474pub struct WaveletMatrixPointAdd<'a, T, M>
1475where
1476 T: Ord + Clone,
1477 M: AbelianGroup,
1478{
1479 wavelet_matrix: &'a WaveletMatrix<T>,
1480 bits: Vec<BinaryIndexedTree<M>>,
1481}
1482
1483impl<'a, T, M> WaveletMatrixPointAdd<'a, T, M>
1484where
1485 T: Ord + Clone,
1486 M: AbelianGroup,
1487{
1488 pub fn update(&mut self, mut index: usize, value: M::T) {
1489 debug_assert!(index < self.wavelet_matrix.len);
1490 for d in (0..self.wavelet_matrix.bit_length).rev() {
1491 let level = self.wavelet_matrix.level(d);
1492 let (bit, rank1) = self.wavelet_matrix.bit_vectors[level].access_rank1(index);
1493 if bit {
1494 index = self.wavelet_matrix.zeros[level] + rank1;
1495 } else {
1496 index -= rank1;
1497 self.bits[level].update(index, value.clone());
1498 }
1499 }
1500 }Sourcefn rank1(&self, k: usize) -> usize
fn rank1(&self, k: usize) -> usize
the number of 1 in [0, k)
Examples found in repository?
crates/competitive/src/data_structure/bit_vector.rs (line 11)
10 fn access_rank1(&self, k: usize) -> (bool, usize) {
11 (self.access(k), self.rank1(k))
12 }
13 /// the number of 1 in [0, k)
14 fn rank1(&self, k: usize) -> usize {
15 (0..k).filter(|&i| self.access(i)).count()
16 }
17 /// the number of 0 in [0, k)
18 fn rank0(&self, k: usize) -> usize {
19 k - self.rank1(k)
20 }
21 /// index of k-th 1
22 fn select1(&self, k: usize) -> Option<usize> {
23 let n = self.bit_length();
24 if self.rank1(n) <= k {
25 return None;
26 }
27 let (mut l, mut r) = (0, n);
28 while r - l > 1 {
29 let m = l.midpoint(r);
30 if self.rank1(m) <= k {
31 l = m;
32 } else {
33 r = m;
34 }
35 }
36 Some(l)
37 }More examples
crates/competitive/src/data_structure/wavelet_matrix.rs (line 646)
645 fn rank1(&self, level: usize, k: usize) -> usize {
646 self.bit_vectors[level].rank1(k)
647 }
648
649 fn reorder<U>(&self, level: usize, current: Vec<U>) -> Vec<U> {
650 assert_eq!(current.len(), self.len);
651 let mut next = Vec::with_capacity(self.len);
652 next.resize_with(self.len, MaybeUninit::uninit);
653 let mut zero = 0;
654 let mut one = self.zeros[level];
655 let mut current = current.into_iter();
656 for block in self.bit_vectors[level].blocks() {
657 let count = current.len().min(64);
658 if block.bits == 0 || block.bits == u64::MAX {
659 let offset = if block.bits == 0 { &mut zero } else { &mut one };
660 for (slot, value) in next[*offset..*offset + count]
661 .iter_mut()
662 .zip(current.by_ref().take(count))
663 {
664 slot.write(value);
665 }
666 *offset += count;
667 } else {
668 for (i, value) in current.by_ref().take(count).enumerate() {
669 let bit = (block.bits >> i) & 1 != 0;
670 next[if bit { one } else { zero }].write(value);
671 zero += !bit as usize;
672 one += bit as usize;
673 }
674 }
675 }
676 // SAFETY: the partition counts fill every slot once, and `MaybeUninit<U>` has `U`'s layout.
677 unsafe {
678 let mut next = mem::ManuallyDrop::new(next);
679 Vec::from_raw_parts(next.as_mut_ptr().cast(), next.len(), next.capacity())
680 }
681 }
682
683 fn range_by_index(&self, idx: usize, mut range: Range<usize>) -> Range<usize> {
684 if !self.quad_vectors.is_empty() {
685 for (level, vector) in self.quad_vectors.iter().enumerate() {
686 if range.is_empty() {
687 break;
688 }
689 let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
690 range = vector.starts[digit] + vector.rank(digit, range.start)
691 ..vector.starts[digit] + vector.rank(digit, range.end);
692 }
693 return range;
694 }
695 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
696 if range.is_empty() {
697 break;
698 }
699 let level = self.level(d);
700 let start1 = self.rank1(level, range.start);
701 let end1 = self.rank1(level, range.end);
702 if ((idx >> d) & 1) != 0 {
703 range.start = self.zeros[level] + start1;
704 range.end = self.zeros[level] + end1;
705 } else {
706 range.start -= start1;
707 range.end -= end1;
708 }
709 }
710 range
711 }
712
713 fn batch<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
714 #[cfg(target_arch = "x86_64")]
715 if !cfg!(target_feature = "popcnt") && is_x86_feature_detected!("popcnt") {
716 // SAFETY: POPCNT is checked above.
717 unsafe {
718 self.batch_popcnt::<OP>(states, count);
719 }
720 return;
721 }
722 self.batch_inner::<OP>(states, count);
723 }
724
725 #[cfg(target_arch = "x86_64")]
726 #[target_feature(enable = "popcnt")]
727 unsafe fn batch_popcnt<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
728 self.batch_inner::<OP>(states, count);
729 }
730
731 #[inline(always)]
732 fn batch_inner<const OP: u8>(&self, states: &mut [[usize; 4]; 16], count: usize) {
733 if OP == RANK_LESSTHAN && self.compress.size() <= 1 {
734 for state in &mut states[..count] {
735 state[3] = if state[2] == 0 {
736 0
737 } else {
738 state[1] - state[0]
739 };
740 }
741 return;
742 }
743 if !self.quad_vectors.is_empty() && matches!(OP, ACCESS | RANK | QUANTILE) {
744 #[cfg(target_arch = "x86_64")]
745 if OP == QUANTILE && count >= 8 && self.backend == super::SimdBackend::Avx512 {
746 // SAFETY: checked batch ranges stay within the quad blocks after each
747 // stable partition. Construction restricts quad counters to u32, and
748 // the cached backend includes AVX512F and VPOPCNTDQ.
749 unsafe {
750 simd::quad_avx512(&self.quad_vectors, &mut states[..count.next_multiple_of(8)]);
751 }
752 return;
753 }
754
755 for (level, vector) in self.quad_vectors.iter().enumerate() {
756 let d = (self.quad_vectors.len() - level - 1) * 2;
757 #[cfg(target_arch = "x86_64")]
758 let next = self
759 .quad_vectors
760 .get(level + 1)
761 .filter(|_| self.len >= 1048576 || (OP == ACCESS && self.len >= 262144));
762 for state in &mut states[..count] {
763 if OP == ACCESS {
764 let (digit, rank) = vector.access_rank(state[0]);
765 state[0] = vector.starts[digit] + rank;
766 state[3] = state[3] * 4 + digit;
767 } else if OP == RANK {
768 let digit = (state[2] >> d) & 3;
769 state[0] = vector.starts[digit] + vector.rank(digit, state[0]);
770 state[1] = vector.starts[digit] + vector.rank(digit, state[1]);
771 } else {
772 let start = vector.ranks(state[0]);
773 let end = vector.ranks(state[1]);
774 let prefix = [
775 0,
776 end[0] - start[0],
777 end[0] + end[1] - start[0] - start[1],
778 state[1] - state[0] - (end[3] - start[3]),
779 ];
780 let digit = (state[2] >= prefix[1]) as usize
781 + (state[2] >= prefix[2]) as usize
782 + (state[2] >= prefix[3]) as usize;
783 state[2] -= prefix[digit];
784 state[0] = vector.starts[digit] + start[digit];
785 state[1] = vector.starts[digit] + end[digit];
786 state[3] = state[3] * 4 + digit;
787 }
788 #[cfg(target_arch = "x86_64")]
789 if let Some(next) = next {
790 // SAFETY: stable partitions keep both endpoints within the next vector.
791 unsafe {
792 std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
793 next.blocks.as_ptr().add(state[0] / 64).cast(),
794 );
795 if OP != ACCESS {
796 std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
797 next.blocks.as_ptr().add(state[1] / 64).cast(),
798 );
799 }
800 }
801 }
802 }
803 }
804 return;
805 }
806 let first = (matches!(OP, ACCESS | RANK | QUANTILE)
807 && self.compress.size().is_power_of_two()) as usize;
808 #[cfg(target_arch = "x86_64")]
809 if OP == RANK_LESSTHAN && count >= 8 && self.len <= u32::MAX as usize {
810 // SAFETY: the caller checked each initial position. Rank transitions remain in
811 // 0..=len, including the sentinel block. Padding uses position zero. The
812 // backend was detected at construction, and x86-64 blocks contain two u64s.
813 unsafe {
814 match self.backend {
815 super::SimdBackend::Avx512 => {
816 simd::rank_lessthan_avx512(
817 &self.bit_vectors[first..],
818 &self.zeros[first..],
819 &mut states[..count.next_multiple_of(8)],
820 );
821 return;
822 }
823 super::SimdBackend::Avx2 if self.len < 1048576 => {
824 simd::rank_lessthan_avx2(
825 &self.bit_vectors[first..],
826 &self.zeros[first..],
827 &mut states[..count.next_multiple_of(4)],
828 );
829 return;
830 }
831 _ => {}
832 }
833 }
834 }
835 for d in (0..self.bit_length - first).rev() {
836 let level = self.level(d);
837 for state in &mut states[..count] {
838 let (bit, start1) = self.bit_vectors[level].access_rank1(state[0]);
839 let start0 = state[0] - start1;
840 let end1 = if OP == ACCESS {
841 0
842 } else {
843 self.rank1(level, state[1])
844 };
845 let end0 = if OP == ACCESS { 0 } else { state[1] - end1 };
846 let count0 = if OP == ACCESS { 0 } else { end0 - start0 };
847 let bit = match OP {
848 ACCESS => bit,
849 RANK | RANK_LESSTHAN => (state[2] >> d) & 1 != 0,
850 _ => state[2] >= count0,
851 };
852 state[0] = if bit {
853 self.zeros[level] + start1
854 } else {
855 start0
856 };
857 if OP != ACCESS {
858 state[1] = if bit { self.zeros[level] + end1 } else { end0 };
859 }
860 if OP == ACCESS || OP == QUANTILE {
861 state[3] |= (bit as usize) << d;
862 }
863 if OP == QUANTILE {
864 state[2] -= if bit { count0 } else { 0 };
865 }
866 if OP == RANK_LESSTHAN {
867 state[3] += if bit { count0 } else { 0 };
868 }
869 }
870 }
871 }
872
873 /// get k-th value
874 pub fn access(&self, mut k: usize) -> T {
875 if !self.quad_vectors.is_empty() {
876 let mut index = 0;
877 for vector in &self.quad_vectors {
878 let (digit, rank) = vector.access_rank(k);
879 index = index * 4 + digit;
880 k = vector.starts[digit] + rank;
881 }
882 return self.compress.values()[index].clone();
883 }
884 let mut idx = 0;
885 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
886 let level = self.level(d);
887 let (bit, rank1) = self.bit_vectors[level].access_rank1(k);
888 idx |= (bit as usize) << d;
889 k = if bit {
890 self.zeros[level] + rank1
891 } else {
892 k - rank1
893 };
894 }
895 self.compress.values()[idx].clone()
896 }
897
898 /// Returns the values at `indices` in input order.
899 pub fn access_batch(&self, indices: impl IntoIterator<Item = usize>) -> Vec<T> {
900 let indices: Vec<_> = indices.into_iter().collect();
901 let mut result = Vec::with_capacity(indices.len());
902 for indices in indices.chunks(16) {
903 let mut states = [[0; 4]; 16];
904 for (state, &index) in states.iter_mut().zip(indices) {
905 assert!(index < self.len);
906 state[0] = index;
907 }
908 self.batch::<ACCESS>(&mut states, indices.len());
909 result.extend(
910 states[..indices.len()]
911 .iter()
912 .map(|state| self.compress.values()[state[3]].clone()),
913 );
914 }
915 result
916 }
917
918 /// the number of val in range
919 pub fn rank(&self, val: T, range: Range<usize>) -> usize {
920 match self.compress.index_exact(&val) {
921 Some(idx) => self.range_by_index(idx, range).len(),
922 None => 0,
923 }
924 }
925
926 /// Returns the number of exact matches for each `(value, range)` query.
927 pub fn rank_batch(&self, queries: impl IntoIterator<Item = (T, Range<usize>)>) -> Vec<usize> {
928 let queries: Vec<_> = queries.into_iter().collect();
929 let mut result = Vec::with_capacity(queries.len());
930 for queries in queries.chunks(16) {
931 let mut states = [[0; 4]; 16];
932 for (state, (value, range)) in states.iter_mut().zip(queries) {
933 assert!(range.start <= range.end && range.end <= self.len);
934 if let Some(index) = self.compress.index_exact(value) {
935 *state = [range.start, range.end, index, 0];
936 }
937 }
938 self.batch::<RANK>(&mut states, queries.len());
939 result.extend(
940 states[..queries.len()]
941 .iter()
942 .map(|state| state[1] - state[0]),
943 );
944 }
945 result
946 }
947
948 /// index of k-th val
949 pub fn select(&self, val: T, k: usize) -> Option<usize> {
950 let idx = self.compress.index_exact(&val)?;
951 let range = self.range_by_index(idx, 0..self.len);
952 if range.len() <= k {
953 return None;
954 }
955 let mut i = range.start + k;
956 if !self.quad_vectors.is_empty() {
957 for (level, vector) in self.quad_vectors.iter().enumerate().rev() {
958 let digit = (idx >> ((self.quad_vectors.len() - level - 1) * 2)) & 3;
959 i = vector.select(digit, i - vector.starts[digit]);
960 }
961 return Some(i);
962 }
963 for level in (self.compress.size().is_power_of_two() as usize..self.bit_length).rev() {
964 if i >= self.zeros[level] {
965 i = self.bit_vectors[level]
966 .select1(i - self.zeros[level])
967 .unwrap();
968 } else {
969 i = self.bit_vectors[level].select0(i).unwrap();
970 }
971 }
972 Some(i)
973 }
974
975 /// get k-th smallest value in range
976 pub fn quantile(&self, mut range: Range<usize>, mut k: usize) -> T {
977 if !self.quad_vectors.is_empty() {
978 return self.quad_quantile(range, k, 0, 0);
979 }
980 let mut idx = 0;
981 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
982 let level = self.level(d);
983 let start1 = self.rank1(level, range.start);
984 let end1 = self.rank1(level, range.end);
985 let start0 = range.start - start1;
986 let end0 = range.end - end1;
987 let z = end0 - start0;
988 let bit = z <= k;
989 k -= if bit { z } else { 0 };
990 idx |= (bit as usize) << d;
991 range.start = if bit {
992 self.zeros[level] + start1
993 } else {
994 start0
995 };
996 range.end = if bit { self.zeros[level] + end1 } else { end0 };
997 }
998 self.compress.values()[idx].clone()
999 }
1000
1001 #[inline(always)]
1002 fn quad_quantile(
1003 &self,
1004 mut range: Range<usize>,
1005 mut k: usize,
1006 level: usize,
1007 mut index: usize,
1008 ) -> T {
1009 for vector in &self.quad_vectors[level..] {
1010 let start = vector.ranks(range.start);
1011 let end = vector.ranks(range.end);
1012 let mut digit = 0;
1013 while digit < 3 && k >= end[digit] - start[digit] {
1014 k -= end[digit] - start[digit];
1015 digit += 1;
1016 }
1017 index = index * 4 + digit;
1018 range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1019 }
1020 self.compress.values()[index].clone()
1021 }
1022
1023 pub fn quantile_batch(
1024 &self,
1025 queries: impl IntoIterator<Item = (Range<usize>, usize)>,
1026 ) -> Vec<T> {
1027 let queries: Vec<_> = queries.into_iter().collect();
1028 let mut result = Vec::with_capacity(queries.len());
1029 for queries in queries.chunks(16) {
1030 let mut states = [[0; 4]; 16];
1031 for (state, (range, k)) in states.iter_mut().zip(queries) {
1032 assert!(range.start <= range.end && range.end <= self.len && *k < range.len());
1033 *state = [range.start, range.end, *k, 0];
1034 }
1035 if let [(range, k)] = queries {
1036 result.push(self.quantile(range.clone(), *k));
1037 continue;
1038 }
1039 self.batch::<QUANTILE>(&mut states, queries.len());
1040 result.extend(
1041 states[..queries.len()]
1042 .iter()
1043 .map(|state| self.compress.values()[state[3]].clone()),
1044 );
1045 }
1046 result
1047 }
1048
1049 /// Returns the requested order statistics of one range. `ranks` must be sorted
1050 /// in nondecreasing order and every rank must be less than the range length.
1051 pub fn quantiles_sorted(&self, range: Range<usize>, ranks: &[usize]) -> Vec<T> {
1052 assert!(range.start <= range.end && range.end <= self.len);
1053 assert!(ranks.windows(2).all(|pair| pair[0] <= pair[1]));
1054 assert!(ranks.last().is_none_or(|&k| k < range.len()));
1055 if ranks.is_empty() {
1056 return Vec::new();
1057 }
1058 if ranks[0] == ranks[ranks.len() - 1] {
1059 return vec![self.quantile(range, ranks[0]); ranks.len()];
1060 }
1061 let use_batch = ranks.len() < 1024;
1062 #[cfg(target_arch = "x86_64")]
1063 let use_batch = use_batch || self.backend == super::SimdBackend::Avx512;
1064 // Sparse requests amortize less of the shared traversal's branching work.
1065 if use_batch && ranks.len() < self.compress.size().div_ceil(16) {
1066 return self.quantile_batch(ranks.iter().map(|&k| (range.clone(), k)));
1067 }
1068 let mut result = Vec::with_capacity(ranks.len());
1069 if !self.quad_vectors.is_empty() {
1070 let mut stack = vec![(range, 0, 0, 0, ranks)];
1071 while let Some((range, base, index, level, ranks)) = stack.pop() {
1072 if level == self.quad_vectors.len() {
1073 result.resize(
1074 result.len() + ranks.len(),
1075 self.compress.values()[index].clone(),
1076 );
1077 continue;
1078 }
1079 if ranks.len() == 1 {
1080 result.push(self.quad_quantile(range, ranks[0] - base, level, index));
1081 continue;
1082 }
1083 let vector = &self.quad_vectors[level];
1084 let start = vector.ranks(range.start);
1085 let end = vector.ranks(range.end);
1086 let mut split = base + range.len();
1087 let mut requested = ranks.len();
1088 for digit in (0..4).rev() {
1089 split -= end[digit] - start[digit];
1090 let boundary = ranks[..requested].partition_point(|&k| k < split);
1091 if boundary < requested {
1092 stack.push((
1093 vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit],
1094 split,
1095 index * 4 + digit,
1096 level + 1,
1097 &ranks[boundary..requested],
1098 ));
1099 }
1100 requested = boundary;
1101 }
1102 }
1103 return result;
1104 }
1105 let mut stack = vec![(range, 0, 0, self.bit_length, ranks)];
1106 while let Some((range, base, index, bits, ranks)) = stack.pop() {
1107 if bits == 0 {
1108 result.resize(
1109 result.len() + ranks.len(),
1110 self.compress.values()[index].clone(),
1111 );
1112 continue;
1113 }
1114 let d = bits - 1;
1115 let level = self.level(d);
1116 let start1 = self.rank1(level, range.start);
1117 let end1 = self.rank1(level, range.end);
1118 let start0 = range.start - start1;
1119 let end0 = range.end - end1;
1120 let split = base + end0 - start0;
1121 let boundary = ranks.partition_point(|&k| k < split);
1122 if boundary < ranks.len() {
1123 stack.push((
1124 self.zeros[level] + start1..self.zeros[level] + end1,
1125 split,
1126 index | (1 << d),
1127 d,
1128 &ranks[boundary..],
1129 ));
1130 }
1131 if boundary != 0 {
1132 stack.push((start0..end0, base, index, d, &ranks[..boundary]));
1133 }
1134 }
1135 result
1136 }
1137
1138 /// get k-th smallest value out of range
1139 pub fn quantile_outer(&self, mut range: Range<usize>, mut k: usize) -> T {
1140 if !self.quad_vectors.is_empty() {
1141 let mut outer = 0..self.len;
1142 let mut index = 0;
1143 for vector in &self.quad_vectors {
1144 let start = vector.ranks(range.start);
1145 let end = vector.ranks(range.end);
1146 let outer_start = vector.ranks(outer.start);
1147 let outer_end = vector.ranks(outer.end);
1148 let mut digit = 0;
1149 while digit < 3 {
1150 let count = outer_end[digit] - outer_start[digit] - (end[digit] - start[digit]);
1151 if k < count {
1152 break;
1153 }
1154 k -= count;
1155 digit += 1;
1156 }
1157 index = index * 4 + digit;
1158 range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1159 outer = vector.starts[digit] + outer_start[digit]
1160 ..vector.starts[digit] + outer_end[digit];
1161 }
1162 return self.compress.values()[index].clone();
1163 }
1164 let mut idx = 0;
1165 let mut orange = 0..self.len;
1166 for d in (0..self.bit_length - self.compress.size().is_power_of_two() as usize).rev() {
1167 let level = self.level(d);
1168 let range_start1 = self.rank1(level, range.start);
1169 let range_end1 = self.rank1(level, range.end);
1170 let outer_start1 = self.rank1(level, orange.start);
1171 let outer_end1 = self.rank1(level, orange.end);
1172 let range_start0 = range.start - range_start1;
1173 let range_end0 = range.end - range_end1;
1174 let outer_start0 = orange.start - outer_start1;
1175 let outer_end0 = orange.end - outer_end1;
1176 let z = (outer_end0 - outer_start0) - (range_end0 - range_start0);
1177 if z <= k {
1178 k -= z;
1179 idx |= 1 << d;
1180 range.start = self.zeros[level] + range_start1;
1181 range.end = self.zeros[level] + range_end1;
1182 orange.start = self.zeros[level] + outer_start1;
1183 orange.end = self.zeros[level] + outer_end1;
1184 } else {
1185 range.start = range_start0;
1186 range.end = range_end0;
1187 orange.start = outer_start0;
1188 orange.end = outer_end0;
1189 }
1190 }
1191 self.compress.values()[idx].clone()
1192 }
1193
1194 /// the number of value less than val in range
1195 pub fn rank_lessthan(&self, val: T, range: Range<usize>) -> usize {
1196 let index = self.compress.index_lower_bound(&val);
1197 if index == self.compress.size() {
1198 return range.len();
1199 }
1200 if self.len >= 262144 && !self.quad_vectors.is_empty() {
1201 self.quad_rank_lessthan_index(index, range, 0)
1202 } else {
1203 let bits = usize::BITS as usize
1204 - self.compress.size().saturating_sub(1).leading_zeros() as usize;
1205 self.rank_lessthan_index(index, range, bits)
1206 }
1207 }
1208
1209 /// Returns the number of values below each query's threshold.
1210 pub fn rank_lessthan_batch(
1211 &self,
1212 queries: impl IntoIterator<Item = (T, Range<usize>)>,
1213 ) -> Vec<usize> {
1214 let queries: Vec<_> = queries.into_iter().collect();
1215 let mut result = Vec::with_capacity(queries.len());
1216 for queries in queries.chunks(16) {
1217 let mut states = [[0; 4]; 16];
1218 for (state, (value, range)) in states.iter_mut().zip(queries) {
1219 assert!(range.start <= range.end && range.end <= self.len);
1220 *state = [
1221 range.start,
1222 range.end,
1223 self.compress.index_lower_bound(value),
1224 0,
1225 ];
1226 }
1227 self.batch::<RANK_LESSTHAN>(&mut states, queries.len());
1228 result.extend(states[..queries.len()].iter().map(|state| state[3]));
1229 }
1230 result
1231 }
1232
1233 fn quad_rank_lessthan_index(
1234 &self,
1235 index: usize,
1236 mut range: Range<usize>,
1237 start_level: usize,
1238 ) -> usize {
1239 let mut result = 0;
1240 for (level, vector) in self.quad_vectors.iter().enumerate().skip(start_level) {
1241 let d = (self.quad_vectors.len() - level - 1) * 2;
1242 if d + 2 <= index.trailing_zeros() as usize {
1243 break;
1244 }
1245 let digit = (index >> d) & 3;
1246 let ranks = |position: usize| {
1247 let block = &vector.blocks[position / 64];
1248 let mask = !(u64::MAX << (position % 64));
1249 let low = if digit & 1 == 0 { !block.lo } else { block.lo };
1250 let high = if digit & 2 == 0 { !block.hi } else { block.hi };
1251 let exact = block.rank[digit] as usize + (low & high & mask).count_ones() as usize;
1252 let less = match digit {
1253 0 => 0,
1254 1 => {
1255 block.rank[0] as usize
1256 + (!(block.lo | block.hi) & mask).count_ones() as usize
1257 }
1258 2 => {
1259 block.rank[0] as usize
1260 + block.rank[1] as usize
1261 + (!block.hi & mask).count_ones() as usize
1262 }
1263 _ => position - exact,
1264 };
1265 (exact, less)
1266 };
1267 let (start, start_less) = ranks(range.start);
1268 let (end, end_less) = ranks(range.end);
1269 result += end_less - start_less;
1270 range = vector.starts[digit] + start..vector.starts[digit] + end;
1271 }
1272 result
1273 }
1274
1275 fn rank_lessthan_index(&self, idx: usize, mut range: Range<usize>, bits: usize) -> usize {
1276 let mut res = 0;
1277 for d in (idx.trailing_zeros() as usize..bits).rev() {
1278 let level = self.level(d);
1279 let start1 = self.rank1(level, range.start);
1280 let end1 = self.rank1(level, range.end);
1281 let bit = (idx >> d) & 1 != 0;
1282 let start0 = range.start - start1;
1283 let end0 = range.end - end1;
1284 res += if bit { end0 - start0 } else { 0 };
1285 range.start = if bit {
1286 self.zeros[level] + start1
1287 } else {
1288 start0
1289 };
1290 range.end = if bit { self.zeros[level] + end1 } else { end0 };
1291 }
1292 res
1293 }
1294
1295 /// the number of valrange in range
1296 pub fn rank_range(&self, valrange: Range<T>, mut range: Range<usize>) -> usize {
1297 let lower = self.compress.index_lower_bound(&valrange.start);
1298 let upper = self.compress.index_lower_bound(&valrange.end);
1299 if lower >= upper {
1300 return 0;
1301 }
1302 if !self.quad_vectors.is_empty() {
1303 if upper == self.compress.size() {
1304 return range.len() - self.quad_rank_lessthan_index(lower, range, 0);
1305 }
1306 for (level, vector) in self.quad_vectors.iter().enumerate() {
1307 let d = (self.quad_vectors.len() - level - 1) * 2;
1308 let low = (lower >> d) & 3;
1309 let high = (upper >> d) & 3;
1310 if low == high {
1311 range = vector.starts[low] + vector.rank(low, range.start)
1312 ..vector.starts[low] + vector.rank(low, range.end);
1313 continue;
1314 }
1315 let start = vector.ranks(range.start);
1316 let end = vector.ranks(range.end);
1317 let count: usize = (low..high).map(|digit| end[digit] - start[digit]).sum();
1318 return count
1319 - self.quad_rank_lessthan_index(
1320 lower,
1321 vector.starts[low] + start[low]..vector.starts[low] + end[low],
1322 level + 1,
1323 )
1324 + self.quad_rank_lessthan_index(
1325 upper,
1326 vector.starts[high] + start[high]..vector.starts[high] + end[high],
1327 level + 1,
1328 );
1329 }
1330 return 0;
1331 }
1332 for d in (0..self.bit_length).rev() {
1333 let level = self.level(d);
1334 let start1 = self.rank1(level, range.start);
1335 let end1 = self.rank1(level, range.end);
1336 let zero_range = range.start - start1..range.end - end1;
1337 let one_range = self.zeros[level] + start1..self.zeros[level] + end1;
1338 if ((lower ^ upper) >> d) & 1 != 0 {
1339 return zero_range.len() - self.rank_lessthan_index(lower, zero_range, d)
1340 + self.rank_lessthan_index(upper, one_range, d);
1341 }
1342 range = if (lower >> d) & 1 != 0 {
1343 one_range
1344 } else {
1345 zero_range
1346 };
1347 }
1348 0
1349 }
1350
1351 /// Returns each query's count in a half-open value range.
1352 pub fn rank_range_batch(
1353 &self,
1354 queries: impl IntoIterator<Item = (Range<T>, Range<usize>)>,
1355 ) -> Vec<usize> {
1356 let queries: Vec<_> = queries.into_iter().collect();
1357 let mut result = Vec::with_capacity(queries.len());
1358 for queries in queries.chunks(8) {
1359 let mut states = [[0; 4]; 16];
1360 for (i, (values, range)) in queries.iter().enumerate() {
1361 assert!(range.start <= range.end && range.end <= self.len);
1362 let lower = self.compress.index_lower_bound(&values.start);
1363 let upper = self.compress.index_lower_bound(&values.end);
1364 if lower < upper {
1365 states[i * 2] = [range.start, range.end, lower, 0];
1366 states[i * 2 + 1] = [range.start, range.end, upper, 0];
1367 }
1368 }
1369 self.batch::<RANK_LESSTHAN>(&mut states, queries.len() * 2);
1370 result.extend(
1371 states[..queries.len() * 2]
1372 .as_chunks::<2>()
1373 .0
1374 .iter()
1375 .map(|pair| pair[1][3] - pair[0][3]),
1376 );
1377 }
1378 result
1379 }
1380
1381 pub fn query_less_than<F>(&self, val: T, mut range: Range<usize>, mut f: F)
1382 where
1383 F: FnMut(usize, Range<usize>),
1384 {
1385 let idx = self.compress.index_lower_bound(&val);
1386 if !self.quad_vectors.is_empty() {
1387 if idx == self.compress.size() && idx.is_power_of_two() {
1388 f(self.bit_length - 1, range);
1389 return;
1390 }
1391 for (level, vector) in self.quad_vectors.iter().enumerate().take(
1392 self.quad_vectors
1393 .len()
1394 .saturating_sub(idx.trailing_zeros() as usize / 2),
1395 ) {
1396 let d = (self.quad_vectors.len() - level - 1) * 2;
1397 let digit = (idx >> d) & 3;
1398 let start = vector.ranks(range.start);
1399 let end = vector.ranks(range.end);
1400 if digit & 2 != 0 {
1401 f(d + 1, start[0] + start[1]..end[0] + end[1]);
1402 }
1403 if digit & 1 != 0 {
1404 let zero = digit & 2;
1405 f(
1406 d,
1407 vector.starts[zero] + start[zero]..vector.starts[zero] + end[zero],
1408 );
1409 }
1410 range = vector.starts[digit] + start[digit]..vector.starts[digit] + end[digit];
1411 }
1412 return;
1413 }
1414 for d in (idx.trailing_zeros() as usize..self.bit_length).rev() {
1415 let level = self.level(d);
1416 let start1 = self.rank1(level, range.start);
1417 let end1 = self.rank1(level, range.end);
1418 let start0 = range.start - start1;
1419 let end0 = range.end - end1;
1420 if ((idx >> d) & 1) != 0 {
1421 f(d, start0..end0);
1422 range.start = self.zeros[level] + start1;
1423 range.end = self.zeros[level] + end1;
1424 } else {
1425 range.start = start0;
1426 range.end = end0;
1427 }
1428 }
1429 }
1430
1431 pub fn build_fold<M>(&self, weights: &[M::T]) -> WaveletMatrixFold<'_, T, M>
1432 where
1433 M: AbelianGroup,
1434 {
1435 assert_eq!(weights.len(), self.len);
1436 let mut offsets = Vec::with_capacity(self.bit_length);
1437 let mut prefix = Vec::with_capacity(self.zeros.iter().map(|&zero| zero + 1).sum());
1438 let mut current: Vec<M::T> = weights.to_vec();
1439 for level in 0..self.bit_length {
1440 current = self.reorder(level, current);
1441 offsets.push(prefix.len());
1442 let mut acc = M::unit();
1443 prefix.push(acc.clone());
1444 for w in ¤t[..self.zeros[level]] {
1445 acc = M::operate(&acc, w);
1446 prefix.push(acc.clone());
1447 }
1448 }
1449 WaveletMatrixFold {
1450 wavelet_matrix: self,
1451 prefix,
1452 offsets,
1453 }
1454 }
1455
1456 pub fn build_point_add<M>(&self, weights: &[M::T]) -> WaveletMatrixPointAdd<'_, T, M>
1457 where
1458 M: AbelianGroup,
1459 {
1460 assert_eq!(weights.len(), self.len);
1461 let mut current = weights.to_vec();
1462 let mut bits = Vec::with_capacity(self.bit_length);
1463 for level in 0..self.bit_length {
1464 current = self.reorder(level, current);
1465 bits.push(BinaryIndexedTree::from_slice(¤t[..self.zeros[level]]));
1466 }
1467 WaveletMatrixPointAdd {
1468 wavelet_matrix: self,
1469 bits,
1470 }
1471 }
1472}
1473
1474pub struct WaveletMatrixPointAdd<'a, T, M>
1475where
1476 T: Ord + Clone,
1477 M: AbelianGroup,
1478{
1479 wavelet_matrix: &'a WaveletMatrix<T>,
1480 bits: Vec<BinaryIndexedTree<M>>,
1481}
1482
1483impl<'a, T, M> WaveletMatrixPointAdd<'a, T, M>
1484where
1485 T: Ord + Clone,
1486 M: AbelianGroup,
1487{
1488 pub fn update(&mut self, mut index: usize, value: M::T) {
1489 debug_assert!(index < self.wavelet_matrix.len);
1490 for d in (0..self.wavelet_matrix.bit_length).rev() {
1491 let level = self.wavelet_matrix.level(d);
1492 let (bit, rank1) = self.wavelet_matrix.bit_vectors[level].access_rank1(index);
1493 if bit {
1494 index = self.wavelet_matrix.zeros[level] + rank1;
1495 } else {
1496 index -= rank1;
1497 self.bits[level].update(index, value.clone());
1498 }
1499 }
1500 }
1501
1502 pub fn fold_lessthan(&self, value: T, range: Range<usize>) -> M::T {
1503 let mut result = M::unit();
1504 self.wavelet_matrix
1505 .query_less_than(value, range, |d, range| {
1506 M::operate_assign(
1507 &mut result,
1508 &self.bits[self.wavelet_matrix.level(d)].fold_abelian(range.start, range.end),
1509 );
1510 });
1511 result
1512 }
1513
1514 pub fn fold_range(&self, values: Range<T>, range: Range<usize>) -> M::T {
1515 let lower = self
1516 .wavelet_matrix
1517 .compress
1518 .index_lower_bound(&values.start);
1519 let upper = self.wavelet_matrix.compress.index_lower_bound(&values.end);
1520 if lower >= upper {
1521 return M::unit();
1522 }
1523 let mut range = range;
1524 for d in (0..self.wavelet_matrix.bit_length).rev() {
1525 let level = self.wavelet_matrix.level(d);
1526 let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1527 let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1528 let start0 = range.start - start1;
1529 let end0 = range.end - end1;
1530 if ((lower >> d) & 1) == ((upper >> d) & 1) {
1531 if ((lower >> d) & 1) == 0 {
1532 range = start0..end0;
1533 } else {
1534 range = self.wavelet_matrix.zeros[level] + start1
1535 ..self.wavelet_matrix.zeros[level] + end1;
1536 }
1537 continue;
1538 }
1539 let zero_range = start0..end0;
1540 let one_range =
1541 self.wavelet_matrix.zeros[level] + start1..self.wavelet_matrix.zeros[level] + end1;
1542 let lower_sum = self.fold_lessthan_index(lower, zero_range.clone(), d);
1543 let upper_sum = self.fold_lessthan_index(upper, one_range, d);
1544 let zero_sum = self.bits[level].fold_abelian(zero_range.start, zero_range.end);
1545 let mut result = M::rinv_operate(&zero_sum, &lower_sum);
1546 M::operate_assign(&mut result, &upper_sum);
1547 return result;
1548 }
1549 M::unit()
1550 }
1551
1552 fn fold_lessthan_index(&self, idx: usize, mut range: Range<usize>, bits: usize) -> M::T {
1553 let mut result = M::unit();
1554 for d in (idx.trailing_zeros() as usize..bits).rev() {
1555 let level = self.wavelet_matrix.level(d);
1556 let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1557 let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1558 let start0 = range.start - start1;
1559 let end0 = range.end - end1;
1560 if ((idx >> d) & 1) != 0 {
1561 M::operate_assign(&mut result, &self.bits[level].fold_abelian(start0, end0));
1562 range.start = self.wavelet_matrix.zeros[level] + start1;
1563 range.end = self.wavelet_matrix.zeros[level] + end1;
1564 } else {
1565 range.start = start0;
1566 range.end = end0;
1567 }
1568 }
1569 result
1570 }
1571}
1572
1573#[derive(Debug, Clone)]
1574pub struct WaveletMatrixFold<'a, T, M>
1575where
1576 T: Ord + Clone,
1577 M: AbelianGroup,
1578{
1579 wavelet_matrix: &'a WaveletMatrix<T>,
1580 prefix: Vec<M::T>,
1581 offsets: Vec<usize>,
1582}
1583
1584impl<'a, T, M> WaveletMatrixFold<'a, T, M>
1585where
1586 T: Ord + Clone,
1587 M: AbelianGroup,
1588{
1589 pub fn fold_lessthan(&self, val: T, range: Range<usize>) -> M::T {
1590 self.fold_lessthan_with_count(val, range).1
1591 }
1592
1593 pub fn fold_lessthan_with_count(&self, val: T, range: Range<usize>) -> (usize, M::T) {
1594 debug_assert!(range.end <= self.wavelet_matrix.len);
1595 let [result] = self.fold_lessthan_indices_with_count(
1596 [self.wavelet_matrix.compress.index_lower_bound(&val)],
1597 [range],
1598 self.wavelet_matrix.bit_length,
1599 );
1600 result
1601 }
1602
1603 pub fn fold_range(&self, valrange: Range<T>, range: Range<usize>) -> M::T {
1604 self.fold_range_with_count(valrange, range).1
1605 }
1606
1607 pub fn fold_range_with_count(
1608 &self,
1609 valrange: Range<T>,
1610 mut range: Range<usize>,
1611 ) -> (usize, M::T) {
1612 debug_assert!(range.end <= self.wavelet_matrix.len);
1613 let lower = self
1614 .wavelet_matrix
1615 .compress
1616 .index_lower_bound(&valrange.start);
1617 let upper = self
1618 .wavelet_matrix
1619 .compress
1620 .index_lower_bound(&valrange.end);
1621 if lower >= upper {
1622 return (0, M::unit());
1623 }
1624 for d in (0..self.wavelet_matrix.bit_length).rev() {
1625 let level = self.wavelet_matrix.level(d);
1626 let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1627 let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1628 let start0 = range.start - start1;
1629 let end0 = range.end - end1;
1630 if ((lower >> d) & 1) == ((upper >> d) & 1) {
1631 if ((lower >> d) & 1) == 0 {
1632 range = start0..end0;
1633 } else {
1634 range = self.wavelet_matrix.zeros[level] + start1
1635 ..self.wavelet_matrix.zeros[level] + end1;
1636 }
1637 continue;
1638 }
1639 let zero_range = start0..end0;
1640 let one_range =
1641 self.wavelet_matrix.zeros[level] + start1..self.wavelet_matrix.zeros[level] + end1;
1642 let [(lower_count, lower_sum), (upper_count, upper_sum)] = self
1643 .fold_lessthan_indices_with_count(
1644 [lower, upper],
1645 [zero_range.clone(), one_range],
1646 d,
1647 );
1648 let zero_sum = self.range_sum(level, zero_range.clone());
1649 return (
1650 zero_range.len() - lower_count + upper_count,
1651 M::operate(&M::rinv_operate(&zero_sum, &lower_sum), &upper_sum),
1652 );
1653 }
1654 (0, M::unit())
1655 }
1656
1657 #[inline]
1658 fn range_sum(&self, level: usize, range: Range<usize>) -> M::T {
1659 let offset = self.offsets[level];
1660 M::rinv_operate(
1661 &self.prefix[offset + range.end],
1662 &self.prefix[offset + range.start],
1663 )
1664 }
1665
1666 fn fold_lessthan_indices_with_count<const N: usize>(
1667 &self,
1668 indices: [usize; N],
1669 mut ranges: [Range<usize>; N],
1670 bits: usize,
1671 ) -> [(usize, M::T); N] {
1672 let mut results = std::array::from_fn(|_| (0, M::unit()));
1673 let last = indices
1674 .iter()
1675 .map(|index| index.trailing_zeros() as usize)
1676 .min()
1677 .unwrap_or(bits);
1678 for d in (last..bits).rev() {
1679 let level = self.wavelet_matrix.level(d);
1680 for ((&index, range), (count, sum)) in indices.iter().zip(&mut ranges).zip(&mut results)
1681 {
1682 let start1 = self.wavelet_matrix.bit_vectors[level].rank1(range.start);
1683 let end1 = self.wavelet_matrix.bit_vectors[level].rank1(range.end);
1684 let start0 = range.start - start1;
1685 let end0 = range.end - end1;
1686 if ((index >> d) & 1) != 0 {
1687 *count += end0 - start0;
1688 M::operate_assign(sum, &self.range_sum(level, start0..end0));
1689 range.start = self.wavelet_matrix.zeros[level] + start1;
1690 range.end = self.wavelet_matrix.zeros[level] + end1;
1691 } else {
1692 range.start = start0;
1693 range.end = end0;
1694 }
1695 }
1696 }
1697 results
1698 }Sourcefn rank0(&self, k: usize) -> usize
fn rank0(&self, k: usize) -> usize
the number of 0 in [0, k)
Examples found in repository?
crates/competitive/src/data_structure/bit_vector.rs (line 41)
39 fn select0(&self, k: usize) -> Option<usize> {
40 let n = self.bit_length();
41 if self.rank0(n) <= k {
42 return None;
43 }
44 let (mut l, mut r) = (0, n);
45 while r - l > 1 {
46 let m = l.midpoint(r);
47 if self.rank0(m) <= k {
48 l = m;
49 } else {
50 r = m;
51 }
52 }
53 Some(l)
54 }More examples
crates/competitive/src/data_structure/wavelet_matrix.rs (line 565)
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 }Sourcefn select1(&self, k: usize) -> Option<usize>
fn select1(&self, k: usize) -> Option<usize>
index of k-th 1
Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 966)
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 }Sourcefn select0(&self, k: usize) -> Option<usize>
fn select0(&self, k: usize) -> Option<usize>
index of k-th 0
Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 969)
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 }Dyn Compatibility§
This trait is dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".