use yo_common::Rng;
const INV_ROOT2: f32 = core::f32::consts::FRAC_1_SQRT_2;
#[derive(Debug)]
struct Round {
signs: Vec<u32>,
pairs: Vec<(u32, u32)>,
}
#[derive(Debug)]
pub struct Rotation {
dim: usize,
seed: u64,
rounds: Vec<Round>,
}
impl Rotation {
#[must_use]
pub fn new(dim: usize, seed: u64) -> Rotation {
assert!(dim > 0, "a vector has at least one dimension");
let mut rng = Rng::new(seed);
let sweeps = sweeps(dim);
let mut rounds = Vec::with_capacity(sweeps);
let mut order: Vec<u32> = (0..dim as u32).collect();
for _ in 0..sweeps {
shuffle(&mut order, &mut rng);
let half = dim / 2;
let pairs: Vec<(u32, u32)> = (0..half)
.map(|k| (order[2 * k], order[2 * k + 1]))
.collect();
let flip = bits(dim, &mut rng);
let turn = bits(half, &mut rng);
let mut signs: Vec<u32> = (0..dim)
.map(|i| u32::from((flip[i / 64] >> (i % 64)) & 1 == 1) << 31)
.collect();
for (k, &(_, j)) in pairs.iter().enumerate() {
if (turn[k / 64] >> (k % 64)) & 1 == 1 {
signs[j as usize] ^= 1 << 31;
}
}
rounds.push(Round { signs, pairs });
}
Rotation { dim, seed, rounds }
}
#[must_use]
pub fn dim(&self) -> usize {
self.dim
}
#[must_use]
pub fn seed(&self) -> u64 {
self.seed
}
pub fn apply(&self, v: &mut [f32]) {
assert_eq!(
v.len(),
self.dim,
"this rotation is for {} dimensions and was handed {}",
self.dim,
v.len()
);
for round in &self.rounds {
for (c, &sign) in v.iter_mut().zip(&round.signs) {
*c = f32::from_bits(c.to_bits() ^ sign);
}
for &(i, j) in &round.pairs {
let (i, j) = (i as usize, j as usize);
let (a, b) = (v[i], v[j]);
v[i] = (a + b) * INV_ROOT2;
v[j] = (a - b) * INV_ROOT2;
}
}
}
}
fn sweeps(dim: usize) -> usize {
((usize::BITS - (dim - 1).leading_zeros()) as usize + 2).max(4)
}
fn shuffle(order: &mut [u32], rng: &mut Rng) {
for i in (1..order.len()).rev() {
order.swap(i, rng.below(i + 1));
}
}
fn bits(n: usize, rng: &mut Rng) -> Vec<u64> {
(0..n.div_ceil(64)).map(|_| rng.next_u64()).collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn dot(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
fn unit(rng: &mut Rng) -> f32 {
(rng.next_u64() >> 40) as f32 / (1u32 << 24) as f32
}
fn sample(dim: usize, n: usize, seed: u64) -> Vec<Vec<f32>> {
let mut rng = Rng::new(seed);
(0..n)
.map(|_| (0..dim).map(|_| unit(&mut rng) * 2.0 - 1.0).collect())
.collect()
}
#[test]
fn a_rotation_keeps_every_length_and_every_angle() {
let r = Rotation::new(64, 7);
let vs = sample(64, 8, 11);
for v in &vs {
let mut spun = v.clone();
r.apply(&mut spun);
let before = dot(v, v).sqrt();
let after = dot(&spun, &spun).sqrt();
assert!((before - after).abs() < 1e-4, "{before} became {after}");
}
for a in 0..vs.len() {
for b in 0..a {
let mut x = vs[a].clone();
let mut y = vs[b].clone();
let before = dot(&x, &y);
r.apply(&mut x);
r.apply(&mut y);
let after = dot(&x, &y);
assert!((before - after).abs() < 1e-3, "{before} became {after}");
}
}
}
#[test]
fn a_seed_gives_the_rotation_it_has_always_given() {
for (dim, want) in [
(8usize, 0xef2f_9c0c_1aad_5cdfu64),
(33, 0x97e7_79cf_5820_4a3f),
(128, 0xde6f_cf6b_69cf_490c),
] {
let mut v: Vec<f32> = (0..dim).map(|i| (i as f32 + 1.0) / 8.0).collect();
Rotation::new(dim, 0xB0A7).apply(&mut v);
let mut got: u64 = 0xcbf2_9ce4_8422_2325;
for c in &v {
for byte in c.to_bits().to_le_bytes() {
got ^= u64::from(byte);
got = got.wrapping_mul(0x0100_0000_01b3);
}
}
assert_eq!(got, want, "the rotation at {dim} dimensions has moved");
}
}
#[test]
fn the_same_seed_is_the_same_rotation() {
let v = sample(32, 1, 3).pop().expect("one vector");
let mut a = v.clone();
let mut b = v;
Rotation::new(32, 99).apply(&mut a);
Rotation::new(32, 99).apply(&mut b);
assert_eq!(a, b);
let mut c = a.clone();
Rotation::new(32, 100).apply(&mut c);
assert_ne!(a, c, "two seeds should not be one rotation");
}
#[test]
fn a_spike_comes_out_flat() {
for dim in [64usize, 256, 768] {
let r = Rotation::new(dim, 5);
let mut v = vec![0.0f32; dim];
v[0] = 1.0;
r.apply(&mut v);
let even = 1.0 / (dim as f32).sqrt();
let biggest = v.iter().fold(0.0f32, |m, x| m.max(x.abs()));
assert!(biggest < even * 5.0, "{dim}: {biggest} against {even}");
let alive = v.iter().filter(|x| x.abs() > even / 4.0).count();
assert!(alive > dim * 3 / 4, "{dim}: only {alive} coordinates moved");
}
}
#[test]
fn an_odd_dimension_still_rotates() {
let r = Rotation::new(7, 1);
let mut v = [1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0];
let before = dot(&v, &v).sqrt();
r.apply(&mut v);
let after = dot(&v, &v).sqrt();
assert!((before - after).abs() < 1e-3);
}
#[test]
fn one_dimension_is_only_a_sign() {
let r = Rotation::new(1, 1);
let mut v = [3.0f32];
r.apply(&mut v);
assert_eq!(v[0].abs(), 3.0);
}
}