pub struct LayerCache {
pub residual_pre_attn: Vec<f32>,
pub normed_pre_attn: Vec<f32>,
pub inv_rms_pre_attn: f32,
pub residual_pre_ffn: Vec<f32>,
pub normed_pre_ffn: Vec<f32>,
pub inv_rms_pre_ffn: f32,
pub gate_pre: Vec<f32>,
pub up_pre: Vec<f32>,
pub attn_out: Vec<f32>,
pub ffn_out: Vec<f32>,
}
pub struct SequenceLayerCache {
pub tokens: Vec<LayerCache>,
pub q_pre_rope: Vec<f32>,
pub k_pre_rope: Vec<f32>,
pub v: Vec<f32>,
pub softmax_probs: Vec<Vec<f32>>,
pub context: Vec<f32>,
pub h_q: Vec<f32>,
pub h_v: Vec<f32>,
pub q_after_rope: Vec<Vec<f32>>,
}
pub struct BackwardTape {
pub layer23_input: Vec<f32>,
pub layer23: SequenceLayerCache,
pub final_normed: Vec<f32>,
pub inv_rms_final: f32,
pub logits: Vec<f32>,
}
pub fn rms_norm_forward(x: &[f32], w: &[f32], eps: f32) -> (Vec<f32>, f32) {
let d = x.len();
let mean_sq: f32 = x.iter().map(|xi| xi * xi).sum::<f32>() / d as f32;
let inv_rms = 1.0 / (mean_sq + eps).sqrt();
let normed: Vec<f32> = x
.iter()
.zip(w.iter())
.map(|(xi, wi)| xi * wi * inv_rms)
.collect();
(normed, inv_rms)
}
pub fn swiglu_forward(
x: &[f32],
w_gate: &[f32],
w_up: &[f32],
w_down: &[f32],
hidden: usize,
inter: usize,
) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let mut gate_pre = vec![0.0f32; inter];
let mut up_pre = vec![0.0f32; inter];
for i in 0..inter {
gate_pre[i] = w_gate[i * hidden..(i + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, b)| a * b)
.sum();
up_pre[i] = w_up[i * hidden..(i + 1) * hidden]
.iter()
.zip(x.iter())
.map(|(a, b)| a * b)
.sum();
}
let mixed: Vec<f32> = gate_pre
.iter()
.zip(up_pre.iter())
.map(|(&g, &u)| {
let s = 1.0 / (1.0 + (-g).exp());
g * s * u
})
.collect();
let mut out = vec![0.0f32; hidden];
for i in 0..hidden {
out[i] = w_down[i * inter..(i + 1) * inter]
.iter()
.zip(mixed.iter())
.map(|(a, b)| a * b)
.sum();
}
(out, gate_pre, up_pre)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rms_norm_forward_roundtrip() {
let x = vec![1.0f32, 2.0, 3.0, 4.0];
let w = vec![1.0f32; 4];
let (normed, inv_rms) = rms_norm_forward(&x, &w, 1e-6);
let mean_sq: f32 = x.iter().map(|xi| xi * xi).sum::<f32>() / 4.0;
let expected_inv = 1.0 / (mean_sq + 1e-6f32).sqrt();
assert!((inv_rms - expected_inv).abs() < 1e-6, "inv_rms mismatch");
let expected_norm: Vec<f32> = x.iter().map(|xi| xi * expected_inv).collect();
for (a, b) in normed.iter().zip(expected_norm.iter()) {
assert!((a - b).abs() < 1e-5, "normed mismatch {a} vs {b}");
}
}
#[test]
fn swiglu_forward_smoke() {
let hidden = 2;
let inter = 3;
let x = vec![1.0f32, -0.5];
let w_gate = vec![1.0f32, 0.0, 0.0, 1.0, 1.0, 0.0];
let w_up = vec![0.5f32, 0.5, 0.5, 0.5, 0.5, 0.5];
let w_down = vec![1.0f32, 0.0, 0.0, 0.0, 1.0, 0.0];
let (out, gate_pre, up_pre) = swiglu_forward(&x, &w_gate, &w_up, &w_down, hidden, inter);
assert_eq!(out.len(), hidden);
assert_eq!(gate_pre.len(), inter);
assert_eq!(up_pre.len(), inter);
}
}