Skip to main content

quad_avx512

Function quad_avx512 

Source
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    }