use kopitiam_core::Result;
use kopitiam_tensor::Tensor;
use crate::kv_cache::KvCache;
use crate::linear::linear;
use crate::rope::RotaryEmbedding;
use crate::weights::LayerWeights;
pub(crate) fn repeat_kv_heads(x: &Tensor, group_size: usize) -> Result<Tensor> {
if group_size == 1 {
return Ok(x.clone());
}
let dims = x.shape().dims();
let (n_kv, seq, head_dim) = (dims[0], dims[1], dims[2]);
x.reshape([n_kv, 1, seq, head_dim])?
.broadcast_to([n_kv, group_size, seq, head_dim])?
.reshape([n_kv * group_size, seq, head_dim])
}
pub(crate) fn causal_mask(seq_q: usize, seq_kv: usize, position_offset: usize) -> Tensor {
let mut data = vec![0f32; seq_q * seq_kv];
for i in 0..seq_q {
let last_visible = position_offset + i;
for (j, cell) in data[i * seq_kv..(i + 1) * seq_kv].iter_mut().enumerate() {
if j > last_visible {
*cell = f32::NEG_INFINITY;
}
}
}
Tensor::from_f32(data, [1, seq_q, seq_kv]).expect("mask data length matches its own shape by construction")
}
#[allow(clippy::too_many_arguments)] pub(crate) fn attention_forward(
x: &Tensor,
weights: &LayerWeights,
rope: &RotaryEmbedding,
n_heads: usize,
n_kv_heads: usize,
head_dim: usize,
layer_index: usize,
position_offset: usize,
cache: &mut KvCache,
) -> Result<Tensor> {
let seq = x.shape().dims()[0];
let positions: Vec<usize> = (position_offset..position_offset + seq).collect();
let q = linear(x, &weights.wq, weights.bq.as_ref())?.reshape([seq, n_heads, head_dim])?.transpose(0, 1)?;
let k = linear(x, &weights.wk, weights.bk.as_ref())?.reshape([seq, n_kv_heads, head_dim])?.transpose(0, 1)?;
let v = linear(x, &weights.wv, weights.bv.as_ref())?.reshape([seq, n_kv_heads, head_dim])?.transpose(0, 1)?;
let q = rope.apply(&q, &positions)?;
let k = rope.apply(&k, &positions)?;
let (k_full, v_full) = cache.append(layer_index, k, v)?;
let seq_kv = k_full.shape().dims()[1];
let group_size = n_heads / n_kv_heads;
let k_rep = repeat_kv_heads(&k_full, group_size)?;
let v_rep = repeat_kv_heads(&v_full, group_size)?;
let scale = Tensor::from_f32(vec![(head_dim as f32).sqrt()], []).expect("scalar tensor");
let scores = q.matmul(&k_rep.transpose(1, 2)?)?.div(&scale)?;
let mask = causal_mask(seq, seq_kv, position_offset);
let scores = scores.add(&mask)?;
let probs = scores.softmax(2)?;
let attn_out = probs.matmul(&v_rep)?; let attn_out = attn_out.transpose(0, 1)?.reshape([seq, n_heads * head_dim])?;
linear(&attn_out, &weights.wo, None)
}
#[allow(dead_code)]
fn debug_assert_divides(n_heads: usize, n_kv_heads: usize) {
debug_assert!(n_kv_heads > 0 && n_heads.is_multiple_of(n_kv_heads));
}
#[cfg(test)]
mod tests {
use super::*;
use kopitiam_core::DType;
#[test]
fn repeat_kv_heads_makes_grouped_query_heads_genuinely_share_kv() {
let kv = Tensor::from_f32(vec![1.0, 2.0], [2, 1, 1]).unwrap();
let expanded = repeat_kv_heads(&kv, 2).unwrap();
assert_eq!(expanded.shape().dims(), &[4, 1, 1]);
let data = expanded.to_vec_f32().unwrap();
assert_eq!(data[0], 1.0);
assert_eq!(data[1], 1.0);
assert_eq!(data[2], 2.0);
assert_eq!(data[3], 2.0);
assert_eq!(data[0], data[1]);
assert_eq!(data[2], data[3]);
assert_ne!(data[0], data[2]);
}
#[test]
fn repeat_kv_heads_group_size_one_is_the_identity() {
let kv = Tensor::from_f32(vec![1.0, 2.0, 3.0, 4.0], [2, 2, 1]).unwrap();
let same = repeat_kv_heads(&kv, 1).unwrap();
assert_eq!(same.to_vec_f32().unwrap(), kv.to_vec_f32().unwrap());
}
#[test]
fn causal_mask_blocks_only_strictly_future_positions() {
let mask = causal_mask(3, 3, 0).to_vec_f32().unwrap();
assert_eq!(mask[0], 0.0);
assert!(mask[1].is_infinite() && mask[1] < 0.0);
assert!(mask[2].is_infinite() && mask[2] < 0.0);
assert_eq!(mask[3], 0.0);
assert_eq!(mask[4], 0.0);
assert!(mask[5].is_infinite() && mask[5] < 0.0);
assert_eq!(mask[6], 0.0);
assert_eq!(mask[7], 0.0);
assert_eq!(mask[8], 0.0);
}
#[test]
fn causal_mask_with_a_position_offset_accounts_for_already_cached_positions() {
let mask = causal_mask(1, 4, 3).to_vec_f32().unwrap();
assert_eq!(mask, vec![0.0, 0.0, 0.0, 0.0]);
}
#[test]
fn masked_softmax_assigns_exactly_zero_weight_to_future_positions() {
let scores = Tensor::from_f32(vec![5.0, 5.0, 5.0], [1, 1, 3]).unwrap();
let mask = causal_mask(1, 3, 0); let masked = scores.add(&mask).unwrap();
let probs = masked.softmax(2).unwrap().to_vec_f32().unwrap();
assert!((probs[0] - 1.0).abs() < 1e-6);
assert_eq!(probs[1], 0.0);
assert_eq!(probs[2], 0.0);
}
#[test]
fn masked_logits_are_exactly_negative_infinity() {
let mask = causal_mask(2, 2, 0).to_vec_f32().unwrap();
assert_eq!(mask[1], f32::NEG_INFINITY);
}
fn dummy_weights(hidden: usize, kv_dim: usize) -> LayerWeights {
let eye = |n: usize| Tensor::from_f32((0..n * n).map(|i| if i / n == i % n { 1.0 } else { 0.0 }).collect(), [n, n]).unwrap();
LayerWeights {
attn_norm: Tensor::from_f32(vec![1.0; hidden], [hidden]).unwrap(),
wq: eye(hidden),
bq: None,
wk: Tensor::from_f32(vec![0.0; kv_dim * hidden], [kv_dim, hidden]).unwrap(),
bk: None,
wv: Tensor::from_f32(vec![0.0; kv_dim * hidden], [kv_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),
}
}
#[test]
fn attention_forward_with_zeroed_kv_weights_produces_zero_output() {
let hidden = 4;
let n_heads = 2;
let n_kv_heads = 1;
let head_dim = hidden / n_heads;
let weights = dummy_weights(hidden, n_kv_heads * head_dim);
let rope = RotaryEmbedding::new(head_dim, 10_000.0, 16);
let mut cache = KvCache::new(1, 16);
let x = Tensor::from_f32((0..2 * hidden).map(|i| i as f32 * 0.1).collect(), [2, hidden]).unwrap();
let out = attention_forward(&x, &weights, &rope, n_heads, n_kv_heads, head_dim, 0, 0, &mut cache).unwrap();
assert_eq!(out.shape().dims(), &[2, hidden]);
for v in out.to_vec_f32().unwrap() {
assert_eq!(v, 0.0);
}
assert_eq!(cache.len(), 2);
}
#[test]
fn attention_forward_rejects_non_f32_input_the_same_way_matmul_does() {
let hidden = 32;
let weights = dummy_weights(hidden, hidden);
let rope = RotaryEmbedding::new(hidden, 10_000.0, 16);
let mut cache = KvCache::new(1, 16);
let x = Tensor::from_quantized(DType::Q4_0, vec![0u8; 18], [1, hidden]).unwrap();
assert!(attention_forward(&x, &weights, &rope, 1, 1, hidden, 0, 0, &mut cache).is_err());
}
}