use rayon::prelude::*;
use statrs::distribution::{Beta, ContinuousCDF};
use crate::rotation::Rotation;
#[doc(hidden)]
pub const VALIDATE_CHUNK: usize = 64 * 1024;
pub(crate) fn par_first_invalid_coord(
values: &[f32],
dim: usize,
max_magnitude: f32,
) -> Option<(usize, usize, f32)> {
let first = values
.par_chunks(VALIDATE_CHUNK)
.enumerate()
.filter_map(|(ci, chunk)| {
first_invalid_in_chunk(chunk, max_magnitude).map(|j| ci * VALIDATE_CHUNK + j)
})
.min()?;
let x = values[first];
let vector_index = if dim == 0 { 0 } else { first / dim };
let coord_index = if dim == 0 { first } else { first % dim };
Some((vector_index, coord_index, x))
}
#[cfg(target_arch = "aarch64")]
#[inline]
fn first_invalid_in_chunk(chunk: &[f32], max_magnitude: f32) -> Option<usize> {
use std::arch::aarch64::*;
let n = chunk.len();
let quads = n / 4;
unsafe {
let bound = vdupq_n_f32(max_magnitude);
for q in 0..quads {
let x = vld1q_f32(chunk.as_ptr().add(q * 4));
let ok = vcaltq_f32(x, bound);
if vminvq_u32(ok) == 0 {
for j in q * 4..n {
let v = chunk[j];
if !(v.abs() < max_magnitude) {
return Some(j);
}
}
unreachable!("vector scan flagged a quad with no invalid element");
}
}
for j in quads * 4..n {
let v = chunk[j];
if !(v.abs() < max_magnitude) {
return Some(j);
}
}
}
None
}
#[cfg(target_arch = "x86_64")]
#[inline]
fn first_invalid_in_chunk(chunk: &[f32], max_magnitude: f32) -> Option<usize> {
if std::arch::is_x86_feature_detected!("avx2") {
unsafe { first_invalid_in_chunk_avx2(chunk, max_magnitude) }
} else {
first_invalid_in_chunk_scalar(chunk, max_magnitude)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn first_invalid_in_chunk_avx2(chunk: &[f32], max_magnitude: f32) -> Option<usize> {
use std::arch::x86_64::*;
let n = chunk.len();
let groups = n / 8;
let bound = _mm256_set1_ps(max_magnitude);
let abs_mask = _mm256_castsi256_ps(_mm256_set1_epi32(0x7fff_ffff));
for g in 0..groups {
let x = _mm256_loadu_ps(chunk.as_ptr().add(g * 8));
let ok = _mm256_cmp_ps::<_CMP_LT_OQ>(_mm256_and_ps(x, abs_mask), bound);
if _mm256_movemask_ps(ok) != 0xff {
for j in g * 8..n {
let v = chunk[j];
if !(v.abs() < max_magnitude) {
return Some(j);
}
}
unreachable!("vector scan flagged a group with no invalid element");
}
}
for j in groups * 8..n {
let v = chunk[j];
if !(v.abs() < max_magnitude) {
return Some(j);
}
}
None
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
#[inline]
fn first_invalid_in_chunk(chunk: &[f32], max_magnitude: f32) -> Option<usize> {
first_invalid_in_chunk_scalar(chunk, max_magnitude)
}
#[cfg_attr(target_arch = "aarch64", allow(dead_code))]
#[inline]
fn first_invalid_in_chunk_scalar(chunk: &[f32], max_magnitude: f32) -> Option<usize> {
chunk
.iter()
.position(|x| !x.is_finite() || x.abs() >= max_magnitude)
}
#[inline(always)]
fn f32_sort_key(x: f32) -> u32 {
let b = x.to_bits();
b ^ ((((b as i32) >> 31) as u32) | 0x8000_0000)
}
#[inline(always)]
fn f32_from_sort_key(k: u32) -> f32 {
let mask = if k & 0x8000_0000 != 0 { 0x8000_0000 } else { 0xFFFF_FFFF };
f32::from_bits(k ^ mask)
}
#[cfg(target_arch = "x86_64")]
const KERNEL_USES_RECON_TABLE: bool = false;
#[cfg(not(target_arch = "x86_64"))]
const KERNEL_USES_RECON_TABLE: bool = true;
#[doc(hidden)]
pub const RECON_TABLE_MIN_ROWS: usize = 16;
#[inline(always)]
fn recon_entry(centroid: f32, inv: f64, sh: f64) -> f64 {
(centroid as f64) * inv - sh
}
fn build_recon_table(
bit_width: usize,
dim: usize,
centroids: &[f32],
inv_scale_tq: &[f32],
shift: &[f32],
) -> Vec<f64> {
let n_codes = 1usize << bit_width;
let mut table = vec![0.0f64; n_codes * dim];
for d in 0..dim {
let inv = inv_scale_tq[d] as f64;
let sh = shift[d] as f64;
let row = &mut table[d * n_codes..(d + 1) * n_codes];
for (c, slot) in row.iter_mut().enumerate() {
*slot = recon_entry(centroids[c], inv, sh);
}
}
table
}
fn tqplus_anchor(beta: &Beta, centroids: &[f32]) -> (f64, f64, f32, f32) {
let c_outer = centroids.iter().fold(0.0f32, |acc, &c| acc.max(c.abs()));
let p_hi = beta.cdf((f64::from(c_outer) + 1.0) / 2.0);
(1.0 - p_hi, p_hi, -c_outer, c_outer)
}
pub const RECOMMENDED_CALIBRATION_ROWS: usize = 1000;
pub const MIN_CALIBRATION_ROWS: usize = 2;
fn rotate_batch_into(
vectors: &[f32],
n: usize,
dim: usize,
rotation: &Rotation,
rotated_scratch: &mut Vec<f32>,
) -> Vec<f32> {
let mut norms = vec![0.0f32; n];
rotated_scratch.clear();
rotated_scratch.reserve(n * dim);
rotated_scratch.spare_capacity_mut()[..n * dim]
.par_chunks_mut(dim)
.zip(norms.par_iter_mut())
.enumerate()
.for_each_init(
|| (vec![0.0f32; dim], vec![0.0f32; dim]),
|(scratch, row), (i, (dst_row, norm))| {
let src = &vectors[i * dim..(i + 1) * dim];
let n_val = simd_norm(src);
*norm = n_val;
let inv = if n_val > crate::MIN_INPUT_NORM { 1.0 / n_val } else { 0.0 };
rotation.apply_scaled_into(src, inv, row, scratch);
for (d, &s) in dst_row.iter_mut().zip(row.iter()) {
d.write(s);
}
},
);
unsafe {
rotated_scratch.set_len(n * dim);
}
norms
}
pub(crate) fn fit_calibration(
vectors: &[f32],
n: usize,
dim: usize,
rotation: &Rotation,
centroids: &[f32],
rotated_scratch: &mut Vec<f32>,
) -> (Vec<f32>, Vec<f32>) {
let _norms = rotate_batch_into(vectors, n, dim, rotation, rotated_scratch);
compute_tqplus_calibration(rotated_scratch, n, dim, centroids)
}
pub(crate) fn encode(
vectors: &[f32],
n: usize,
dim: usize,
rotation: &Rotation,
boundaries: &[f32],
centroids: &[f32],
bit_width: usize,
calibration: Option<(&[f32], &[f32])>,
rotated_scratch: &mut Vec<f32>,
packed_out: &mut Vec<u8>,
scales_out: &mut Vec<f32>,
) {
assert!(
dim != 0 && dim % 8 == 0,
"encode requires dim to be a nonzero multiple of 8, got {dim}",
);
let norms = rotate_batch_into(vectors, n, dim, rotation, rotated_scratch);
let rotated = std::mem::take(rotated_scratch);
encode_prerotated(
&rotated, &norms, n, dim, boundaries, centroids, bit_width, calibration, packed_out,
scales_out,
);
*rotated_scratch = rotated;
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn encode_prerotated(
rotated: &[f32],
norms: &[f32],
n: usize,
dim: usize,
boundaries: &[f32],
centroids: &[f32],
bit_width: usize,
calibration: Option<(&[f32], &[f32])>,
packed_out: &mut Vec<u8>,
scales_out: &mut Vec<f32>,
) {
let identity;
let (shift, scale_tq): (&[f32], &[f32]) = match calibration {
Some((s, sc)) => {
assert_eq!(s.len(), dim, "shift length must equal dim");
assert_eq!(sc.len(), dim, "scale_tq length must equal dim");
(s, sc)
}
None => {
identity = (vec![0.0f32; dim], vec![1.0f32; dim]);
(&identity.0, &identity.1)
}
};
let inv_scale_tq: Vec<f32> = scale_tq.iter().map(|s| 1.0 / s).collect();
let centroid_orig: Option<Vec<f64>> = (KERNEL_USES_RECON_TABLE
&& n >= RECON_TABLE_MIN_ROWS)
.then(|| build_recon_table(bit_width, dim, centroids, &inv_scale_tq, shift));
let bytes_per_plane = dim / 8;
let bytes_per_row = bit_width * bytes_per_plane;
let packed_old = packed_out.len();
let scales_old = scales_out.len();
crate::reserve_mostly_exact(packed_out, n * bytes_per_row);
crate::reserve_mostly_exact(scales_out, n);
scales_out.resize(scales_old + n, 0.0f32);
let packed = &mut packed_out.spare_capacity_mut()[..n * bytes_per_row];
let scales = &mut scales_out[scales_old..];
match bit_width {
2 => quantize_batch::<2>(
packed, scales, rotated, shift, scale_tq, &inv_scale_tq,
centroid_orig.as_deref(), boundaries, centroids, norms, dim,
bytes_per_row, bytes_per_plane,
),
3 => quantize_batch::<3>(
packed, scales, rotated, shift, scale_tq, &inv_scale_tq,
centroid_orig.as_deref(), boundaries, centroids, norms, dim,
bytes_per_row, bytes_per_plane,
),
4 => quantize_batch::<4>(
packed, scales, rotated, shift, scale_tq, &inv_scale_tq,
centroid_orig.as_deref(), boundaries, centroids, norms, dim,
bytes_per_row, bytes_per_plane,
),
other => unreachable!("unsupported bit_width {other}"),
}
unsafe {
packed_out.set_len(packed_old + n * bytes_per_row);
}
#[cfg(test)]
if FORCE_PANIC_AFTER_APPEND.with(|f| f.replace(false)) {
panic!("forced post-append encode panic (test)");
}
}
#[cfg(test)]
thread_local! {
static FORCE_PANIC_AFTER_APPEND: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
}
#[cfg(test)]
pub(crate) fn force_panic_after_append(on: bool) {
FORCE_PANIC_AFTER_APPEND.with(|f| f.set(on));
}
#[allow(clippy::too_many_arguments)]
fn quantize_batch<const BITS: usize>(
packed: &mut [std::mem::MaybeUninit<u8>],
scales: &mut [f32],
rotated: &[f32],
shift: &[f32],
scale_tq: &[f32],
inv_scale_tq: &[f32],
centroid_orig: Option<&[f64]>,
boundaries: &[f32],
centroids: &[f32],
norms: &[f32],
dim: usize,
bytes_per_row: usize,
bytes_per_plane: usize,
) {
packed.par_chunks_mut(bytes_per_row)
.zip(scales.par_iter_mut())
.enumerate()
.for_each_init(
|| vec![0u8; bytes_per_row],
|row_buf, (i, (packed_row, scale))| {
let rot_orig = &rotated[i * dim..(i + 1) * dim];
*scale = fused_quantize_scale_pack::<BITS>(
rot_orig, shift, scale_tq, inv_scale_tq,
centroid_orig, boundaries, centroids, norms[i],
row_buf, dim, bytes_per_plane,
);
for (d, &s) in packed_row.iter_mut().zip(row_buf.iter()) {
d.write(s);
}
},
);
}
fn compute_tqplus_calibration(
rotated: &[f32],
n: usize,
dim: usize,
centroids: &[f32],
) -> (Vec<f32>, Vec<f32>) {
let mut shift = vec![0.0f32; dim];
let mut scale = vec![1.0f32; dim];
debug_assert!(
n >= MIN_CALIBRATION_ROWS,
"fit needs two distinct order statistics per coordinate"
);
let a = (dim as f64 - 1.0) / 2.0;
let beta = Beta::new(a, a).expect("Beta(a, a) is valid for a > 0");
let (p_lo, p_hi, qc_lo, qc_hi) = tqplus_anchor(&beta, centroids);
let qc_span = qc_hi - qc_lo;
let lo_idx = ((n as f64) * p_lo) as usize;
let hi_idx = (((n as f64) * p_hi) as usize).min(n - 1).max(lo_idx + 1);
let workers = rayon::current_num_threads().max(1);
let mut tile_size = 128usize;
while tile_size > 32 && dim / tile_size < 2 * workers {
tile_size /= 2;
}
shift
.par_chunks_mut(tile_size)
.zip(scale.par_chunks_mut(tile_size))
.enumerate()
.for_each(|(tile_idx, (sh_tile, sc_tile))| {
let d0 = tile_idx * tile_size;
let tile = sh_tile.len();
let mut cols: Vec<u32> = Vec::with_capacity(tile * n);
{
let spare = &mut cols.spare_capacity_mut()[..tile * n];
for i in 0..n {
let row = &rotated[i * dim + d0..i * dim + d0 + tile];
for (c, &v) in row.iter().enumerate() {
spare[c * n + i].write(f32_sort_key(v));
}
}
}
unsafe {
cols.set_len(tile * n);
}
for (c, (sh, sc)) in sh_tile.iter_mut().zip(sc_tile.iter_mut()).enumerate() {
let coord = &mut cols[c * n..(c + 1) * n];
let (_, lo_val, right) = coord.select_nth_unstable(lo_idx);
let qe_lo = f32_from_sort_key(*lo_val);
let (_, hi_val, _) = right.select_nth_unstable(hi_idx - lo_idx - 1);
let qe_hi = f32_from_sort_key(*hi_val);
let qe_span = qe_hi - qe_lo;
if qe_span > 1e-6 {
*sc = qc_span / qe_span;
*sh = qc_lo / *sc - qe_lo;
}
}
});
(shift, scale)
}
const NORM_CHAINS: usize = 8;
#[inline(always)]
fn simd_norm(row: &[f32]) -> f32 {
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx") {
return unsafe { norm_sq_avx(row) }.sqrt();
}
}
#[cfg(target_arch = "aarch64")]
{
return unsafe { norm_sq_neon(row) }.sqrt();
}
#[allow(unreachable_code)]
{
norm_sq_scalar(row).sqrt()
}
}
#[inline]
fn norm_sq_scalar(row: &[f32]) -> f32 {
let mut chains = [0.0f32; NORM_CHAINS];
for (j, &x) in row.iter().enumerate() {
chains[j % NORM_CHAINS] += x * x;
}
combine_norm_chains(&chains)
}
#[inline(always)]
fn combine_norm_chains(c: &[f32; NORM_CHAINS]) -> f32 {
((c[0] + c[1]) + (c[2] + c[3])) + ((c[4] + c[5]) + (c[6] + c[7]))
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx")]
unsafe fn norm_sq_avx(row: &[f32]) -> f32 {
use std::arch::x86_64::*;
let n = row.len();
let mut acc = _mm256_setzero_ps();
let mut i = 0;
while i + NORM_CHAINS <= n {
let v = _mm256_loadu_ps(row.as_ptr().add(i));
acc = _mm256_add_ps(acc, _mm256_mul_ps(v, v));
i += NORM_CHAINS;
}
let mut chains = [0.0f32; NORM_CHAINS];
_mm256_storeu_ps(chains.as_mut_ptr(), acc);
while i < n {
let x = *row.get_unchecked(i);
chains[i % NORM_CHAINS] += x * x;
i += 1;
}
combine_norm_chains(&chains)
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn norm_sq_neon(row: &[f32]) -> f32 {
use std::arch::aarch64::*;
let n = row.len();
let mut acc_lo = vdupq_n_f32(0.0);
let mut acc_hi = vdupq_n_f32(0.0);
let mut i = 0;
while i + NORM_CHAINS <= n {
let lo = vld1q_f32(row.as_ptr().add(i));
let hi = vld1q_f32(row.as_ptr().add(i + 4));
acc_lo = vaddq_f32(acc_lo, vmulq_f32(lo, lo));
acc_hi = vaddq_f32(acc_hi, vmulq_f32(hi, hi));
i += NORM_CHAINS;
}
let mut chains = [0.0f32; NORM_CHAINS];
vst1q_f32(chains.as_mut_ptr(), acc_lo);
vst1q_f32(chains.as_mut_ptr().add(4), acc_hi);
while i < n {
let x = *row.get_unchecked(i);
chains[i % NORM_CHAINS] += x * x;
i += 1;
}
combine_norm_chains(&chains)
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
fn fused_quantize_scale_pack<const BITS: usize>(
rot_orig: &[f32],
shift: &[f32],
scale_tq: &[f32],
inv_scale_tq: &[f32],
centroid_orig: Option<&[f64]>,
boundaries: &[f32],
centroids: &[f32],
norm: f32,
packed_row: &mut [u8],
dim: usize,
bytes_per_plane: usize,
) -> f32 {
use std::arch::aarch64::*;
let chunks = dim / 8;
let mut acc_a;
let mut acc_b;
unsafe {
acc_a = vdupq_n_f64(0.0);
acc_b = vdupq_n_f64(0.0);
}
unsafe {
for c in 0..chunks {
let offset = c * 8;
let vals_lo = vmulq_f32(
vaddq_f32(
vld1q_f32(rot_orig.as_ptr().add(offset)),
vld1q_f32(shift.as_ptr().add(offset)),
),
vld1q_f32(scale_tq.as_ptr().add(offset)),
);
let vals_hi = vmulq_f32(
vaddq_f32(
vld1q_f32(rot_orig.as_ptr().add(offset + 4)),
vld1q_f32(shift.as_ptr().add(offset + 4)),
),
vld1q_f32(scale_tq.as_ptr().add(offset + 4)),
);
let mut acc_lo = vdupq_n_u32(0);
let mut acc_hi = vdupq_n_u32(0);
if BITS == 4 {
let mid = vdupq_n_f32(boundaries[7]);
let m_lo = vcgtq_f32(vals_lo, mid);
let m_hi = vcgtq_f32(vals_hi, mid);
acc_lo = vshlq_n_u32::<3>(vshrq_n_u32::<31>(m_lo));
acc_hi = vshlq_n_u32::<3>(vshrq_n_u32::<31>(m_hi));
for k in 0..7 {
let b_low = vdupq_n_f32(boundaries[k]);
let b_high = vdupq_n_f32(boundaries[8 + k]);
let bv_lo = vbslq_f32(m_lo, b_high, b_low);
let bv_hi = vbslq_f32(m_hi, b_high, b_low);
acc_lo =
vaddq_u32(acc_lo, vshrq_n_u32::<31>(vcgtq_f32(vals_lo, bv_lo)));
acc_hi =
vaddq_u32(acc_hi, vshrq_n_u32::<31>(vcgtq_f32(vals_hi, bv_hi)));
}
} else {
for bi in 0..(1usize << BITS) - 1 {
let bv = vdupq_n_f32(boundaries[bi]);
acc_lo = vaddq_u32(acc_lo, vshrq_n_u32::<31>(vcgtq_f32(vals_lo, bv)));
acc_hi = vaddq_u32(acc_hi, vshrq_n_u32::<31>(vcgtq_f32(vals_hi, bv)));
}
}
let counts: [u8; 8] = [
vgetq_lane_u32::<0>(acc_lo) as u8,
vgetq_lane_u32::<1>(acc_lo) as u8,
vgetq_lane_u32::<2>(acc_lo) as u8,
vgetq_lane_u32::<3>(acc_lo) as u8,
vgetq_lane_u32::<0>(acc_hi) as u8,
vgetq_lane_u32::<1>(acc_hi) as u8,
vgetq_lane_u32::<2>(acc_hi) as u8,
vgetq_lane_u32::<3>(acc_hi) as u8,
];
let mut terms = [0.0f64; 8];
match centroid_orig {
Some(table) => {
for k in 0..8 {
let d = offset + k;
terms[k] = (rot_orig[d] as f64)
* table[d * (1 << BITS) + counts[k] as usize];
}
}
None => {
for k in 0..8 {
let d = offset + k;
let centroid_in_orig = recon_entry(
centroids[counts[k] as usize],
inv_scale_tq[d] as f64,
shift[d] as f64,
);
terms[k] = (rot_orig[d] as f64) * centroid_in_orig;
}
}
}
acc_a = vaddq_f64(acc_a, vld1q_f64(terms.as_ptr()));
acc_b = vaddq_f64(acc_b, vld1q_f64(terms.as_ptr().add(2)));
acc_a = vaddq_f64(acc_a, vld1q_f64(terms.as_ptr().add(4)));
acc_b = vaddq_f64(acc_b, vld1q_f64(terms.as_ptr().add(6)));
let codes_vec = vld1_u8(counts.as_ptr());
let weights: [u8; 8] = [128, 64, 32, 16, 8, 4, 2, 1];
let wv = vld1_u8(weights.as_ptr());
for p in 0..BITS {
let mask = vdup_n_u8(1u8 << p);
let hit = vcgt_u8(vand_u8(codes_vec, mask), vdup_n_u8(0));
packed_row[p * bytes_per_plane + offset / 8] = vaddv_u8(vand_u8(hit, wv));
}
}
}
let inner = unsafe {
(vgetq_lane_f64::<0>(acc_a) + vgetq_lane_f64::<1>(acc_a))
+ (vgetq_lane_f64::<0>(acc_b) + vgetq_lane_f64::<1>(acc_b))
};
scale_from_inner(inner, norm)
}
const DEGENERATE_INNER_EPS: f64 = 0.1;
#[inline(always)]
fn scale_from_inner(inner: f64, norm: f32) -> f32 {
if inner > DEGENERATE_INNER_EPS {
norm / inner as f32
} else {
0.0
}
}
#[cfg(target_arch = "x86_64")]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
fn fused_quantize_scale_pack<const BITS: usize>(
rot_orig: &[f32],
shift: &[f32],
scale_tq: &[f32],
inv_scale_tq: &[f32],
centroid_orig: Option<&[f64]>,
boundaries: &[f32],
centroids: &[f32],
norm: f32,
packed_row: &mut [u8],
dim: usize,
bytes_per_plane: usize,
) -> f32 {
if std::arch::is_x86_feature_detected!("avx2") {
unsafe {
fused_quantize_scale_pack_avx2::<BITS>(
rot_orig, shift, scale_tq, inv_scale_tq, centroid_orig,
boundaries, centroids, norm, packed_row, dim, bytes_per_plane,
)
}
} else {
fused_quantize_scale_pack_scalar::<BITS>(
rot_orig, shift, scale_tq, inv_scale_tq, centroid_orig,
boundaries, centroids, norm, packed_row, dim, bytes_per_plane,
)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
#[allow(clippy::too_many_arguments)]
unsafe fn fused_quantize_scale_pack_avx2<const BITS: usize>(
rot_orig: &[f32],
shift: &[f32],
scale_tq: &[f32],
inv_scale_tq: &[f32],
centroid_orig: Option<&[f64]>,
boundaries: &[f32],
centroids: &[f32],
norm: f32,
packed_row: &mut [u8],
dim: usize,
bytes_per_plane: usize,
) -> f32 {
use std::arch::x86_64::*;
let mut acc4 = _mm256_setzero_pd();
let chunks = dim / 8;
let _ = centroid_orig;
let mut cpad = [0.0f32; 8];
let mut cpad_hi = [0.0f32; 8];
for (i, slot) in cpad.iter_mut().enumerate() {
*slot = centroids[i.min((1usize << BITS) - 1)];
}
if BITS == 4 {
for (i, slot) in cpad_hi.iter_mut().enumerate() {
*slot = centroids[8 + i];
}
}
let cvec = _mm256_loadu_ps(cpad.as_ptr());
let cvec_hi = _mm256_loadu_ps(cpad_hi.as_ptr());
for c in 0..chunks {
let offset = c * 8;
let vals = _mm256_mul_ps(
_mm256_add_ps(
_mm256_loadu_ps(rot_orig.as_ptr().add(offset)),
_mm256_loadu_ps(shift.as_ptr().add(offset)),
),
_mm256_loadu_ps(scale_tq.as_ptr().add(offset)),
);
let mut acc = _mm256_setzero_si256();
if BITS == 4 {
let mid = _mm256_set1_ps(boundaries[7]);
let m = _mm256_cmp_ps::<_CMP_GT_OQ>(vals, mid);
acc = _mm256_slli_epi32::<3>(_mm256_srli_epi32::<31>(_mm256_castps_si256(m)));
for k in 0..7 {
let b_low = _mm256_set1_ps(boundaries[k]);
let b_high = _mm256_set1_ps(boundaries[8 + k]);
let bv = _mm256_blendv_ps(b_low, b_high, m);
let gt = _mm256_cmp_ps::<_CMP_GT_OQ>(vals, bv);
acc = _mm256_sub_epi32(acc, _mm256_castps_si256(gt));
}
} else if BITS == 2 {
let mid = _mm256_set1_ps(boundaries[1]);
let m = _mm256_cmp_ps::<_CMP_GT_OQ>(vals, mid);
acc = _mm256_slli_epi32::<1>(_mm256_srli_epi32::<31>(_mm256_castps_si256(m)));
let bv = _mm256_blendv_ps(
_mm256_set1_ps(boundaries[0]),
_mm256_set1_ps(boundaries[2]),
m,
);
let gt = _mm256_cmp_ps::<_CMP_GT_OQ>(vals, bv);
acc = _mm256_sub_epi32(acc, _mm256_castps_si256(gt));
} else {
for bi in 0..(1usize << BITS) - 1 {
let bv = _mm256_set1_ps(boundaries[bi]);
let gt = _mm256_cmp_ps::<_CMP_GT_OQ>(vals, bv);
acc = _mm256_sub_epi32(acc, _mm256_castps_si256(gt));
}
}
let rev = _mm256_permutevar8x32_epi32(
acc,
_mm256_setr_epi32(7, 6, 5, 4, 3, 2, 1, 0),
);
for p in 0..BITS {
let bit = _mm256_sll_epi32(rev, _mm_cvtsi32_si128(31 - p as i32));
let m = _mm256_movemask_ps(_mm256_castsi256_ps(bit)) as u8;
*packed_row.get_unchecked_mut(p * bytes_per_plane + c) = m;
}
let sel = if BITS == 4 {
let low3 = _mm256_and_si256(acc, _mm256_set1_epi32(7));
let lo = _mm256_permutevar8x32_ps(cvec, low3);
let hi = _mm256_permutevar8x32_ps(cvec_hi, low3);
let use_hi = _mm256_cmpgt_epi32(acc, _mm256_set1_epi32(7));
_mm256_blendv_ps(lo, hi, _mm256_castsi256_ps(use_hi))
} else {
_mm256_permutevar8x32_ps(cvec, acc)
};
let x_lo = _mm256_sub_pd(
_mm256_mul_pd(
_mm256_cvtps_pd(_mm256_castps256_ps128(sel)),
_mm256_cvtps_pd(_mm_loadu_ps(inv_scale_tq.as_ptr().add(offset))),
),
_mm256_cvtps_pd(_mm_loadu_ps(shift.as_ptr().add(offset))),
);
let x_hi = _mm256_sub_pd(
_mm256_mul_pd(
_mm256_cvtps_pd(_mm256_extractf128_ps::<1>(sel)),
_mm256_cvtps_pd(_mm_loadu_ps(inv_scale_tq.as_ptr().add(offset + 4))),
),
_mm256_cvtps_pd(_mm_loadu_ps(shift.as_ptr().add(offset + 4))),
);
let rot_lo = _mm256_cvtps_pd(_mm_loadu_ps(rot_orig.as_ptr().add(offset)));
let rot_hi = _mm256_cvtps_pd(_mm_loadu_ps(rot_orig.as_ptr().add(offset + 4)));
acc4 = _mm256_add_pd(acc4, _mm256_mul_pd(rot_lo, x_lo));
acc4 = _mm256_add_pd(acc4, _mm256_mul_pd(rot_hi, x_hi));
}
let mut chains = [0.0f64; 4];
_mm256_storeu_pd(chains.as_mut_ptr(), acc4);
let inner = (chains[0] + chains[1]) + (chains[2] + chains[3]);
scale_from_inner(inner, norm)
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
fn fused_quantize_scale_pack<const BITS: usize>(
rot_orig: &[f32],
shift: &[f32],
scale_tq: &[f32],
inv_scale_tq: &[f32],
centroid_orig: Option<&[f64]>,
boundaries: &[f32],
centroids: &[f32],
norm: f32,
packed_row: &mut [u8],
dim: usize,
bytes_per_plane: usize,
) -> f32 {
fused_quantize_scale_pack_scalar::<BITS>(
rot_orig, shift, scale_tq, inv_scale_tq, centroid_orig,
boundaries, centroids, norm, packed_row, dim, bytes_per_plane,
)
}
#[cfg_attr(target_arch = "aarch64", allow(dead_code))]
#[allow(clippy::too_many_arguments)]
#[inline(always)]
fn fused_quantize_scale_pack_scalar<const BITS: usize>(
rot_orig: &[f32],
shift: &[f32],
scale_tq: &[f32],
inv_scale_tq: &[f32],
centroid_orig: Option<&[f64]>,
boundaries: &[f32],
centroids: &[f32],
norm: f32,
packed_row: &mut [u8],
dim: usize,
bytes_per_plane: usize,
) -> f32 {
let mut chains = [0.0f64; 4];
let chunks = dim / 8;
for c in 0..chunks {
let offset = c * 8;
let mut codes = [0u8; 8];
for (k, code) in codes.iter_mut().enumerate() {
let j = offset + k;
let calib = (rot_orig[j] + shift[j]) * scale_tq[j];
let mut v = 0u8;
for bi in 0..(1usize << BITS) - 1 {
if calib > boundaries[bi] { v += 1; }
}
*code = v;
let centroid_in_orig = match centroid_orig {
Some(table) => table[j * (1 << BITS) + v as usize],
None => {
recon_entry(centroids[v as usize], inv_scale_tq[j] as f64, shift[j] as f64)
}
};
chains[j % 4] += (rot_orig[j] as f64) * centroid_in_orig;
}
for p in 0..BITS {
let mut byte = 0u8;
for (k, &code) in codes.iter().enumerate() {
byte |= ((code >> p) & 1) << (7 - k);
}
packed_row[p * bytes_per_plane + c] = byte;
}
}
let inner = (chains[0] + chains[1]) + (chains[2] + chains[3]);
scale_from_inner(inner, norm)
}
#[cfg(test)]
mod simd_identity_tests {
use super::*;
use crate::codebook;
use crate::rotation::Rotation;
fn pseudo_rows(n: usize, dim: usize, seed: u64) -> Vec<f32> {
let mut x = seed;
(0..n * dim)
.map(|_| {
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
(x as f64 / u64::MAX as f64) as f32 - 0.5
})
.collect()
}
#[test]
fn quantize_kernel_matches_scalar_bit_exactly() {
crate::rotation::tests::require_simd_features();
fn run<const BITS: usize>(dim: usize) {
let rotation = Rotation::new(dim);
let (boundaries, centroids) = codebook::codebook(BITS, dim);
let shift: Vec<f32> = (0..dim).map(|d| (d as f32 * 0.001) - 0.01).collect();
let scale_tq: Vec<f32> = (0..dim).map(|d| 1.0 + (d as f32 * 0.0005)).collect();
let inv_scale_tq: Vec<f32> = scale_tq.iter().map(|s| 1.0 / s).collect();
let table = build_recon_table(BITS, dim, ¢roids, &inv_scale_tq, &shift);
let n = 4;
let raw = pseudo_rows(n, dim, 0xD1536 + BITS as u64);
let bytes_per_plane = dim / 8;
let bytes_per_row = BITS * bytes_per_plane;
let mut scratch = vec![0.0f32; dim];
for i in 0..n {
let mut rot = vec![0.0f32; dim];
let src = &raw[i * dim..(i + 1) * dim];
let norm = src.iter().map(|x| x * x).sum::<f32>().sqrt();
rotation.apply_scaled_into(src, 1.0 / norm, &mut rot, &mut scratch);
let mut by_setting: Vec<(Vec<u8>, f32)> = Vec::new();
for table_opt in [None, Some(table.as_slice())] {
let mut expect = vec![0u8; bytes_per_row];
let scale_ref = fused_quantize_scale_pack_scalar::<BITS>(
&rot, &shift, &scale_tq, &inv_scale_tq, table_opt,
&boundaries, ¢roids, norm, &mut expect, dim,
bytes_per_plane,
);
let mut paths: Vec<(&str, Vec<u8>, f32)> = Vec::new();
{
let mut p = vec![0u8; bytes_per_row];
let sc = fused_quantize_scale_pack::<BITS>(
&rot, &shift, &scale_tq, &inv_scale_tq, table_opt,
&boundaries, ¢roids, norm, &mut p, dim,
bytes_per_plane,
);
paths.push(("dispatch", p, sc));
}
#[cfg(target_arch = "x86_64")]
if std::arch::is_x86_feature_detected!("avx2") {
let mut p = vec![0u8; bytes_per_row];
let sc = unsafe {
fused_quantize_scale_pack_avx2::<BITS>(
&rot, &shift, &scale_tq, &inv_scale_tq, table_opt,
&boundaries, ¢roids, norm, &mut p, dim,
bytes_per_plane,
)
};
paths.push(("avx2", p, sc));
}
for (name, packed, scale) in &paths {
assert_eq!(
packed, &expect,
"BITS={BITS} row {i} table={} path={name} packed bytes diverge",
table_opt.is_some()
);
assert_eq!(
scale.to_bits(),
scale_ref.to_bits(),
"BITS={BITS} row {i} table={} path={name} scale diverges: \
{scale} vs {scale_ref}",
table_opt.is_some()
);
}
by_setting.push((expect, scale_ref));
}
let (ref inline_packed, inline_scale) = by_setting[0];
let (ref table_packed, table_scale) = by_setting[1];
assert_eq!(
inline_packed, table_packed,
"BITS={BITS} row {i}: recon-table and inline paths packed \
different bytes. RECON_TABLE_MIN_ROWS would then be a \
format switch — a 15-row batch and a 16-row batch would \
encode the same vector differently.",
);
assert_eq!(
inline_scale.to_bits(),
table_scale.to_bits(),
"BITS={BITS} row {i}: recon-table and inline paths produced \
different stored scales ({inline_scale} vs {table_scale}). \
The table must apply exactly the ops the kernel applies \
inline, in the same order and the same widths.",
);
}
}
run::<2>(1536);
run::<3>(128);
run::<4>(1536);
run::<2>(3072);
run::<4>(3072);
}
#[test]
fn recon_table_entries_are_bit_identical_to_the_inline_expression() {
fn run<const BITS: usize>(dim: usize) {
let (_, centroids) = codebook::codebook(BITS, dim);
let shift: Vec<f32> = (0..dim).map(|d| (d as f32 * 0.001) - 0.01).collect();
let inv_scale_tq: Vec<f32> =
(0..dim).map(|d| 1.0 / (1.0 + (d as f32 * 0.0005))).collect();
let table = build_recon_table(BITS, dim, ¢roids, &inv_scale_tq, &shift);
let n_codes = 1usize << BITS;
assert_eq!(
table.len(),
n_codes * dim,
"BITS={BITS} dim={dim}: table is not coordinate-major with \
2^BITS entries per coordinate",
);
for d in 0..dim {
for c in 0..n_codes {
let want =
(centroids[c] as f64) * (inv_scale_tq[d] as f64) - (shift[d] as f64);
let got = table[d * n_codes + c];
assert_eq!(
got.to_bits(),
want.to_bits(),
"BITS={BITS} dim={dim} d={d} code={c}: table entry \
{got:e} is not bit-identical to the inline \
expression {want:e}. RECON_TABLE_MIN_ROWS is then a \
format switch at f64 precision even if the packed \
bytes happen to round the same way today.",
);
}
}
}
run::<2>(1536);
run::<3>(128);
run::<4>(1536);
run::<2>(3072);
run::<4>(3072);
}
#[test]
fn f32_sort_key_is_order_preserving_and_invertible() {
let mut vals: Vec<f32> = pseudo_rows(1, 4096, 0xC0FFEE);
vals.extend_from_slice(&[
0.0, -0.0, 1.0, -1.0, f32::MIN_POSITIVE, -f32::MIN_POSITIVE,
f32::MAX, f32::MIN, 1e-30, -1e-30, 3.5, -3.5,
]);
for &x in &vals {
assert_eq!(
f32_from_sort_key(f32_sort_key(x)).to_bits(),
x.to_bits(),
"round-trip failed for {x}",
);
}
for &a in &vals {
for &b in &vals {
if a == 0.0 && b == 0.0 {
continue;
}
assert_eq!(
f32_sort_key(a) < f32_sort_key(b),
a < b,
"key order disagrees for ({a}, {b})",
);
}
}
}
#[test]
fn key_selection_matches_float_selection() {
for n in [1usize, 2, 17, 1000, 4096] {
let vals = pseudo_rows(1, n, 0xBEEF ^ n as u64);
for &rank in &[0usize, n / 20, n / 2, n - 1] {
let mut by_val = vals.clone();
let (_, v, _) = by_val.select_nth_unstable_by(rank, |a: &f32, b: &f32| {
a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal)
});
let expect = *v;
let mut by_key: Vec<u32> = vals.iter().copied().map(f32_sort_key).collect();
let (_, k, _) = by_key.select_nth_unstable(rank);
assert_eq!(
f32_from_sort_key(*k).to_bits(),
expect.to_bits(),
"n={n} rank={rank}",
);
}
}
}
#[test]
fn norm_simd_matches_scalar_bit_exactly() {
crate::rotation::tests::require_simd_features();
for len in [8usize, 16, 24, 64, 200, 768, 1000, 1536, 3072, 1, 5, 7, 9, 15] {
for seed in [1u64, 0xD1536, 0xFFFF_FFFF] {
let row = pseudo_rows(1, len, seed);
let scalar = norm_sq_scalar(&row).sqrt();
let mut paths: Vec<(&str, f32)> = vec![("dispatch", simd_norm(&row))];
#[cfg(target_arch = "x86_64")]
if std::arch::is_x86_feature_detected!("avx") {
paths.push(("avx", unsafe { norm_sq_avx(&row) }.sqrt()));
}
#[cfg(target_arch = "aarch64")]
paths.push(("neon", unsafe { norm_sq_neon(&row) }.sqrt()));
for (name, got) in &paths {
assert_eq!(
got.to_bits(),
scalar.to_bits(),
"len={len} seed={seed} path={name}: norm {got} != scalar {scalar}",
);
}
}
}
}
#[test]
fn norm_reduction_order_is_frozen() {
let row: Vec<f32> = (0..16).map(|i| (i as f32) + 1.0).collect();
let mut chains = [0.0f32; NORM_CHAINS];
for (j, &x) in row.iter().enumerate() {
chains[j % NORM_CHAINS] += x * x;
}
let expect =
(((chains[0] + chains[1]) + (chains[2] + chains[3]))
+ ((chains[4] + chains[5]) + (chains[6] + chains[7])))
.sqrt();
assert_eq!(simd_norm(&row).to_bits(), expect.to_bits());
assert_eq!(NORM_CHAINS, 8, "NORM_CHAINS is part of the encode contract");
assert!((simd_norm(&row) - 1496.0f32.sqrt()).abs() < 1e-3);
}
#[test]
fn validation_matches_scalar_exactly() {
let n = 100;
let clean = pseudo_rows(1, n, 7);
assert_eq!(
first_invalid_in_chunk(&clean, 1e16),
first_invalid_in_chunk_scalar(&clean, 1e16)
);
for bad in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY, 1e16, -2e16] {
for pos in [0usize, 1, 7, 8, 15, 63, 64, 96, 99] {
let mut v = clean.clone();
v[pos] = bad;
assert_eq!(
first_invalid_in_chunk(&v, 1e16),
first_invalid_in_chunk_scalar(&v, 1e16),
"bad={bad} pos={pos}"
);
assert_eq!(first_invalid_in_chunk(&v, 1e16), Some(pos));
}
}
}
}
#[cfg(test)]
mod anchor_tests {
use super::*;
use crate::codebook;
fn heavy_tailed(n: usize, dim: usize, seed: u64) -> Vec<f32> {
let mut x = seed;
let mut next = || {
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
x as f64 / u64::MAX as f64
};
let sd = 1.0 / (dim as f64).sqrt();
(0..n * dim)
.map(|_| {
let u = next();
let g = (next() - 0.5) * 2.0;
let scale = if u < 0.02 { 12.0 } else { 1.0 };
(g * sd * scale) as f32
})
.collect()
}
#[test]
fn overload_does_not_worsen_with_more_bits() {
let (n, dim) = (4000usize, 64usize);
let rotated = heavy_tailed(n, dim, 0xC0FFEE_99);
let mut prev = f32::INFINITY;
for bits in [2usize, 3, 4] {
let (_, centroids) = codebook::codebook(bits, dim);
let c_outer = centroids.iter().fold(0.0f32, |a, &c| a.max(c.abs()));
let (shift, scale) = compute_tqplus_calibration(&rotated, n, dim, ¢roids);
let mut cal: Vec<f32> = Vec::with_capacity(n * dim);
for i in 0..n {
for d in 0..dim {
cal.push(((rotated[i * dim + d] + shift[d]) * scale[d]).abs());
}
}
cal.sort_by(|a, b| a.partial_cmp(b).expect("values are finite"));
let p999 = cal[(cal.len() as f64 * 0.999) as usize];
let ratio = p999 / c_outer;
assert!(
ratio <= prev,
"bits={bits}: tail overload ratio rose to {ratio} from {prev} — \
adding bits made the fit worse, which is the #454 failure mode"
);
prev = ratio;
}
}
#[test]
fn the_anchor_probability_widens_with_bit_width() {
let dim = 1536usize;
let a = (dim as f64 - 1.0) / 2.0;
let beta = Beta::new(a, a).expect("Beta(a, a) is valid for a > 0");
let mut prev = 0.0f64;
for bits in [2usize, 3, 4] {
let (_, centroids) = codebook::codebook(bits, dim);
let (p_lo, p_hi, qc_lo, qc_hi) = tqplus_anchor(&beta, ¢roids);
assert!(
p_hi > prev,
"bits={bits}: anchor probability {p_hi} did not widen past {prev}"
);
assert!((p_lo + p_hi - 1.0).abs() < 1e-12, "anchor must stay symmetric");
assert_eq!(qc_lo, -qc_hi, "canonical targets must stay symmetric");
let c_outer = centroids.iter().fold(0.0f32, |a, &c| a.max(c.abs()));
assert_eq!(qc_hi, c_outer);
prev = p_hi;
}
}
}