#![cfg(feature = "cpu")]
use rlx_ir::ops::vq::VqMetric;
use rlx_ir::{DType, Graph, NodeId, Op, Shape};
use rlx_runtime::{Device, Session};
use std::time::Instant;
fn const_f32(g: &mut Graph, xs: &[f32], dims: &[usize]) -> NodeId {
let mut b = Vec::with_capacity(xs.len() * 4);
for x in xs {
b.extend_from_slice(&x.to_le_bytes());
}
g.add_node(
Op::Constant { data: b },
vec![],
Shape::new(dims, DType::F32),
)
}
fn bytes_to_f32s(b: &[u8]) -> Vec<f32> {
b.chunks_exact(4)
.map(|c| f32::from_le_bytes(c.try_into().unwrap()))
.collect()
}
fn vq_assign_l2_naive(x: &[f32], cb: &[f32], n: usize, d: usize, k: usize) -> Vec<f32> {
let mut out = vec![0f32; n];
for i in 0..n {
let xi = &x[i * d..(i + 1) * d];
let mut best = f32::INFINITY;
let mut best_j = 0usize;
for j in 0..k {
let cj = &cb[j * d..(j + 1) * d];
let mut dist = 0.0f32;
for t in 0..d {
let diff = xi[t] - cj[t];
dist += diff * diff;
}
if dist < best {
best = dist;
best_j = j;
}
}
out[i] = best_j as f32;
}
out
}
fn vq_assign_l2_fast(
x: &[f32],
cb: &[f32],
cb_norm: &[f32],
n: usize,
d: usize,
k: usize,
) -> Vec<f32> {
use rayon::prelude::*;
let mut out = vec![0f32; n];
out.par_iter_mut().enumerate().for_each(|(i, o)| {
let xi = &x[i * d..(i + 1) * d];
let mut best = f32::INFINITY;
let mut best_j = 0usize;
for j in 0..k {
let cj = &cb[j * d..(j + 1) * d];
let mut dot = 0.0f32;
for t in 0..d {
dot += xi[t] * cj[t];
}
let dist = cb_norm[j] - 2.0 * dot; if dist < best {
best = dist;
best_j = j;
}
}
*o = best_j as f32;
});
out
}
fn seeded(n: usize, salt: u32) -> Vec<f32> {
(0..n)
.map(|i| {
let z = (i as u32).wrapping_mul(2654435761).wrapping_add(salt);
((z >> 9) as f32 / (1u32 << 23) as f32) - 0.5
})
.collect()
}
#[test]
#[ignore = "hardware-dependent perf benchmark: the fused-vs-composition VQ \
winner depends on the platform BLAS. The N*K≥2M `speedup > 1.0` \
gate assumes a generic BLAS where the composition's [N,K] GEMM \
spills L2; Apple's Accelerate GEMM is fast enough that the \
composition wins there, so this hard-asserts falsely on Apple \
Silicon. Run with `--ignored --nocapture` (valid on the OpenBLAS \
rig). Correctness stays covered by fused_matches_composition_*."]
fn bench_vq_composition_vs_fused() {
for &(n, d, k) in &[
(256usize, 128usize, 1024usize),
(512, 128, 4096),
(256, 64, 8192),
] {
let x = seeded(n * d, 1);
let cb = seeded(k * d, 2);
let mut g = Graph::new("vq");
let xn = const_f32(&mut g, &x, &[n, d]);
let cbn = const_f32(&mut g, &cb, &[k, d]);
let (idx, _q) = g.vector_quantize(xn, cbn, VqMetric::L2);
g.set_outputs(vec![idx]);
let mut compiled = Session::new(Device::Cpu).compile(g);
let iters = 30;
let warmup = 5;
let comp_idx = bytes_to_f32s(&compiled.run_typed(&[])[0].0);
for _ in 0..warmup {
let _ = compiled.run_typed(&[]);
}
let mut comp_ns = u128::MAX;
for _ in 0..iters {
let t = Instant::now();
let _ = compiled.run_typed(&[]);
comp_ns = comp_ns.min(t.elapsed().as_nanos());
}
let cb_norm: Vec<f32> = (0..k)
.map(|j| cb[j * d..(j + 1) * d].iter().map(|&v| v * v).sum())
.collect();
let fused_idx = vq_assign_l2_fast(&x, &cb, &cb_norm, n, d, k);
for _ in 0..warmup {
let _ = vq_assign_l2_fast(&x, &cb, &cb_norm, n, d, k);
}
let mut fused_ns = u128::MAX;
for _ in 0..iters {
let t = Instant::now();
let _ = vq_assign_l2_fast(&x, &cb, &cb_norm, n, d, k);
fused_ns = fused_ns.min(t.elapsed().as_nanos());
}
let mut naive_ns = u128::MAX;
for _ in 0..iters {
let t = Instant::now();
let _ = vq_assign_l2_naive(&x, &cb, n, d, k);
naive_ns = naive_ns.min(t.elapsed().as_nanos());
}
let mism = comp_idx
.iter()
.zip(fused_idx.iter())
.filter(|(a, b)| a != b)
.count();
let speedup = comp_ns as f64 / fused_ns as f64;
println!(
"N={n:>4} D={d:>3} K={k:>5} composition={:>9}ns fused(fast)={:>9}ns naive={:>9}ns speedup={speedup:>5.2}x divergent={mism}/{n}",
comp_ns, fused_ns, naive_ns,
);
assert!(
fused_ns < naive_ns,
"fused ({fused_ns}ns) should beat naive ({naive_ns}ns)"
);
if n * k >= 2_000_000 {
assert!(
speedup > 1.0,
"fused should beat the composition at N={n} D={d} K={k} \
(N*K={}, got {speedup:.2}x)",
n * k
);
}
}
}
#[test]
fn fused_matches_composition_on_separable_data() {
let (n, d, k) = (64usize, 16usize, 64usize);
let mut cb = vec![0f32; k * d];
for j in 0..k {
cb[j * d + (j % d)] = 10.0 + j as f32;
}
let mut x = vec![0f32; n * d];
for i in 0..n {
let j = i % k;
x[i * d + (j % d)] = 10.0 + j as f32 + 0.01;
}
let cb_norm: Vec<f32> = (0..k)
.map(|j| cb[j * d..(j + 1) * d].iter().map(|&v| v * v).sum())
.collect();
let mut g = Graph::new("vq_ok");
let xn = const_f32(&mut g, &x, &[n, d]);
let cbn = const_f32(&mut g, &cb, &[k, d]);
let (idx, _q) = g.vector_quantize(xn, cbn, VqMetric::L2);
g.set_outputs(vec![idx]);
let comp = bytes_to_f32s(&Session::new(Device::Cpu).compile(g).run_typed(&[])[0].0);
let fused = vq_assign_l2_fast(&x, &cb, &cb_norm, n, d, k);
assert_eq!(
comp, fused,
"fused must match composition on separable data"
);
for i in 0..n {
assert_eq!(fused[i] as usize, i % k, "nearest code should be i%k");
}
}