pub unsafe fn quad_avx512(
layers: &[WaveletMatrixQuadVector],
states: &mut [[usize; 4]],
)Examples found in repository?
crates/competitive/src/data_structure/wavelet_matrix.rs (line 750)
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 }