use rand::{RngCore, SeedableRng};
use rand_chacha::ChaCha8Rng;
pub const K: usize = 2;
const ROTATION_SEED: [u8; 32] = [
164, 143, 161, 123, 88, 50, 61, 10, 234, 184, 161, 204, 105, 1, 20, 184, 43, 140, 200,
117, 24, 180, 247, 84, 141, 68, 110, 161, 228, 223, 32, 242,
];
fn block_size(dim: usize) -> usize {
debug_assert!(dim > 0 && dim % 8 == 0);
dim & dim.wrapping_neg()
}
#[derive(Debug, Clone)]
pub struct Rotation {
dim: usize,
block: usize,
inv_sqrt_block: f32,
signs: Vec<Vec<f32>>,
perms: Vec<Vec<u32>>,
signs1_pre: Vec<f32>,
}
impl Rotation {
pub fn new(dim: usize) -> Self {
assert!(dim > 0 && dim % 8 == 0, "rotation dim must be a positive multiple of 8");
assert!(
dim <= crate::MAX_DIM,
"rotation dim {dim} exceeds MAX_DIM ({})",
crate::MAX_DIM,
);
let block = block_size(dim);
let inv_sqrt_block = 1.0 / (block as f32).sqrt();
let mut rng = ChaCha8Rng::from_seed(ROTATION_SEED);
let mut signs = Vec::with_capacity(K);
let mut perms = Vec::with_capacity(K);
for _ in 0..K {
let sign_row: Vec<f32> = (0..dim)
.map(|_| if rng.next_u32() & 1 == 1 { -1.0 } else { 1.0 })
.collect();
signs.push(sign_row);
perms.push(fisher_yates(dim, &mut rng));
}
let mut signs1_pre = vec![1.0f32; dim];
for (i, &p) in perms[1].iter().enumerate() {
signs1_pre[p as usize] = signs[1][i];
}
Self { dim, block, inv_sqrt_block, signs, perms, signs1_pre }
}
pub fn dim(&self) -> usize {
self.dim
}
pub fn apply(&self, row: &mut [f32]) {
let mut scratch = vec![0.0f32; self.dim];
self.apply_with_scratch(row, &mut scratch);
}
pub fn apply_scaled_into(
&self,
src: &[f32],
inv: f32,
dst: &mut [f32],
scratch: &mut [f32],
) {
assert_eq!(src.len(), self.dim, "rotation input row must have length dim");
assert_eq!(dst.len(), self.dim, "rotation output row must have length dim");
assert_eq!(scratch.len(), self.dim, "rotation scratch must have length dim");
let dim = self.dim;
let block = self.block;
const _: () = assert!(K == 2, "buffer schedule below is written for K = 2");
let wht = |buf: &mut [f32]| {
let mut offset = 0;
while offset < dim {
wht_block(&mut buf[offset..offset + block], block, self.inv_sqrt_block);
offset += block;
}
};
permute_gather::<2>(src, &self.perms[0], &self.signs[0], inv, scratch);
wht(scratch);
for (x, &sg) in scratch.iter_mut().zip(self.signs1_pre.iter()) {
*x *= sg;
}
permute_gather::<0>(scratch, &self.perms[1], &self.signs1_pre, 1.0, dst);
wht(dst);
}
pub fn apply_with_scratch(&self, row: &mut [f32], scratch: &mut [f32]) {
assert_eq!(row.len(), self.dim, "rotation input row must have length dim");
assert_eq!(scratch.len(), self.dim, "rotation scratch must have length dim");
let dim = self.dim;
let block = self.block;
const _: () = assert!(K % 2 == 0, "ping-pong ends in `row` only for even K");
let (mut input, mut output): (&mut [f32], &mut [f32]) = (row, scratch);
for round in 0..K {
let perm = &self.perms[round];
let sign_row = &self.signs[round];
permute_gather::<1>(input, perm, sign_row, 1.0, output);
let mut offset = 0;
while offset < dim {
let blk = &mut output[offset..offset + block];
wht_block(blk, block, self.inv_sqrt_block);
offset += block;
}
std::mem::swap(&mut input, &mut output);
}
}
}
#[inline(always)]
fn permute_gather<const MODE: usize>(
src: &[f32],
perm: &[u32],
signs: &[f32],
inv: f32,
dst: &mut [f32],
) {
assert_eq!(src.len(), dst.len(), "permute_gather: src length must equal dst");
assert_eq!(perm.len(), dst.len(), "permute_gather: perm length must equal dst");
assert_eq!(signs.len(), dst.len(), "permute_gather: signs length must equal dst");
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx512f")
&& std::arch::is_x86_feature_detected!("avx2")
{
unsafe { return permute_gather_avx512::<MODE>(src, perm, signs, inv, dst) }
} else if std::arch::is_x86_feature_detected!("avx2") {
unsafe { return permute_gather_avx2::<MODE>(src, perm, signs, inv, dst) }
}
}
permute_gather_scalar::<MODE>(src, perm, signs, inv, dst)
}
#[inline(always)]
fn permute_gather_scalar<const MODE: usize>(
src: &[f32],
perm: &[u32],
signs: &[f32],
inv: f32,
dst: &mut [f32],
) {
for ((d, &p), &s) in dst.iter_mut().zip(perm.iter()).zip(signs.iter()) {
let v = src[p as usize];
*d = match MODE {
0 => v,
1 => v * s,
_ => (v * inv) * s,
};
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn permute_gather_avx2<const MODE: usize>(
src: &[f32],
perm: &[u32],
signs: &[f32],
inv: f32,
dst: &mut [f32],
) {
use std::arch::x86_64::*;
let n = dst.len();
let base = src.as_ptr();
let invv = _mm256_set1_ps(inv);
let mut i = 0;
while i + 8 <= n {
let idx = _mm256_loadu_si256(perm.as_ptr().add(i) as *const __m256i);
let mut v = _mm256_i32gather_ps::<4>(base, idx);
if MODE == 2 {
v = _mm256_mul_ps(v, invv);
}
if MODE != 0 {
v = _mm256_mul_ps(v, _mm256_loadu_ps(signs.as_ptr().add(i)));
}
_mm256_storeu_ps(dst.as_mut_ptr().add(i), v);
i += 8;
}
while i < n {
let v = *src.get_unchecked(*perm.get_unchecked(i) as usize);
let s = *signs.get_unchecked(i);
*dst.get_unchecked_mut(i) = match MODE {
0 => v,
1 => v * s,
_ => (v * inv) * s,
};
i += 1;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f", enable = "avx2")]
unsafe fn permute_gather_avx512<const MODE: usize>(
src: &[f32],
perm: &[u32],
signs: &[f32],
inv: f32,
dst: &mut [f32],
) {
use std::arch::x86_64::*;
let n = dst.len();
let base = src.as_ptr();
let invv = _mm512_set1_ps(inv);
let mut i = 0;
while i + 16 <= n {
let idx = _mm512_loadu_si512(perm.as_ptr().add(i) as *const __m512i);
let mut v = _mm512_i32gather_ps::<4>(idx, base);
if MODE == 2 {
v = _mm512_mul_ps(v, invv);
}
if MODE != 0 {
v = _mm512_mul_ps(v, _mm512_loadu_ps(signs.as_ptr().add(i)));
}
_mm512_storeu_ps(dst.as_mut_ptr().add(i), v);
i += 16;
}
while i + 8 <= n {
let idx = _mm256_loadu_si256(perm.as_ptr().add(i) as *const __m256i);
let mut v = _mm256_i32gather_ps::<4>(src.as_ptr(), idx);
if MODE == 2 {
v = _mm256_mul_ps(v, _mm256_set1_ps(inv));
}
if MODE != 0 {
v = _mm256_mul_ps(v, _mm256_loadu_ps(signs.as_ptr().add(i)));
}
_mm256_storeu_ps(dst.as_mut_ptr().add(i), v);
i += 8;
}
while i < n {
let v = *src.get_unchecked(*perm.get_unchecked(i) as usize);
let s = *signs.get_unchecked(i);
*dst.get_unchecked_mut(i) = match MODE {
0 => v,
1 => v * s,
_ => (v * inv) * s,
};
i += 1;
}
}
#[cfg(target_arch = "aarch64")]
#[inline(always)]
fn wht_block(blk: &mut [f32], block: usize, inv_sqrt_block: f32) {
use std::arch::aarch64::*;
debug_assert!(block >= 8 && block.is_power_of_two());
let p = blk.as_mut_ptr();
unsafe {
let mut j = 0;
while j < block {
let ab = vld2q_f32(p.add(j));
let s1 = vaddq_f32(ab.0, ab.1); let d1 = vsubq_f32(ab.0, ab.1); let s_even = vtrn1q_f32(s1, d1); let s_odd = vtrn2q_f32(s1, d1); let sum2 = vaddq_f32(s_even, s_odd); let dif2 = vsubq_f32(s_even, s_odd); let out0 = vcombine_f32(vget_low_f32(sum2), vget_low_f32(dif2));
let out1 = vcombine_f32(vget_high_f32(sum2), vget_high_f32(dif2));
vst1q_f32(p.add(j), out0);
vst1q_f32(p.add(j + 4), out1);
j += 8;
}
let mut len = 4;
while 4 * len < block {
let oct = 8 * len;
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = vld1q_f32(p.add(j));
let b = vld1q_f32(p.add(j + len));
let c = vld1q_f32(p.add(j + 2 * len));
let d = vld1q_f32(p.add(j + 3 * len));
let e = vld1q_f32(p.add(j + 4 * len));
let f = vld1q_f32(p.add(j + 5 * len));
let g = vld1q_f32(p.add(j + 6 * len));
let h = vld1q_f32(p.add(j + 7 * len));
let apb = vaddq_f32(a, b);
let amb = vsubq_f32(a, b);
let cpd = vaddq_f32(c, d);
let cmd = vsubq_f32(c, d);
let epf = vaddq_f32(e, f);
let emf = vsubq_f32(e, f);
let gph = vaddq_f32(g, h);
let gmh = vsubq_f32(g, h);
let s0 = vaddq_f32(apb, cpd);
let s1 = vaddq_f32(amb, cmd);
let s2 = vsubq_f32(apb, cpd);
let s3 = vsubq_f32(amb, cmd);
let s4 = vaddq_f32(epf, gph);
let s5 = vaddq_f32(emf, gmh);
let s6 = vsubq_f32(epf, gph);
let s7 = vsubq_f32(emf, gmh);
vst1q_f32(p.add(j), vaddq_f32(s0, s4));
vst1q_f32(p.add(j + len), vaddq_f32(s1, s5));
vst1q_f32(p.add(j + 2 * len), vaddq_f32(s2, s6));
vst1q_f32(p.add(j + 3 * len), vaddq_f32(s3, s7));
vst1q_f32(p.add(j + 4 * len), vsubq_f32(s0, s4));
vst1q_f32(p.add(j + 5 * len), vsubq_f32(s1, s5));
vst1q_f32(p.add(j + 6 * len), vsubq_f32(s2, s6));
vst1q_f32(p.add(j + 7 * len), vsubq_f32(s3, s7));
j += 4;
}
i += oct;
}
len <<= 3;
}
if 2 * len < block {
let quad = 4 * len;
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = vld1q_f32(p.add(j));
let b = vld1q_f32(p.add(j + len));
let c = vld1q_f32(p.add(j + 2 * len));
let d = vld1q_f32(p.add(j + 3 * len));
let apb = vaddq_f32(a, b);
let amb = vsubq_f32(a, b);
let cpd = vaddq_f32(c, d);
let cmd = vsubq_f32(c, d);
vst1q_f32(p.add(j), vaddq_f32(apb, cpd));
vst1q_f32(p.add(j + len), vaddq_f32(amb, cmd));
vst1q_f32(p.add(j + 2 * len), vsubq_f32(apb, cpd));
vst1q_f32(p.add(j + 3 * len), vsubq_f32(amb, cmd));
j += 4;
}
i += quad;
}
len <<= 2;
}
if len < block {
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = vld1q_f32(p.add(j));
let b = vld1q_f32(p.add(j + len));
vst1q_f32(p.add(j), vaddq_f32(a, b));
vst1q_f32(p.add(j + len), vsubq_f32(a, b));
j += 4;
}
i += 2 * len;
}
}
let sv = vdupq_n_f32(inv_sqrt_block);
let mut j = 0;
while j < block {
vst1q_f32(p.add(j), vmulq_f32(vld1q_f32(p.add(j)), sv));
j += 4;
}
}
}
#[cfg(target_arch = "x86_64")]
#[inline(always)]
fn wht_block(blk: &mut [f32], block: usize, inv_sqrt_block: f32) {
if std::arch::is_x86_feature_detected!("avx512f")
&& std::arch::is_x86_feature_detected!("avx2")
{
unsafe { wht_block_avx512(blk, block, inv_sqrt_block) }
} else if std::arch::is_x86_feature_detected!("avx2") {
unsafe { wht_block_avx2(blk, block, inv_sqrt_block) }
} else {
wht_block_scalar(blk, block, inv_sqrt_block)
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn wht_block_avx2(blk: &mut [f32], block: usize, inv_sqrt_block: f32) {
use std::arch::x86_64::*;
debug_assert!(block >= 8 && block.is_power_of_two());
let p = blk.as_mut_ptr();
let mut j = 0;
while j < block {
let v = _mm256_loadu_ps(p.add(j));
let a = _mm256_shuffle_ps::<0b10_10_00_00>(v, v);
let b = _mm256_shuffle_ps::<0b11_11_01_01>(v, v);
let r1 = _mm256_blend_ps::<0b1010_1010>(
_mm256_add_ps(a, b),
_mm256_sub_ps(a, b),
);
let a = _mm256_shuffle_ps::<0b01_00_01_00>(r1, r1);
let b = _mm256_shuffle_ps::<0b11_10_11_10>(r1, r1);
let r2 = _mm256_blend_ps::<0b1100_1100>(
_mm256_add_ps(a, b),
_mm256_sub_ps(a, b),
);
let a = _mm256_permute2f128_ps::<0x00>(r2, r2);
let b = _mm256_permute2f128_ps::<0x11>(r2, r2);
let r4 = _mm256_blend_ps::<0b1111_0000>(
_mm256_add_ps(a, b),
_mm256_sub_ps(a, b),
);
_mm256_storeu_ps(p.add(j), r4);
j += 8;
}
let mut len = 8;
while 4 * len < block {
let oct = 8 * len;
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = _mm256_loadu_ps(p.add(j));
let b = _mm256_loadu_ps(p.add(j + len));
let c = _mm256_loadu_ps(p.add(j + 2 * len));
let d = _mm256_loadu_ps(p.add(j + 3 * len));
let e = _mm256_loadu_ps(p.add(j + 4 * len));
let f = _mm256_loadu_ps(p.add(j + 5 * len));
let g = _mm256_loadu_ps(p.add(j + 6 * len));
let h = _mm256_loadu_ps(p.add(j + 7 * len));
let apb = _mm256_add_ps(a, b);
let amb = _mm256_sub_ps(a, b);
let cpd = _mm256_add_ps(c, d);
let cmd = _mm256_sub_ps(c, d);
let epf = _mm256_add_ps(e, f);
let emf = _mm256_sub_ps(e, f);
let gph = _mm256_add_ps(g, h);
let gmh = _mm256_sub_ps(g, h);
let s0 = _mm256_add_ps(apb, cpd);
let s1 = _mm256_add_ps(amb, cmd);
let s2 = _mm256_sub_ps(apb, cpd);
let s3 = _mm256_sub_ps(amb, cmd);
let s4 = _mm256_add_ps(epf, gph);
let s5 = _mm256_add_ps(emf, gmh);
let s6 = _mm256_sub_ps(epf, gph);
let s7 = _mm256_sub_ps(emf, gmh);
_mm256_storeu_ps(p.add(j), _mm256_add_ps(s0, s4));
_mm256_storeu_ps(p.add(j + len), _mm256_add_ps(s1, s5));
_mm256_storeu_ps(p.add(j + 2 * len), _mm256_add_ps(s2, s6));
_mm256_storeu_ps(p.add(j + 3 * len), _mm256_add_ps(s3, s7));
_mm256_storeu_ps(p.add(j + 4 * len), _mm256_sub_ps(s0, s4));
_mm256_storeu_ps(p.add(j + 5 * len), _mm256_sub_ps(s1, s5));
_mm256_storeu_ps(p.add(j + 6 * len), _mm256_sub_ps(s2, s6));
_mm256_storeu_ps(p.add(j + 7 * len), _mm256_sub_ps(s3, s7));
j += 8;
}
i += oct;
}
len <<= 3;
}
if 2 * len < block {
let quad = 4 * len;
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = _mm256_loadu_ps(p.add(j));
let b = _mm256_loadu_ps(p.add(j + len));
let c = _mm256_loadu_ps(p.add(j + 2 * len));
let d = _mm256_loadu_ps(p.add(j + 3 * len));
let apb = _mm256_add_ps(a, b);
let amb = _mm256_sub_ps(a, b);
let cpd = _mm256_add_ps(c, d);
let cmd = _mm256_sub_ps(c, d);
_mm256_storeu_ps(p.add(j), _mm256_add_ps(apb, cpd));
_mm256_storeu_ps(p.add(j + len), _mm256_add_ps(amb, cmd));
_mm256_storeu_ps(p.add(j + 2 * len), _mm256_sub_ps(apb, cpd));
_mm256_storeu_ps(p.add(j + 3 * len), _mm256_sub_ps(amb, cmd));
j += 8;
}
i += quad;
}
len <<= 2;
}
if len < block {
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = _mm256_loadu_ps(p.add(j));
let b = _mm256_loadu_ps(p.add(j + len));
_mm256_storeu_ps(p.add(j), _mm256_add_ps(a, b));
_mm256_storeu_ps(p.add(j + len), _mm256_sub_ps(a, b));
j += 8;
}
i += 2 * len;
}
}
let sv = _mm256_set1_ps(inv_sqrt_block);
let mut j = 0;
while j < block {
_mm256_storeu_ps(p.add(j), _mm256_mul_ps(_mm256_loadu_ps(p.add(j)), sv));
j += 8;
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f", enable = "avx2")]
unsafe fn wht_block_avx512(blk: &mut [f32], block: usize, inv_sqrt_block: f32) {
use std::arch::x86_64::*;
debug_assert!(block >= 8 && block.is_power_of_two());
let p = blk.as_mut_ptr();
if block < 16 {
return wht_block_avx2(blk, block, inv_sqrt_block);
}
let mut j = 0;
while j < block {
let v = _mm512_loadu_ps(p.add(j));
let a = _mm512_shuffle_ps::<0b10_10_00_00>(v, v);
let b = _mm512_shuffle_ps::<0b11_11_01_01>(v, v);
let r1 = _mm512_mask_blend_ps(
0b1010_1010_1010_1010,
_mm512_add_ps(a, b),
_mm512_sub_ps(a, b),
);
let a = _mm512_shuffle_ps::<0b01_00_01_00>(r1, r1);
let b = _mm512_shuffle_ps::<0b11_10_11_10>(r1, r1);
let r2 = _mm512_mask_blend_ps(
0b1100_1100_1100_1100,
_mm512_add_ps(a, b),
_mm512_sub_ps(a, b),
);
let a = _mm512_shuffle_f32x4::<0b10_10_00_00>(r2, r2);
let b = _mm512_shuffle_f32x4::<0b11_11_01_01>(r2, r2);
let r4 = _mm512_mask_blend_ps(
0b1111_0000_1111_0000,
_mm512_add_ps(a, b),
_mm512_sub_ps(a, b),
);
let a = _mm512_shuffle_f32x4::<0b01_00_01_00>(r4, r4);
let b = _mm512_shuffle_f32x4::<0b11_10_11_10>(r4, r4);
let r8 = _mm512_mask_blend_ps(
0b1111_1111_0000_0000,
_mm512_add_ps(a, b),
_mm512_sub_ps(a, b),
);
_mm512_storeu_ps(p.add(j), r8);
j += 16;
}
let mut len = 16;
while 4 * len < block {
let oct = 8 * len;
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = _mm512_loadu_ps(p.add(j));
let b = _mm512_loadu_ps(p.add(j + len));
let c = _mm512_loadu_ps(p.add(j + 2 * len));
let d = _mm512_loadu_ps(p.add(j + 3 * len));
let e = _mm512_loadu_ps(p.add(j + 4 * len));
let f = _mm512_loadu_ps(p.add(j + 5 * len));
let g = _mm512_loadu_ps(p.add(j + 6 * len));
let h = _mm512_loadu_ps(p.add(j + 7 * len));
let apb = _mm512_add_ps(a, b);
let amb = _mm512_sub_ps(a, b);
let cpd = _mm512_add_ps(c, d);
let cmd = _mm512_sub_ps(c, d);
let epf = _mm512_add_ps(e, f);
let emf = _mm512_sub_ps(e, f);
let gph = _mm512_add_ps(g, h);
let gmh = _mm512_sub_ps(g, h);
let s0 = _mm512_add_ps(apb, cpd);
let s1 = _mm512_add_ps(amb, cmd);
let s2 = _mm512_sub_ps(apb, cpd);
let s3 = _mm512_sub_ps(amb, cmd);
let s4 = _mm512_add_ps(epf, gph);
let s5 = _mm512_add_ps(emf, gmh);
let s6 = _mm512_sub_ps(epf, gph);
let s7 = _mm512_sub_ps(emf, gmh);
_mm512_storeu_ps(p.add(j), _mm512_add_ps(s0, s4));
_mm512_storeu_ps(p.add(j + len), _mm512_add_ps(s1, s5));
_mm512_storeu_ps(p.add(j + 2 * len), _mm512_add_ps(s2, s6));
_mm512_storeu_ps(p.add(j + 3 * len), _mm512_add_ps(s3, s7));
_mm512_storeu_ps(p.add(j + 4 * len), _mm512_sub_ps(s0, s4));
_mm512_storeu_ps(p.add(j + 5 * len), _mm512_sub_ps(s1, s5));
_mm512_storeu_ps(p.add(j + 6 * len), _mm512_sub_ps(s2, s6));
_mm512_storeu_ps(p.add(j + 7 * len), _mm512_sub_ps(s3, s7));
j += 16;
}
i += oct;
}
len <<= 3;
}
if 2 * len < block {
let quad = 4 * len;
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = _mm512_loadu_ps(p.add(j));
let b = _mm512_loadu_ps(p.add(j + len));
let c = _mm512_loadu_ps(p.add(j + 2 * len));
let d = _mm512_loadu_ps(p.add(j + 3 * len));
let apb = _mm512_add_ps(a, b);
let amb = _mm512_sub_ps(a, b);
let cpd = _mm512_add_ps(c, d);
let cmd = _mm512_sub_ps(c, d);
_mm512_storeu_ps(p.add(j), _mm512_add_ps(apb, cpd));
_mm512_storeu_ps(p.add(j + len), _mm512_add_ps(amb, cmd));
_mm512_storeu_ps(p.add(j + 2 * len), _mm512_sub_ps(apb, cpd));
_mm512_storeu_ps(p.add(j + 3 * len), _mm512_sub_ps(amb, cmd));
j += 16;
}
i += quad;
}
len <<= 2;
}
if len < block {
let mut i = 0;
while i < block {
let mut j = i;
while j < i + len {
let a = _mm512_loadu_ps(p.add(j));
let b = _mm512_loadu_ps(p.add(j + len));
_mm512_storeu_ps(p.add(j), _mm512_add_ps(a, b));
_mm512_storeu_ps(p.add(j + len), _mm512_sub_ps(a, b));
j += 16;
}
i += 2 * len;
}
}
let sv = _mm512_set1_ps(inv_sqrt_block);
let mut j = 0;
while j < block {
_mm512_storeu_ps(p.add(j), _mm512_mul_ps(_mm512_loadu_ps(p.add(j)), sv));
j += 16;
}
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
#[inline(always)]
fn wht_block(blk: &mut [f32], block: usize, inv_sqrt_block: f32) {
wht_block_scalar(blk, block, inv_sqrt_block)
}
#[cfg_attr(target_arch = "aarch64", allow(dead_code))]
#[inline(always)]
fn wht_block_scalar(blk: &mut [f32], block: usize, inv_sqrt_block: f32) {
let mut len = 1;
while len < block {
let mut i = 0;
while i < block {
for j in i..i + len {
let a = blk[j];
let b = blk[j + len];
blk[j] = a + b;
blk[j + len] = a - b;
}
i += 2 * len;
}
len <<= 1;
}
for x in blk.iter_mut() {
*x *= inv_sqrt_block;
}
}
fn fisher_yates(dim: usize, rng: &mut ChaCha8Rng) -> Vec<u32> {
let mut perm: Vec<u32> = (0..dim as u32).collect();
for i in (1..dim).rev() {
let j = (rng.next_u64() % (i as u64 + 1)) as usize;
perm.swap(i, j);
}
perm
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
#[cfg(target_arch = "x86_64")]
pub(crate) fn require_simd_features() {
let Ok(list) = std::env::var("TURBOVEC_REQUIRE_SIMD") else {
return;
};
for feat in list.split(',').map(str::trim).filter(|f| !f.is_empty()) {
let present = match feat {
"avx" => std::arch::is_x86_feature_detected!("avx"),
"avx2" => std::arch::is_x86_feature_detected!("avx2"),
"avx512f" => std::arch::is_x86_feature_detected!("avx512f"),
"avx512bw" => std::arch::is_x86_feature_detected!("avx512bw"),
"avx512vbmi" => std::arch::is_x86_feature_detected!("avx512vbmi"),
other => panic!(
"TURBOVEC_REQUIRE_SIMD lists unknown feature {other:?}"
),
};
assert!(
present,
"TURBOVEC_REQUIRE_SIMD demands {feat:?} but this host does \
not have it — the kernels gated on it would be skipped, \
leaving them untested rather than failing",
);
}
}
#[cfg(not(target_arch = "x86_64"))]
pub(crate) fn require_simd_features() {}
#[test]
fn block_size_is_largest_power_of_two_divisor() {
assert_eq!(block_size(8), 8);
assert_eq!(block_size(200), 8); assert_eq!(block_size(768), 256); assert_eq!(block_size(1000), 8); assert_eq!(block_size(1536), 512); assert_eq!(block_size(3072), 1024); assert_eq!(block_size(1024), 1024); }
#[test]
fn wht_simd_matches_scalar_bit_exactly() {
require_simd_features();
for block in [8usize, 16, 32, 64, 128, 256, 512, 1024, 2048, 4096, 8192, 16384] {
let inv = 1.0 / (block as f32).sqrt();
let mut x = 0x9E3779B97F4A7C15u64;
let buf: Vec<f32> = (0..block)
.map(|_| {
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
(x as f64 / u64::MAX as f64) as f32 - 0.5
})
.collect();
let mut expect = buf.clone();
wht_block_scalar(&mut expect, block, inv);
#[cfg_attr(not(target_arch = "x86_64"), allow(unused_mut))]
let mut checked = vec![("dispatch", {
let mut b = buf.clone();
wht_block(&mut b, block, inv);
b
})];
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx2") {
let mut b = buf.clone();
unsafe { wht_block_avx2(&mut b, block, inv) };
checked.push(("avx2", b));
}
if std::arch::is_x86_feature_detected!("avx512f")
&& std::arch::is_x86_feature_detected!("avx2")
{
let mut b = buf.clone();
unsafe { wht_block_avx512(&mut b, block, inv) };
checked.push(("avx512", b));
}
}
for (name, got) in &checked {
for (i, (a, b)) in got.iter().zip(expect.iter()).enumerate() {
assert_eq!(
a.to_bits(),
b.to_bits(),
"block {block} lane {i}: {name} {a} != scalar {b}"
);
}
}
}
}
#[test]
fn permute_gather_paths_match_scalar_bit_exactly() {
require_simd_features();
for dim in [8usize, 12, 13, 16, 24, 29, 64, 200, 1536] {
let mut x = 0x243F_6A88_85A3_08D3u64 ^ dim as u64;
let mut next = || {
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
x
};
let src: Vec<f32> =
(0..dim).map(|_| (next() as f64 / u64::MAX as f64) as f32 - 0.5).collect();
let signs: Vec<f32> =
(0..dim).map(|_| if next() & 1 == 1 { -1.0 } else { 1.0 }).collect();
let mut rng = ChaCha8Rng::from_seed(ROTATION_SEED);
let perm = fisher_yates(dim, &mut rng);
let inv = 0.812_5_f32;
macro_rules! check_mode {
($mode:literal) => {{
let mut expect = vec![0.0f32; dim];
permute_gather_scalar::<$mode>(&src, &perm, &signs, inv, &mut expect);
let mut got = vec![0.0f32; dim];
permute_gather::<$mode>(&src, &perm, &signs, inv, &mut got);
assert_eq!(got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
expect.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"dim {} mode {} dispatch", dim, $mode);
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx2") {
let mut g = vec![0.0f32; dim];
unsafe {
permute_gather_avx2::<$mode>(&src, &perm, &signs, inv, &mut g)
};
assert_eq!(g.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
expect.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"dim {} mode {} avx2", dim, $mode);
}
if std::arch::is_x86_feature_detected!("avx512f")
&& std::arch::is_x86_feature_detected!("avx2")
{
let mut g = vec![0.0f32; dim];
unsafe {
permute_gather_avx512::<$mode>(&src, &perm, &signs, inv, &mut g)
};
assert_eq!(g.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
expect.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
"dim {} mode {} avx512", dim, $mode);
}
}
}};
}
check_mode!(0);
check_mode!(1);
check_mode!(2);
}
}
#[test]
fn golden_rotation_dim128() {
let rot = Rotation::new(128);
let mut row = vec![0.0f32; 128];
row[0] = 1.0;
rot.apply(&mut row);
let fold = row.iter().fold(0u32, |acc, v| acc.rotate_left(1) ^ v.to_bits());
let head: Vec<u32> = row[..4].iter().map(|v| v.to_bits()).collect();
assert_eq!(
(fold, head[0], head[1], head[2], head[3]),
GOLDEN_DIM128,
"dim=128 rotation output drifted from the frozen v5 bytes",
);
}
#[test]
fn apply_scaled_into_is_bit_identical_to_apply_of_the_scaled_row() {
require_simd_features();
for &dim in &[8usize, 24, 64, 128, 200, 768, 1000, 1024, 1536, 3072] {
let rot = Rotation::new(dim);
let mut state = 0x517C_C1B7_2722_0A95u64 ^ dim as u64;
let mut next = || {
state ^= state << 13;
state ^= state >> 7;
state ^= state << 17;
((state >> 33) as f64 / (1u64 << 31) as f64 - 1.0) as f32
};
for trial in 0..3 {
let src: Vec<f32> = (0..dim).map(|_| next()).collect();
let norm = src.iter().map(|x| x * x).sum::<f32>().sqrt();
for &inv in &[1.0f32, 1.0 / norm, 0.0, 0.8125, -1.0] {
let mut expect: Vec<f32> = src.iter().map(|x| x * inv).collect();
rot.apply(&mut expect);
let mut dst = vec![f32::NAN; dim];
let mut scratch = vec![f32::NAN; dim];
rot.apply_scaled_into(&src, inv, &mut dst, &mut scratch);
for i in 0..dim {
assert_eq!(
dst[i].to_bits(),
expect[i].to_bits(),
"dim={dim} trial={trial} inv={inv} coord {i}: \
apply_scaled_into diverged from apply of the \
pre-scaled row ({} vs {}). This is the function \
that writes every encoded byte — a difference \
here is a format break.",
dst[i],
expect[i],
);
}
}
let mut dst = vec![0.0f32; dim];
let mut scratch = vec![0.0f32; dim];
let before = src.clone();
rot.apply_scaled_into(&src, 0.5, &mut dst, &mut scratch);
assert_eq!(src, before, "dim={dim}: apply_scaled_into mutated src");
}
}
}
#[test]
fn preserves_norm_and_is_deterministic() {
for &dim in &[8usize, 200, 768, 1000, 1536] {
let rot = Rotation::new(dim);
let mut state = 0x1234_5678u64 ^ dim as u64;
let mut v: Vec<f32> = (0..dim)
.map(|_| {
state = state.wrapping_mul(6364136223846793005).wrapping_add(1);
((state >> 33) as f64 / (1u64 << 31) as f64 - 1.0) as f32
})
.collect();
let before = (v.iter().map(|x| x * x).sum::<f32>()).sqrt();
let orig = v.clone();
rot.apply(&mut v);
let after = (v.iter().map(|x| x * x).sum::<f32>()).sqrt();
assert!(
(before - after).abs() / before < 1e-4,
"norm changed at dim={dim}: {before} -> {after}"
);
let mut again = orig;
Rotation::new(dim).apply(&mut again);
assert_eq!(v, again, "rotation not deterministic at dim={dim}");
}
}
}
#[cfg(test)]
const GOLDEN_DIM128: (u32, u32, u32, u32, u32) =
(186507913, 1033895935, 3175088127, 1027604479, 3162505217);