use kopitiam_core::Result;
use kopitiam_tensor::Tensor;
use crate::linear::linear;
pub(crate) fn swiglu_mlp(x: &Tensor, gate_weight: &Tensor, up_weight: &Tensor, down_weight: &Tensor) -> Result<Tensor> {
let gate = linear(x, gate_weight, None)?;
let up = linear(x, up_weight, None)?;
let gated = gate.silu()?.mul(&up)?;
linear(&gated, down_weight, None)
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_close(a: f32, b: f32) {
assert!((a - b).abs() < 1e-5, "expected {b}, got {a}");
}
#[test]
fn swiglu_matches_hand_computation() {
let x = Tensor::from_f32(vec![1.0, 2.0], [1, 2]).unwrap();
let gate_w = Tensor::from_f32(vec![2.0, 0.0, 0.0, 3.0], [2, 2]).unwrap();
let up_w = Tensor::from_f32(vec![1.0, 1.0, 1.0, -1.0], [2, 2]).unwrap();
let down_w = Tensor::from_f32(vec![1.0, 0.0, 0.0, 1.0], [2, 2]).unwrap();
let out = swiglu_mlp(&x, &gate_w, &up_w, &down_w).unwrap().to_vec_f32().unwrap();
assert_close(out[0], 5.284_782_5);
assert_close(out[1], -5.985_164);
}
#[test]
fn a_zero_gate_projection_zeroes_the_whole_block() {
let x = Tensor::from_f32(vec![1.0, 2.0], [1, 2]).unwrap();
let zero_gate = Tensor::from_f32(vec![0.0; 4], [2, 2]).unwrap();
let up_w = Tensor::from_f32(vec![1.0, 1.0, 1.0, -1.0], [2, 2]).unwrap();
let down_w = Tensor::from_f32(vec![1.0, 0.0, 0.0, 1.0], [2, 2]).unwrap();
let out = swiglu_mlp(&x, &zero_gate, &up_w, &down_w).unwrap().to_vec_f32().unwrap();
assert_close(out[0], 0.0);
assert_close(out[1], 0.0);
}
}