use hanzo_ml::{DType, Device, Result, Tensor, D};
#[derive(Clone, Copy, Debug)]
pub struct SigReg {
pub t_max: f64,
pub n_knots: usize,
}
impl Default for SigReg {
fn default() -> Self {
Self {
t_max: 3.0,
n_knots: 17,
}
}
}
struct Quadrature {
t: Tensor,
phi: Tensor,
weights: Tensor,
}
impl SigReg {
fn quadrature(&self, dtype: DType, device: &Device) -> Result<Quadrature> {
let k = self.n_knots;
if k < 2 {
hanzo_ml::bail!("SIGReg needs at least 2 quadrature knots, got {k}");
}
let dt = self.t_max / (k - 1) as f64;
let mut t = Vec::with_capacity(k);
let mut phi = Vec::with_capacity(k);
let mut weights = Vec::with_capacity(k);
for i in 0..k {
let ti = i as f64 * dt;
let wi = if i == 0 || i == k - 1 { dt } else { 2.0 * dt };
let pi = (-0.5 * ti * ti).exp();
t.push(ti);
phi.push(pi);
weights.push(wi * pi);
}
Ok(Quadrature {
t: Tensor::from_vec(t, k, device)?.to_dtype(dtype)?,
phi: Tensor::from_vec(phi, k, device)?.to_dtype(dtype)?,
weights: Tensor::from_vec(weights, k, device)?.to_dtype(dtype)?,
})
}
pub fn statistic(&self, proj: &Tensor) -> Result<Tensor> {
let (n, _s) = proj.dims2()?;
let q = self.quadrature(proj.dtype(), proj.device())?;
let x_t = proj.unsqueeze(D::Minus1)?.broadcast_mul(&q.t)?;
let cos_mean = x_t.cos()?.mean(0)?;
let sin_mean = x_t.sin()?.mean(0)?;
let real = cos_mean.broadcast_sub(&q.phi)?.sqr()?;
let imag = sin_mean.sqr()?;
let err = (real + imag)?;
let stat = err.broadcast_mul(&q.weights)?.sum(D::Minus1)?;
stat * n as f64
}
pub fn loss_with_directions(&self, embeddings: &Tensor, directions: &Tensor) -> Result<Tensor> {
let (_n, d) = embeddings.dims2()?;
let (dd, _s) = directions.dims2()?;
if d != dd {
hanzo_ml::bail!("embeddings dim ({d}) and directions dim ({dd}) mismatch");
}
let proj = embeddings.matmul(directions)?;
self.statistic(&proj)?.mean_all()
}
pub fn loss(&self, embeddings: &Tensor, n_slices: usize) -> Result<Tensor> {
let (_n, d) = embeddings.dims2()?;
let directions = random_directions(d, n_slices, embeddings.dtype(), embeddings.device())?;
self.loss_with_directions(embeddings, &directions)
}
}
pub fn random_directions(
d: usize,
n_slices: usize,
dtype: DType,
device: &Device,
) -> Result<Tensor> {
let a = Tensor::randn(0f32, 1f32, (d, n_slices), device)?.to_dtype(dtype)?;
let norm = a.sqr()?.sum_keepdim(0)?.sqrt()?;
a.broadcast_div(&norm)
}
#[cfg(test)]
mod tests {
use super::*;
use hanzo_ml::Device;
include!("sigreg_fixture.rs");
fn max_abs_diff(a: &Tensor, b: &[f64]) -> f64 {
let a = a.to_dtype(DType::F64).unwrap().flatten_all().unwrap();
let a = a.to_vec1::<f64>().unwrap();
a.iter()
.zip(b)
.map(|(x, y)| (x - y).abs())
.fold(0.0, f64::max)
}
#[test]
fn quadrature_matches_official() {
let dev = Device::Cpu;
let cfg = SigReg {
t_max: T_MAX,
n_knots: N_KNOTS,
};
let q = cfg.quadrature(DType::F64, &dev).unwrap();
assert!(max_abs_diff(&q.t, REF_T) < 1e-12, "t buffer mismatch");
assert!(max_abs_diff(&q.phi, REF_PHI) < 1e-12, "phi buffer mismatch");
assert!(
max_abs_diff(&q.weights, REF_WEIGHTS) < 1e-12,
"weights buffer mismatch"
);
}
#[test]
fn statistic_matches_official_f64() {
let dev = Device::Cpu;
let cfg = SigReg {
t_max: T_MAX,
n_knots: N_KNOTS,
};
let proj = Tensor::from_slice(PROJ, (N, S), &dev).unwrap();
let stat = cfg.statistic(&proj).unwrap();
let d = max_abs_diff(&stat, REF_PER_SLICE);
assert!(d < 1e-5, "per-slice statistic diff {d:e} exceeds 1e-5");
let x = Tensor::from_slice(X, (N, D), &dev).unwrap();
let a = Tensor::from_slice(A, (D, S), &dev).unwrap();
let loss = cfg.loss_with_directions(&x, &a).unwrap();
let dl = (loss.to_scalar::<f64>().unwrap() - REF_SIGREG).abs();
assert!(dl < 1e-5, "sigreg loss diff {dl:e} exceeds 1e-5");
}
#[test]
fn statistic_matches_official_f32() {
let dev = Device::Cpu;
let cfg = SigReg {
t_max: T_MAX,
n_knots: N_KNOTS,
};
let proj = Tensor::from_slice(PROJ, (N, S), &dev)
.unwrap()
.to_dtype(DType::F32)
.unwrap();
let stat = cfg.statistic(&proj).unwrap();
let d = max_abs_diff(&stat, REF_PER_SLICE_F32);
assert!(d < 1e-3, "f32 per-slice statistic diff {d:e} exceeds 1e-3");
let x = Tensor::from_slice(X, (N, D), &dev)
.unwrap()
.to_dtype(DType::F32)
.unwrap();
let a = Tensor::from_slice(A, (D, S), &dev)
.unwrap()
.to_dtype(DType::F32)
.unwrap();
let loss = cfg.loss_with_directions(&x, &a).unwrap();
let dl = (loss.to_scalar::<f32>().unwrap() as f64 - REF_SIGREG_F32).abs();
assert!(dl < 1e-3, "f32 sigreg loss diff {dl:e} exceeds 1e-3");
}
#[test]
fn random_directions_are_unit_norm() {
let dev = Device::Cpu;
let dirs = random_directions(16, 64, DType::F32, &dev).unwrap();
let norms = dirs
.sqr()
.unwrap()
.sum_keepdim(0)
.unwrap()
.sqrt()
.unwrap()
.flatten_all()
.unwrap()
.to_vec1::<f32>()
.unwrap();
for nrm in norms {
assert!((nrm - 1.0).abs() < 1e-5, "column norm {nrm} != 1");
}
}
fn eigenvalue_spread(z: &Tensor) -> f64 {
let (n, d) = z.dims2().unwrap();
let mean = z.mean_keepdim(0).unwrap();
let zc = z.broadcast_sub(&mean).unwrap();
let cov = (zc.t().unwrap().matmul(&zc).unwrap() / n as f64).unwrap();
let fro_sq = cov
.sqr()
.unwrap()
.sum_all()
.unwrap()
.to_scalar::<f64>()
.unwrap();
let trace = (zc
.sqr()
.unwrap()
.sum_all()
.unwrap()
.to_scalar::<f64>()
.unwrap())
/ n as f64;
fro_sq - trace * trace / d as f64
}
#[test]
fn demo_training_drives_embeddings_toward_isotropy() {
use crate::optim::{AdamW, Optimizer, ParamsAdamW};
use hanzo_ml::Var;
use rand::{rngs::StdRng, SeedableRng};
use rand_distr::{Distribution, StandardNormal};
let dev = Device::Cpu;
let mut rng = StdRng::seed_from_u64(42);
let mut randn = |rows: usize, cols: usize| {
let v: Vec<f32> = (0..rows * cols)
.map(|_| StandardNormal.sample(&mut rng))
.collect();
Tensor::from_vec(v, (rows, cols), &dev).unwrap()
};
let (n, in_dim, hidden, embed) = (64usize, 6usize, 16usize, 8usize);
let sigreg = SigReg::default();
let lambda = 0.7;
let scale = Tensor::from_vec(
(0..in_dim)
.map(|i| 0.3 + 0.5 * i as f32)
.collect::<Vec<_>>(),
(1, in_dim),
&dev,
)
.unwrap();
let x_base = randn(n, in_dim).broadcast_mul(&scale).unwrap();
let noise = (randn(n, in_dim) * 0.1).unwrap();
let v1 = (&x_base + &noise).unwrap();
let v2 = (&x_base - &noise).unwrap();
let w1 =
Var::from_tensor(&(randn(in_dim, hidden) * (1.0 / (in_dim as f64).sqrt())).unwrap())
.unwrap();
let b1 = Var::from_tensor(&Tensor::zeros((1, hidden), DType::F32, &dev).unwrap()).unwrap();
let w2 =
Var::from_tensor(&(randn(hidden, embed) * (1.0 / (hidden as f64).sqrt())).unwrap())
.unwrap();
let b2 = Var::from_tensor(&Tensor::zeros((1, embed), DType::F32, &dev).unwrap()).unwrap();
let params = vec![w1.clone(), b1.clone(), w2.clone(), b2.clone()];
let dirs = {
let a = randn(embed, 48);
let norm = a.sqr().unwrap().sum_keepdim(0).unwrap().sqrt().unwrap();
a.broadcast_div(&norm).unwrap()
};
let forward = |x: &Tensor| -> Tensor {
x.matmul(w1.as_tensor())
.unwrap()
.broadcast_add(b1.as_tensor())
.unwrap()
.relu()
.unwrap()
.matmul(w2.as_tensor())
.unwrap()
.broadcast_add(b2.as_tensor())
.unwrap()
};
let embeddings = || {
let e1 = forward(&v1);
let e2 = forward(&v2);
Tensor::cat(&[&e1, &e2], 0)
.unwrap()
.to_dtype(DType::F64)
.unwrap()
};
let step_loss = || -> (Tensor, f64) {
let e1 = forward(&v1);
let e2 = forward(&v2);
let inv = (&e1 - &e2).unwrap().sqr().unwrap().mean_all().unwrap();
let emb = Tensor::cat(&[&e1, &e2], 0).unwrap();
let sig = sigreg.loss_with_directions(&emb, &dirs).unwrap();
let loss = ((&sig * lambda).unwrap() + (inv * (1.0 - lambda)).unwrap()).unwrap();
let sig_v = sig.to_scalar::<f32>().unwrap() as f64;
(loss, sig_v)
};
let mut opt = AdamW::new(
params,
ParamsAdamW {
lr: 1e-2,
..Default::default()
},
)
.unwrap();
let (loss0_t, sig0) = step_loss();
let loss0 = loss0_t.to_scalar::<f32>().unwrap();
let spread0 = eigenvalue_spread(&embeddings());
let mut last = loss0;
for _ in 0..150 {
let (loss, _) = step_loss();
last = loss.to_scalar::<f32>().unwrap();
opt.backward_step(&loss).unwrap();
}
let (_, sig1) = step_loss();
let spread1 = eigenvalue_spread(&embeddings());
eprintln!(
"loss {loss0:.4} -> {last:.4} sigreg {sig0:.4} -> {sig1:.4} eig-spread {spread0:.4} -> {spread1:.4}"
);
assert!(
last < loss0,
"total loss did not decrease: {loss0} -> {last}"
);
assert!(
sig1 < sig0,
"sigreg statistic did not decrease: {sig0} -> {sig1}"
);
assert!(
spread1 < 0.7 * spread0,
"eigenvalue spread did not measurably shrink: {spread0} -> {spread1}"
);
}
}