use kopitiam_core::Result;
use kopitiam_tensor::Tensor;
use crate::attention::attention_forward;
use crate::kv_cache::KvCache;
use crate::mlp::swiglu_mlp;
use crate::rope::RotaryEmbedding;
use crate::weights::LayerWeights;
#[allow(clippy::too_many_arguments)] pub(crate) fn block_forward(
x: &Tensor,
weights: &LayerWeights,
rope: &RotaryEmbedding,
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
norm_eps: f32,
layer_index: usize,
position_offset: usize,
cache: &mut KvCache,
) -> Result<Tensor> {
let normed = x.rms_norm(&weights.attn_norm, norm_eps)?;
let attn_out =
attention_forward(&normed, weights, rope, n_heads, n_kv_heads, head_dim, layer_index, position_offset, cache)?;
let x = x.add(&attn_out)?;
let normed = x.rms_norm(&weights.ffn_norm, norm_eps)?;
let mlp_out = swiglu_mlp(&normed, &weights.w_gate, &weights.w_up, &weights.w_down)?;
x.add(&mlp_out)
}
#[cfg(test)]
mod tests {
use super::*;
fn eye(n: usize) -> Tensor {
Tensor::from_f32((0..n * n).map(|i| if i / n == i % n { 1.0 } else { 0.0 }).collect(), [n, n]).unwrap()
}
#[test]
fn block_forward_produces_a_finite_correctly_shaped_residual_update() {
let hidden = 4;
let n_heads = 2;
let n_kv_heads = 1;
let head_dim = hidden / n_heads;
let weights = LayerWeights {
attn_norm: Tensor::from_f32(vec![1.0; hidden], [hidden]).unwrap(),
wq: eye(hidden),
bq: None,
wk: Tensor::from_f32(vec![0.1; n_kv_heads * head_dim * hidden], [n_kv_heads * head_dim, hidden]).unwrap(),
bk: None,
wv: Tensor::from_f32(vec![0.1; n_kv_heads * head_dim * hidden], [n_kv_heads * head_dim, hidden]).unwrap(),
bv: None,
wo: eye(hidden),
ffn_norm: Tensor::from_f32(vec![1.0; hidden], [hidden]).unwrap(),
w_gate: eye(hidden),
w_up: eye(hidden),
w_down: eye(hidden),
};
let rope = RotaryEmbedding::new(head_dim, 10_000.0, 16);
let mut cache = KvCache::new(1, 16);
let x = Tensor::from_f32(vec![0.5, -0.3, 0.8, 0.1, 0.2, 0.4, -0.6, 0.9], [2, hidden]).unwrap();
let out = block_forward(&x, &weights, &rope, n_heads, n_kv_heads, head_dim, 1e-6, 0, 0, &mut cache).unwrap();
assert_eq!(out.shape().dims(), &[2, hidden]);
for v in out.to_vec_f32().unwrap() {
assert!(v.is_finite());
}
assert_ne!(out.to_vec_f32().unwrap(), x.to_vec_f32().unwrap());
}
}