|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| use meganeura::{Graph, NodeId};
|
|
|
|
|
| pub const COMPONENTS: usize = 3;
|
|
|
|
|
| #[derive(Debug, Clone)]
|
| pub struct Basis {
|
|
|
| pub mean: Vec<f32>,
|
|
|
| pub components: Vec<f32>,
|
|
|
| pub scale: [f32; COMPONENTS],
|
|
|
| pub offset: [f32; COMPONENTS],
|
| pub dim: usize,
|
| }
|
|
|
| impl Basis {
|
|
|
|
|
|
|
| pub fn placeholder(dim: usize) -> Self {
|
| let mut components = vec![0.0; COMPONENTS * dim];
|
| for c in 0..COMPONENTS {
|
| components[c * dim + c] = 1.0;
|
| }
|
| Self {
|
| mean: vec![0.0; dim],
|
| components,
|
| scale: [1.0; COMPONENTS],
|
| offset: [0.5; COMPONENTS],
|
| dim,
|
| }
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| pub fn fit(features: &[f32], tokens: usize, dim: usize, skip: usize) -> Self {
|
| assert_eq!(features.len(), tokens * dim, "feature matrix shape mismatch");
|
| assert!(skip < tokens, "nothing left after skipping {skip} tokens");
|
| let rows = tokens - skip;
|
| let data = &features[skip * dim..];
|
|
|
|
|
| let mut mean = vec![0.0f32; dim];
|
| for r in 0..rows {
|
| for d in 0..dim {
|
| mean[d] += data[r * dim + d];
|
| }
|
| }
|
| for m in &mut mean {
|
| *m /= rows as f32;
|
| }
|
| let mut centred: Vec<f32> = (0..rows * dim)
|
| .map(|i| data[i] - mean[i % dim])
|
| .collect();
|
|
|
| let mut components = vec![0.0f32; COMPONENTS * dim];
|
| let mut projected = vec![0.0f32; COMPONENTS * rows];
|
|
|
| for c in 0..COMPONENTS {
|
|
|
|
|
| let mut v: Vec<f32> = (0..dim)
|
| .map(|d| ((d * 2654435761usize) % 1024) as f32 / 1024.0 - 0.5)
|
| .collect();
|
| normalize(&mut v);
|
|
|
| let mut scores = vec![0.0f32; rows];
|
| for _ in 0..48 {
|
|
|
|
|
| for r in 0..rows {
|
| let row = ¢red[r * dim..(r + 1) * dim];
|
| scores[r] = row.iter().zip(&v).map(|(a, b)| a * b).sum();
|
| }
|
| let mut next = vec![0.0f32; dim];
|
| for r in 0..rows {
|
| let s = scores[r];
|
| let row = ¢red[r * dim..(r + 1) * dim];
|
| for d in 0..dim {
|
| next[d] += s * row[d];
|
| }
|
| }
|
| if normalize(&mut next) < 1e-12 {
|
| break;
|
| }
|
| v = next;
|
| }
|
|
|
|
|
|
|
| for r in 0..rows {
|
| let row = ¢red[r * dim..(r + 1) * dim];
|
| let s: f32 = row.iter().zip(&v).map(|(a, b)| a * b).sum();
|
| scores[r] = s;
|
| projected[c * rows + r] = s;
|
| }
|
| for r in 0..rows {
|
| let s = scores[r];
|
| for d in 0..dim {
|
| centred[r * dim + d] -= s * v[d];
|
| }
|
| }
|
| components[c * dim..(c + 1) * dim].copy_from_slice(&v);
|
| }
|
|
|
|
|
|
|
|
|
| let mut lo = [0.0f32; COMPONENTS];
|
| let mut span = [0.0f32; COMPONENTS];
|
| for c in 0..COMPONENTS {
|
| let mut col: Vec<f32> = projected[c * rows..(c + 1) * rows].to_vec();
|
| col.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
|
| lo[c] = col[(rows as f32 * 0.02) as usize];
|
| let hi = col[((rows as f32 * 0.98) as usize).min(rows - 1)];
|
| span[c] = hi - lo[c];
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
| let dominant = span.iter().copied().fold(0.0f32, f32::max);
|
| let floor = dominant * 1e-3;
|
| let mut scale = [0.0f32; COMPONENTS];
|
| let mut offset = [0.5f32; COMPONENTS];
|
| for c in 0..COMPONENTS {
|
| if span[c] > floor && span[c] > f32::MIN_POSITIVE {
|
| scale[c] = 1.0 / span[c];
|
| offset[c] = -lo[c] / span[c];
|
| }
|
| }
|
|
|
| Self {
|
| mean,
|
| components,
|
| scale,
|
| offset,
|
| dim,
|
| }
|
| }
|
|
|
|
|
|
|
| pub fn weight_matrix(&self) -> Vec<f32> {
|
| let mut w = vec![0.0f32; self.dim * COMPONENTS];
|
| for c in 0..COMPONENTS {
|
| for d in 0..self.dim {
|
| w[d * COMPONENTS + c] = self.components[c * self.dim + d] * self.scale[c];
|
| }
|
| }
|
| w
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
| pub fn bias_vector(&self) -> Vec<f32> {
|
| let mut b = [0.0f32; COMPONENTS];
|
| for c in 0..COMPONENTS {
|
| let dot: f32 = (0..self.dim)
|
| .map(|d| self.mean[d] * self.components[c * self.dim + d])
|
| .sum();
|
| b[c] = self.offset[c] - dot * self.scale[c];
|
| }
|
| b.to_vec()
|
| }
|
|
|
|
|
|
|
| pub fn project(&self, features: &[f32], tokens: usize) -> Vec<f32> {
|
| let w = self.weight_matrix();
|
| let b = self.bias_vector();
|
| let mut out = vec![0.0f32; tokens * COMPONENTS];
|
| for t in 0..tokens {
|
| for c in 0..COMPONENTS {
|
| let mut acc = b[c];
|
| for d in 0..self.dim {
|
| acc += features[t * self.dim + d] * w[d * COMPONENTS + c];
|
| }
|
| out[t * COMPONENTS + c] = acc;
|
| }
|
| }
|
| out
|
| }
|
| }
|
|
|
| fn normalize(v: &mut [f32]) -> f32 {
|
| let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
|
| if norm > 1e-12 {
|
| for x in v.iter_mut() {
|
| *x /= norm;
|
| }
|
| }
|
| norm
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| pub fn add_projection(g: &mut Graph, features: NodeId, hidden: usize) -> NodeId {
|
| let w = g.parameter("pca.weight", &[hidden, COMPONENTS]);
|
| let b = g.parameter("pca.bias", &[COMPONENTS]);
|
| let projected = g.matmul(features, w);
|
| g.bias_add(projected, b)
|
| }
|
|
|
| #[cfg(test)]
|
| mod tests {
|
| use super::*;
|
|
|
|
|
|
|
| #[test]
|
| fn recovers_a_planted_subspace() {
|
| let dim = 32;
|
| let tokens = 200;
|
| let mut features = vec![0.0f32; tokens * dim];
|
| for t in 0..tokens {
|
| let a = (t as f32 / tokens as f32) * 2.0 - 1.0;
|
| let b = ((t * 7 % tokens) as f32 / tokens as f32) * 2.0 - 1.0;
|
| for d in 0..dim {
|
|
|
| features[t * dim + d] = if d == 3 {
|
| 5.0 * a
|
| } else if d == 11 {
|
| 4.0 * b
|
| } else {
|
| 0.01 * ((d as f32) * 0.1 + a)
|
| };
|
| }
|
| }
|
|
|
| let basis = Basis::fit(&features, tokens, dim, 0);
|
|
|
| let c0 = &basis.components[0..dim];
|
| let c1 = &basis.components[dim..2 * dim];
|
| let strongest = |c: &[f32]| {
|
| c.iter()
|
| .enumerate()
|
| .max_by(|a, b| a.1.abs().partial_cmp(&b.1.abs()).unwrap())
|
| .unwrap()
|
| .0
|
| };
|
| let (s0, s1) = (strongest(c0), strongest(c1));
|
| assert!(
|
| (s0 == 3 && s1 == 11) || (s0 == 11 && s1 == 3),
|
| "expected components along dims 3 and 11, got {s0} and {s1}"
|
| );
|
| }
|
|
|
| #[test]
|
| fn components_are_orthonormal() {
|
| let dim = 24;
|
| let tokens = 90;
|
| let features: Vec<f32> = (0..tokens * dim)
|
| .map(|i| ((i * 37 % 101) as f32 / 101.0 - 0.5) * (1.0 + (i % 5) as f32))
|
| .collect();
|
| let basis = Basis::fit(&features, tokens, dim, 0);
|
|
|
| for a in 0..COMPONENTS {
|
| let va = &basis.components[a * dim..(a + 1) * dim];
|
| let norm: f32 = va.iter().map(|x| x * x).sum::<f32>().sqrt();
|
| assert!((norm - 1.0).abs() < 1e-3, "component {a} norm {norm}");
|
| for b in (a + 1)..COMPONENTS {
|
| let vb = &basis.components[b * dim..(b + 1) * dim];
|
| let dot: f32 = va.iter().zip(vb).map(|(x, y)| x * y).sum();
|
| assert!(dot.abs() < 1e-2, "components {a},{b} not orthogonal: {dot}");
|
| }
|
| }
|
| }
|
|
|
|
|
|
|
| #[test]
|
| fn folded_weights_match_explicit_form() {
|
| let dim = 16;
|
| let tokens = 60;
|
| let features: Vec<f32> = (0..tokens * dim)
|
| .map(|i| ((i * 13 % 97) as f32 / 97.0 - 0.5) * 3.0)
|
| .collect();
|
| let basis = Basis::fit(&features, tokens, dim, 0);
|
| let folded = basis.project(&features, tokens);
|
|
|
| for t in 0..tokens {
|
| for c in 0..COMPONENTS {
|
| let explicit: f32 = (0..dim)
|
| .map(|d| {
|
| (features[t * dim + d] - basis.mean[d]) * basis.components[c * dim + d]
|
| })
|
| .sum::<f32>()
|
| * basis.scale[c]
|
| + basis.offset[c];
|
| let got = folded[t * COMPONENTS + c];
|
| assert!(
|
| (got - explicit).abs() < 1e-3,
|
| "token {t} component {c}: folded {got} vs explicit {explicit}"
|
| );
|
| }
|
| }
|
| }
|
|
|
|
|
|
|
| #[test]
|
| fn projection_spans_the_display_range() {
|
| let dim = 48;
|
| let tokens = 300;
|
| let features: Vec<f32> = (0..tokens * dim)
|
| .map(|i| {
|
| let t = (i / dim) as f32;
|
| let d = (i % dim) as f32;
|
| (t * 0.017 + d * 0.31).sin() * 2.0
|
| })
|
| .collect();
|
| let basis = Basis::fit(&features, tokens, dim, 0);
|
| let out = basis.project(&features, tokens);
|
|
|
| for c in 0..COMPONENTS {
|
| let vals: Vec<f32> = (0..tokens).map(|t| out[t * COMPONENTS + c]).collect();
|
| let inside = vals.iter().filter(|v| (0.0..=1.0).contains(*v)).count();
|
| let frac = inside as f32 / tokens as f32;
|
| assert!(frac > 0.9, "component {c}: only {:.0}% inside [0,1]", frac * 100.0);
|
| }
|
| }
|
|
|
|
|
|
|
| #[test]
|
| fn degenerate_components_stay_neutral() {
|
| let dim = 32;
|
| let tokens = 120;
|
|
|
| let mut features = vec![0.0f32; tokens * dim];
|
| for t in 0..tokens {
|
| let a = t as f32 / tokens as f32;
|
| for d in 0..dim {
|
| features[t * dim + d] = a * (d as f32 * 0.05).cos();
|
| }
|
| }
|
|
|
| let basis = Basis::fit(&features, tokens, dim, 0);
|
| let out = basis.project(&features, tokens);
|
| assert!(out.iter().all(|v| v.is_finite()), "projection produced non-finite values");
|
|
|
|
|
|
|
| for c in 1..COMPONENTS {
|
| let vals: Vec<f32> = (0..tokens).map(|t| out[t * COMPONENTS + c]).collect();
|
| let lo = vals.iter().copied().fold(f32::INFINITY, f32::min);
|
| let hi = vals.iter().copied().fold(f32::NEG_INFINITY, f32::max);
|
| assert!(
|
| (hi - lo) < 1e-3 && (lo - 0.5).abs() < 1e-3,
|
| "component {c} should be flat mid-grey, spans [{lo}, {hi}]"
|
| );
|
| }
|
| }
|
|
|
| #[test]
|
| fn placeholder_is_usable_before_the_first_fit() {
|
| let basis = Basis::placeholder(384);
|
| let features = vec![0.25f32; 8 * 384];
|
| let out = basis.project(&features, 8);
|
| assert_eq!(out.len(), 8 * COMPONENTS);
|
| assert!(out.iter().all(|v| v.is_finite()));
|
| }
|
| }
|
|
|