use crate::forward::cpu::{add_bias, matmul_bt, softmax_attention};
use crate::lora_hook::{LoraHook, apply_lora_rows};
use crate::weights::TransformerLayerWeights;
#[derive(Debug, Clone)]
pub struct AttentionBuffers {
pub q: Vec<f32>,
pub k: Vec<f32>,
pub v: Vec<f32>,
pub scores: Vec<f32>,
pub context: Vec<f32>,
pub concat: Vec<f32>,
pub ffn_intermediate: Vec<f32>,
pub temp: Vec<f32>,
qkv: Vec<f32>,
q_head: Vec<f32>,
k_head: Vec<f32>,
v_all_t: Vec<f32>,
scores_head: Vec<f32>,
context_head: Vec<f32>,
}
impl AttentionBuffers {
pub fn new(
max_seq_len: usize,
hidden_size: usize,
num_heads: usize,
intermediate_size: usize,
) -> Self {
let head_dim = hidden_size / num_heads;
Self {
q: vec![0.0; max_seq_len * hidden_size],
k: vec![0.0; max_seq_len * hidden_size],
v: vec![0.0; max_seq_len * hidden_size],
scores: vec![0.0; num_heads * max_seq_len * max_seq_len],
context: vec![0.0; num_heads * max_seq_len * head_dim],
concat: vec![0.0; max_seq_len * hidden_size],
ffn_intermediate: vec![0.0; max_seq_len * intermediate_size],
temp: vec![0.0; max_seq_len * hidden_size],
qkv: vec![0.0; max_seq_len * 3 * hidden_size],
q_head: vec![0.0; max_seq_len * head_dim],
k_head: vec![0.0; max_seq_len * head_dim],
v_all_t: vec![0.0; hidden_size * max_seq_len],
scores_head: vec![0.0; max_seq_len * max_seq_len],
context_head: vec![0.0; max_seq_len * head_dim],
}
}
#[cfg(test)]
pub(crate) fn total_scratch_len(&self) -> usize {
self.q.len()
+ self.k.len()
+ self.v.len()
+ self.scores.len()
+ self.context.len()
+ self.concat.len()
+ self.ffn_intermediate.len()
+ self.temp.len()
+ self.qkv.len()
+ self.q_head.len()
+ self.k_head.len()
+ self.v_all_t.len()
+ self.scores_head.len()
+ self.context_head.len()
}
}
pub fn multi_head_attention(
hidden_states: &[f32],
layer_weights: &TransformerLayerWeights<'_>,
attention_mask: &[u32],
seq_len: usize,
hidden_size: usize,
num_heads: usize,
head_dim: usize,
buffers: &mut AttentionBuffers,
lora: &dyn LoraHook,
layer_idx: usize,
) -> Vec<f32> {
let fused_qkv_weight: Vec<f32> = layer_weights
.query_weight
.data
.iter()
.chain(layer_weights.key_weight.data.iter())
.chain(layer_weights.value_weight.data.iter())
.copied()
.collect();
let fused_qkv_bias: Vec<f32> = layer_weights
.query_bias
.data
.iter()
.chain(layer_weights.key_bias.data.iter())
.chain(layer_weights.value_bias.data.iter())
.copied()
.collect();
multi_head_attention_in_place(
hidden_states,
layer_weights,
&fused_qkv_weight,
&fused_qkv_bias,
attention_mask,
seq_len,
hidden_size,
num_heads,
head_dim,
buffers,
lora,
layer_idx,
);
buffers.temp[..seq_len * hidden_size].to_vec()
}
#[inline]
fn assert_standard_no_overflow(
seq_len: usize,
hidden_size: usize,
num_heads: usize,
head_dim: usize,
) {
assert!(num_heads > 0, "standard: num_heads must be non-zero");
assert!(head_dim > 0, "standard: head_dim must be non-zero");
assert!(
num_heads.checked_mul(head_dim).is_some(),
"standard shape overflow: num_heads * head_dim"
);
assert_eq!(
hidden_size,
num_heads * head_dim,
"standard: hidden_size must equal num_heads * head_dim"
);
assert!(
seq_len.checked_mul(hidden_size).is_some(),
"standard shape overflow: seq_len * hidden_size"
);
assert!(
num_heads.checked_mul(seq_len).is_some(),
"standard shape overflow: num_heads * seq_len"
);
let nh_sl = num_heads * seq_len;
assert!(
nh_sl.checked_mul(seq_len).is_some(),
"standard shape overflow: num_heads * seq_len * seq_len"
);
assert!(
nh_sl.checked_mul(head_dim).is_some(),
"standard shape overflow: num_heads * seq_len * head_dim"
);
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn multi_head_attention_in_place(
hidden_states: &[f32],
layer_weights: &TransformerLayerWeights<'_>,
fused_qkv_weight: &[f32],
fused_qkv_bias: &[f32],
attention_mask: &[u32],
seq_len: usize,
hidden_size: usize,
num_heads: usize,
head_dim: usize,
buffers: &mut AttentionBuffers,
lora: &dyn LoraHook,
layer_idx: usize,
) {
assert_standard_no_overflow(seq_len, hidden_size, num_heads, head_dim);
assert_eq!(
hidden_states.len(),
seq_len * hidden_size,
"standard: hidden_states length must equal seq_len * hidden_size"
);
assert_eq!(
attention_mask.len(),
seq_len,
"standard: attention_mask length must equal seq_len"
);
let used_hidden = seq_len * hidden_size;
let used_scores = num_heads * seq_len * seq_len;
{
let AttentionBuffers { qkv, q, k, v, .. } = &mut *buffers;
let qkv = &mut qkv[..seq_len * 3 * hidden_size];
matmul_bt(
hidden_states,
fused_qkv_weight,
qkv,
seq_len,
hidden_size,
3 * hidden_size,
);
add_bias(qkv, fused_qkv_bias, 3 * hidden_size);
for i in 0..seq_len {
let src = i * 3 * hidden_size;
q[i * hidden_size..(i + 1) * hidden_size].copy_from_slice(&qkv[src..src + hidden_size]);
k[i * hidden_size..(i + 1) * hidden_size]
.copy_from_slice(&qkv[src + hidden_size..src + 2 * hidden_size]);
v[i * hidden_size..(i + 1) * hidden_size]
.copy_from_slice(&qkv[src + 2 * hidden_size..src + 3 * hidden_size]);
}
}
apply_lora_rows(
lora,
layer_idx,
"query",
hidden_states,
&mut buffers.q[..used_hidden],
hidden_size,
hidden_size,
);
apply_lora_rows(
lora,
layer_idx,
"key",
hidden_states,
&mut buffers.k[..used_hidden],
hidden_size,
hidden_size,
);
apply_lora_rows(
lora,
layer_idx,
"value",
hidden_states,
&mut buffers.v[..used_hidden],
hidden_size,
hidden_size,
);
let scale = 1.0 / (head_dim as f32).sqrt();
{
let (q_buf, rest) = buffers.q.split_at(used_hidden);
let _ = rest;
for h in 0..num_heads {
let head_offset = h * head_dim;
for i in 0..seq_len {
let src_start = i * hidden_size + head_offset;
let dst_start = i * head_dim;
buffers.q_head[dst_start..dst_start + head_dim]
.copy_from_slice(&q_buf[src_start..src_start + head_dim]);
}
for i in 0..seq_len {
let src_start = i * hidden_size + head_offset;
let dst_start = i * head_dim;
buffers.k_head[dst_start..dst_start + head_dim]
.copy_from_slice(&buffers.k[src_start..src_start + head_dim]);
}
let q_head = &buffers.q_head[..seq_len * head_dim];
let k_head = &buffers.k_head[..seq_len * head_dim];
let scores_head = &mut buffers.scores_head[..seq_len * seq_len];
matmul_bt(q_head, k_head, scores_head, seq_len, head_dim, seq_len);
let scores_offset = h * seq_len * seq_len;
for (idx, &score) in scores_head.iter().enumerate() {
buffers.scores[scores_offset + idx] = score * scale;
}
}
}
{
let scores = &mut buffers.scores[..used_scores];
for h in 0..num_heads {
for i in 0..seq_len {
let row = &mut scores[(h * seq_len + i) * seq_len..(h * seq_len + i + 1) * seq_len];
for j in 0..seq_len {
if attention_mask[j] == 0 {
row[j] = f32::NEG_INFINITY;
}
}
}
}
softmax_attention(scores, seq_len, num_heads);
}
{
let AttentionBuffers { v, v_all_t, .. } = &mut *buffers;
let v_all_t = &mut v_all_t[..hidden_size * seq_len];
for i in 0..seq_len {
let row_start = i * hidden_size;
for d in 0..hidden_size {
v_all_t[d * seq_len + i] = v[row_start + d];
}
}
}
{
let AttentionBuffers {
scores,
v_all_t,
context_head,
concat,
..
} = &mut *buffers;
let concat = &mut concat[..used_hidden];
for h in 0..num_heads {
let head_offset = h * head_dim;
let scores_offset = h * seq_len * seq_len;
let scores_head = &scores[scores_offset..scores_offset + seq_len * seq_len];
let v_head_t = &v_all_t[head_offset * seq_len..(head_offset + head_dim) * seq_len];
let context_head = &mut context_head[..seq_len * head_dim];
matmul_bt(
scores_head,
v_head_t,
context_head,
seq_len,
seq_len,
head_dim,
);
for i in 0..seq_len {
let dst = i * hidden_size + head_offset;
concat[dst..dst + head_dim]
.copy_from_slice(&context_head[i * head_dim..(i + 1) * head_dim]);
}
}
}
{
let concat = &buffers.concat[..used_hidden];
let output = &mut buffers.temp[..used_hidden];
matmul_bt(
concat,
layer_weights.attn_output_weight.data,
output,
seq_len,
hidden_size,
hidden_size,
);
add_bias(output, layer_weights.attn_output_bias.data, hidden_size);
apply_lora_rows(
lora,
layer_idx,
"attn_output",
concat,
output,
hidden_size,
hidden_size,
);
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn multi_head_attention_batched(
hidden_states: &[f32],
layer_weights: &TransformerLayerWeights<'_>,
fused_qkv_weight: &[f32],
fused_qkv_bias: &[f32],
cu_seqlens: &[usize],
hidden_size: usize,
num_heads: usize,
head_dim: usize,
q: &mut [f32],
k: &mut [f32],
v: &mut [f32],
qkv: &mut [f32],
concat: &mut [f32],
output: &mut [f32],
lora: &dyn LoraHook,
layer_idx: usize,
) {
assert!(
cu_seqlens.len() >= 2,
"standard: cu_seqlens must have at least 2 entries (batch + 1)"
);
assert_eq!(cu_seqlens[0], 0, "standard: cu_seqlens must start at 0");
let batch = cu_seqlens.len() - 1;
let total = cu_seqlens[batch];
assert!(num_heads > 0, "standard: num_heads must be non-zero");
assert!(head_dim > 0, "standard: head_dim must be non-zero");
assert_eq!(
hidden_size,
num_heads * head_dim,
"standard: hidden_size must equal num_heads * head_dim"
);
assert!(
total.checked_mul(hidden_size).is_some(),
"standard: total * hidden_size overflow"
);
let used_hidden = total * hidden_size;
assert_eq!(
hidden_states.len(),
used_hidden,
"standard: hidden_states length must equal total * hidden_size"
);
assert!(q.len() >= used_hidden, "standard: q scratch too small");
assert!(k.len() >= used_hidden, "standard: k scratch too small");
assert!(v.len() >= used_hidden, "standard: v scratch too small");
assert!(
concat.len() >= used_hidden,
"standard: concat scratch too small"
);
assert!(
output.len() >= used_hidden,
"standard: output scratch too small"
);
assert!(
qkv.len() >= used_hidden * 3,
"standard: qkv scratch too small"
);
{
let qkv = &mut qkv[..used_hidden * 3];
matmul_bt(
hidden_states,
fused_qkv_weight,
qkv,
total,
hidden_size,
3 * hidden_size,
);
add_bias(qkv, fused_qkv_bias, 3 * hidden_size);
for r in 0..total {
let src = r * 3 * hidden_size;
q[r * hidden_size..(r + 1) * hidden_size].copy_from_slice(&qkv[src..src + hidden_size]);
k[r * hidden_size..(r + 1) * hidden_size]
.copy_from_slice(&qkv[src + hidden_size..src + 2 * hidden_size]);
v[r * hidden_size..(r + 1) * hidden_size]
.copy_from_slice(&qkv[src + 2 * hidden_size..src + 3 * hidden_size]);
}
}
apply_lora_rows(
lora,
layer_idx,
"query",
hidden_states,
&mut q[..used_hidden],
hidden_size,
hidden_size,
);
apply_lora_rows(
lora,
layer_idx,
"key",
hidden_states,
&mut k[..used_hidden],
hidden_size,
hidden_size,
);
apply_lora_rows(
lora,
layer_idx,
"value",
hidden_states,
&mut v[..used_hidden],
hidden_size,
hidden_size,
);
let scale = 1.0 / (head_dim as f32).sqrt();
let q = &q[..used_hidden];
let k = &k[..used_hidden];
let v = &v[..used_hidden];
let concat = &mut concat[..used_hidden];
for b in 0..batch {
let start = cu_seqlens[b];
let end = cu_seqlens[b + 1];
assert!(
end >= start,
"standard: cu_seqlens must be non-decreasing (segment {b})"
);
let seq_len = end - start;
if seq_len == 0 {
continue;
}
assert_standard_no_overflow(seq_len, hidden_size, num_heads, head_dim);
let row_start = start * hidden_size;
let concat_b = &mut concat[row_start..row_start + seq_len * hidden_size];
let mut q_head = vec![0.0f32; seq_len * head_dim];
let mut k_head = vec![0.0f32; seq_len * head_dim];
let mut v_all_t = vec![0.0f32; hidden_size * seq_len];
let mut scores_head = vec![0.0f32; seq_len * seq_len];
let mut scores = vec![0.0f32; num_heads * seq_len * seq_len];
let mut context_head = vec![0.0f32; seq_len * head_dim];
for h in 0..num_heads {
let head_offset = h * head_dim;
for i in 0..seq_len {
let src_start = row_start + i * hidden_size + head_offset;
let dst_start = i * head_dim;
q_head[dst_start..dst_start + head_dim]
.copy_from_slice(&q[src_start..src_start + head_dim]);
}
for i in 0..seq_len {
let src_start = row_start + i * hidden_size + head_offset;
let dst_start = i * head_dim;
k_head[dst_start..dst_start + head_dim]
.copy_from_slice(&k[src_start..src_start + head_dim]);
}
matmul_bt(
&q_head[..seq_len * head_dim],
&k_head[..seq_len * head_dim],
&mut scores_head[..seq_len * seq_len],
seq_len,
head_dim,
seq_len,
);
let scores_offset = h * seq_len * seq_len;
for (idx, &score) in scores_head.iter().enumerate() {
scores[scores_offset + idx] = score * scale;
}
}
softmax_attention(&mut scores, seq_len, num_heads);
for i in 0..seq_len {
let v_row_start = row_start + i * hidden_size;
for d in 0..hidden_size {
v_all_t[d * seq_len + i] = v[v_row_start + d];
}
}
for h in 0..num_heads {
let head_offset = h * head_dim;
let scores_offset = h * seq_len * seq_len;
let scores_head = &scores[scores_offset..scores_offset + seq_len * seq_len];
let v_head_t = &v_all_t[head_offset * seq_len..(head_offset + head_dim) * seq_len];
matmul_bt(
scores_head,
v_head_t,
&mut context_head[..seq_len * head_dim],
seq_len,
seq_len,
head_dim,
);
for i in 0..seq_len {
let dst = i * hidden_size + head_offset;
concat_b[dst..dst + head_dim]
.copy_from_slice(&context_head[i * head_dim..(i + 1) * head_dim]);
}
}
}
let concat = &concat[..used_hidden];
let output = &mut output[..used_hidden];
matmul_bt(
concat,
layer_weights.attn_output_weight.data,
output,
total,
hidden_size,
hidden_size,
);
add_bias(output, layer_weights.attn_output_bias.data, hidden_size);
apply_lora_rows(
lora,
layer_idx,
"attn_output",
concat,
output,
hidden_size,
hidden_size,
);
}
#[cfg(test)]
#[allow(clippy::too_many_arguments)]
pub(crate) fn multi_head_attention_batched_padded_reference(
hidden_states: &[f32],
layer_weights: &TransformerLayerWeights<'_>,
fused_qkv_weight: &[f32],
fused_qkv_bias: &[f32],
attention_mask: &[u32],
batch: usize,
seq_len: usize,
hidden_size: usize,
num_heads: usize,
head_dim: usize,
q: &mut [f32],
k: &mut [f32],
v: &mut [f32],
qkv: &mut [f32],
concat: &mut [f32],
output: &mut [f32],
lora: &dyn LoraHook,
layer_idx: usize,
) {
assert_standard_no_overflow(seq_len, hidden_size, num_heads, head_dim);
assert!(
batch.checked_mul(seq_len).is_some(),
"standard: batch * seq_len overflow"
);
let rows = batch * seq_len;
assert!(
rows.checked_mul(hidden_size).is_some(),
"standard: rows * hidden_size overflow"
);
let used_hidden = rows * hidden_size;
assert_eq!(
hidden_states.len(),
used_hidden,
"standard: hidden_states length must equal batch * seq_len * hidden_size"
);
assert_eq!(
attention_mask.len(),
rows,
"standard: attention_mask length must equal batch * seq_len"
);
assert!(q.len() >= used_hidden, "standard: q scratch too small");
assert!(k.len() >= used_hidden, "standard: k scratch too small");
assert!(v.len() >= used_hidden, "standard: v scratch too small");
assert!(
concat.len() >= used_hidden,
"standard: concat scratch too small"
);
assert!(
output.len() >= used_hidden,
"standard: output scratch too small"
);
assert!(
qkv.len() >= used_hidden * 3,
"standard: qkv scratch too small"
);
{
let qkv = &mut qkv[..used_hidden * 3];
matmul_bt(
hidden_states,
fused_qkv_weight,
qkv,
rows,
hidden_size,
3 * hidden_size,
);
add_bias(qkv, fused_qkv_bias, 3 * hidden_size);
for r in 0..rows {
let src = r * 3 * hidden_size;
q[r * hidden_size..(r + 1) * hidden_size].copy_from_slice(&qkv[src..src + hidden_size]);
k[r * hidden_size..(r + 1) * hidden_size]
.copy_from_slice(&qkv[src + hidden_size..src + 2 * hidden_size]);
v[r * hidden_size..(r + 1) * hidden_size]
.copy_from_slice(&qkv[src + 2 * hidden_size..src + 3 * hidden_size]);
}
}
apply_lora_rows(
lora,
layer_idx,
"query",
hidden_states,
&mut q[..used_hidden],
hidden_size,
hidden_size,
);
apply_lora_rows(
lora,
layer_idx,
"key",
hidden_states,
&mut k[..used_hidden],
hidden_size,
hidden_size,
);
apply_lora_rows(
lora,
layer_idx,
"value",
hidden_states,
&mut v[..used_hidden],
hidden_size,
hidden_size,
);
let scale = 1.0 / (head_dim as f32).sqrt();
let q = &q[..used_hidden];
let k = &k[..used_hidden];
let v = &v[..used_hidden];
let concat = &mut concat[..used_hidden];
concat
.chunks_mut(seq_len * hidden_size)
.enumerate()
.for_each(|(b, concat_b)| {
let seq_offset = b * seq_len;
let row_start = seq_offset * hidden_size;
let mask_b = &attention_mask[seq_offset..seq_offset + seq_len];
let mut q_head = vec![0.0f32; seq_len * head_dim];
let mut k_head = vec![0.0f32; seq_len * head_dim];
let mut v_all_t = vec![0.0f32; hidden_size * seq_len];
let mut scores_head = vec![0.0f32; seq_len * seq_len];
let mut scores = vec![0.0f32; num_heads * seq_len * seq_len];
let mut context_head = vec![0.0f32; seq_len * head_dim];
for h in 0..num_heads {
let head_offset = h * head_dim;
for i in 0..seq_len {
let src_start = row_start + i * hidden_size + head_offset;
let dst_start = i * head_dim;
q_head[dst_start..dst_start + head_dim]
.copy_from_slice(&q[src_start..src_start + head_dim]);
}
for i in 0..seq_len {
let src_start = row_start + i * hidden_size + head_offset;
let dst_start = i * head_dim;
k_head[dst_start..dst_start + head_dim]
.copy_from_slice(&k[src_start..src_start + head_dim]);
}
matmul_bt(
&q_head[..seq_len * head_dim],
&k_head[..seq_len * head_dim],
&mut scores_head[..seq_len * seq_len],
seq_len,
head_dim,
seq_len,
);
let scores_offset = h * seq_len * seq_len;
for (idx, &score) in scores_head.iter().enumerate() {
scores[scores_offset + idx] = score * scale;
}
}
for h in 0..num_heads {
for i in 0..seq_len {
let row_off = (h * seq_len + i) * seq_len;
let row = &mut scores[row_off..row_off + seq_len];
for j in 0..seq_len {
if mask_b[j] == 0 {
row[j] = f32::NEG_INFINITY;
}
}
}
}
softmax_attention(&mut scores, seq_len, num_heads);
for i in 0..seq_len {
let v_row_start = row_start + i * hidden_size;
for d in 0..hidden_size {
v_all_t[d * seq_len + i] = v[v_row_start + d];
}
}
for h in 0..num_heads {
let head_offset = h * head_dim;
let scores_offset = h * seq_len * seq_len;
let scores_head = &scores[scores_offset..scores_offset + seq_len * seq_len];
let v_head_t = &v_all_t[head_offset * seq_len..(head_offset + head_dim) * seq_len];
matmul_bt(
scores_head,
v_head_t,
&mut context_head[..seq_len * head_dim],
seq_len,
seq_len,
head_dim,
);
for i in 0..seq_len {
let dst = i * hidden_size + head_offset;
concat_b[dst..dst + head_dim]
.copy_from_slice(&context_head[i * head_dim..(i + 1) * head_dim]);
}
}
});
let concat = &concat[..used_hidden];
let output = &mut output[..used_hidden];
matmul_bt(
concat,
layer_weights.attn_output_weight.data,
output,
rows,
hidden_size,
hidden_size,
);
add_bias(output, layer_weights.attn_output_bias.data, hidden_size);
apply_lora_rows(
lora,
layer_idx,
"attn_output",
concat,
output,
hidden_size,
hidden_size,
);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::lora_hook::NoopLoraHook;
use crate::weights::{Tensor1D, Tensor2D, TransformerLayerWeights};
use std::sync::atomic::{AtomicUsize, Ordering};
struct RowShapeHook {
row_width: usize,
calls: AtomicUsize,
}
impl LoraHook for RowShapeHook {
fn apply(&self, _layer_idx: usize, _module: &str, x: &[f32], output: &mut [f32]) {
assert_eq!(x.len(), self.row_width);
assert_eq!(output.len(), self.row_width);
self.calls.fetch_add(1, Ordering::Relaxed);
}
}
#[test]
fn test_attention_simd_matches_expected() {
let seq_len = 2;
let num_heads = 2;
let head_dim = 2;
let hidden_size = num_heads * head_dim;
let hidden_states = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
let identity_4x4: Vec<f32> = vec![
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, ];
let zero_bias_4: Vec<f32> = vec![0.0; 4];
let ones_4: Vec<f32> = vec![1.0; 4];
let intermediate_size = hidden_size;
let layer = TransformerLayerWeights {
query_weight: Tensor2D {
data: &identity_4x4,
rows: hidden_size,
cols: hidden_size,
},
query_bias: Tensor1D {
data: &zero_bias_4,
len: hidden_size,
},
key_weight: Tensor2D {
data: &identity_4x4,
rows: hidden_size,
cols: hidden_size,
},
key_bias: Tensor1D {
data: &zero_bias_4,
len: hidden_size,
},
value_weight: Tensor2D {
data: &identity_4x4,
rows: hidden_size,
cols: hidden_size,
},
value_bias: Tensor1D {
data: &zero_bias_4,
len: hidden_size,
},
attn_output_weight: Tensor2D {
data: &identity_4x4,
rows: hidden_size,
cols: hidden_size,
},
attn_output_bias: Tensor1D {
data: &zero_bias_4,
len: hidden_size,
},
attn_layer_norm_weight: Tensor1D {
data: &ones_4,
len: hidden_size,
},
attn_layer_norm_bias: Tensor1D {
data: &zero_bias_4,
len: hidden_size,
},
ffn_intermediate_weight: Tensor2D {
data: &identity_4x4,
rows: intermediate_size,
cols: hidden_size,
},
ffn_intermediate_bias: Tensor1D {
data: &zero_bias_4,
len: intermediate_size,
},
ffn_output_weight: Tensor2D {
data: &identity_4x4,
rows: hidden_size,
cols: intermediate_size,
},
ffn_output_bias: Tensor1D {
data: &zero_bias_4,
len: hidden_size,
},
ffn_layer_norm_weight: Tensor1D {
data: &ones_4,
len: hidden_size,
},
ffn_layer_norm_bias: Tensor1D {
data: &zero_bias_4,
len: hidden_size,
},
};
let attention_mask = vec![1u32; seq_len];
let mut buffers = AttentionBuffers::new(seq_len, hidden_size, num_heads, intermediate_size);
let result = multi_head_attention(
&hidden_states,
&layer,
&attention_mask,
seq_len,
hidden_size,
num_heads,
head_dim,
&mut buffers,
&NoopLoraHook,
0,
);
assert_eq!(result.len(), seq_len * hidden_size);
for (i, &val) in result.iter().enumerate() {
assert!(val.is_finite(), "result[{i}] = {val} is not finite");
}
let mut buffers2 =
AttentionBuffers::new(seq_len, hidden_size, num_heads, intermediate_size);
let result2 = multi_head_attention(
&hidden_states,
&layer,
&attention_mask,
seq_len,
hidden_size,
num_heads,
head_dim,
&mut buffers2,
&NoopLoraHook,
0,
);
assert_eq!(result, result2, "attention must be deterministic");
}
#[test]
fn test_attention_mask_suppresses_tokens() {
let seq_len = 3;
let num_heads = 1;
let head_dim = 2;
let hidden_size = num_heads * head_dim;
let hidden_states = vec![1.0, 0.0, 0.0, 1.0, 1.0, 1.0];
let identity_2x2: Vec<f32> = vec![1.0, 0.0, 0.0, 1.0];
let zero_bias_2: Vec<f32> = vec![0.0; 2];
let ones_2: Vec<f32> = vec![1.0; 2];
let intermediate_size = hidden_size;
let layer = TransformerLayerWeights {
query_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: hidden_size,
},
query_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
key_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: hidden_size,
},
key_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
value_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: hidden_size,
},
value_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
attn_output_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: hidden_size,
},
attn_output_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
attn_layer_norm_weight: Tensor1D {
data: &ones_2,
len: hidden_size,
},
attn_layer_norm_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
ffn_intermediate_weight: Tensor2D {
data: &identity_2x2,
rows: intermediate_size,
cols: hidden_size,
},
ffn_intermediate_bias: Tensor1D {
data: &zero_bias_2,
len: intermediate_size,
},
ffn_output_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: intermediate_size,
},
ffn_output_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
ffn_layer_norm_weight: Tensor1D {
data: &ones_2,
len: hidden_size,
},
ffn_layer_norm_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
};
let mask_all = vec![1u32, 1, 1];
let mask_partial = vec![1u32, 1, 0];
let mut buf1 = AttentionBuffers::new(seq_len, hidden_size, num_heads, intermediate_size);
let mut buf2 = AttentionBuffers::new(seq_len, hidden_size, num_heads, intermediate_size);
let result_all = multi_head_attention(
&hidden_states,
&layer,
&mask_all,
seq_len,
hidden_size,
num_heads,
head_dim,
&mut buf1,
&NoopLoraHook,
0,
);
let result_masked = multi_head_attention(
&hidden_states,
&layer,
&mask_partial,
seq_len,
hidden_size,
num_heads,
head_dim,
&mut buf2,
&NoopLoraHook,
0,
);
assert_ne!(
result_all, result_masked,
"masking a token should change attention output"
);
for &v in result_all.iter().chain(result_masked.iter()) {
assert!(v.is_finite());
}
}
#[test]
fn masked_token_value_does_not_leak_when_valid_score_below_sentinel() {
let seq_len = 2;
let num_heads = 1;
let head_dim = 2;
let hidden_size = num_heads * head_dim;
let hidden_states = vec![1.0, 0.0, 500.0, 500.0];
let query_w: Vec<f32> = vec![200.0, 0.0, 0.0, 0.0];
let key_w: Vec<f32> = vec![-100.0, 0.0, 0.0, 0.0];
let identity_2x2: Vec<f32> = vec![1.0, 0.0, 0.0, 1.0];
let zero_bias_2: Vec<f32> = vec![0.0; 2];
let ones_2: Vec<f32> = vec![1.0; 2];
let intermediate_size = hidden_size;
let layer = TransformerLayerWeights {
query_weight: Tensor2D {
data: &query_w,
rows: hidden_size,
cols: hidden_size,
},
query_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
key_weight: Tensor2D {
data: &key_w,
rows: hidden_size,
cols: hidden_size,
},
key_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
value_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: hidden_size,
},
value_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
attn_output_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: hidden_size,
},
attn_output_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
attn_layer_norm_weight: Tensor1D {
data: &ones_2,
len: hidden_size,
},
attn_layer_norm_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
ffn_intermediate_weight: Tensor2D {
data: &identity_2x2,
rows: intermediate_size,
cols: hidden_size,
},
ffn_intermediate_bias: Tensor1D {
data: &zero_bias_2,
len: intermediate_size,
},
ffn_output_weight: Tensor2D {
data: &identity_2x2,
rows: hidden_size,
cols: intermediate_size,
},
ffn_output_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
ffn_layer_norm_weight: Tensor1D {
data: &ones_2,
len: hidden_size,
},
ffn_layer_norm_bias: Tensor1D {
data: &zero_bias_2,
len: hidden_size,
},
};
let mask = vec![1u32, 0];
let mut buf = AttentionBuffers::new(seq_len, hidden_size, num_heads, intermediate_size);
let out = multi_head_attention(
&hidden_states,
&layer,
&mask,
seq_len,
hidden_size,
num_heads,
head_dim,
&mut buf,
&NoopLoraHook,
0,
);
assert!(
out.iter().all(|v| v.is_finite()),
"output must be finite: {out:?}"
);
assert!(
out[0].abs() < 50.0 && out[1].abs() < 50.0,
"masked token value leaked into row 0 output: {:?} (expected ~[1,0])",
&out[0..2]
);
}
#[test]
fn standard_no_overflow_accepts_valid_shape() {
assert_standard_no_overflow(8, 64, 8, 8);
}
#[test]
#[should_panic(expected = "hidden_size must equal num_heads * head_dim")]
fn standard_no_overflow_rejects_layout_mismatch() {
assert_standard_no_overflow(1, 4, 1, 2);
}
#[test]
#[should_panic(expected = "num_heads * seq_len * seq_len")]
fn standard_no_overflow_rejects_wrapping_product() {
assert_standard_no_overflow(1usize << 32, 2, 2, 1);
}
#[test]
fn batched_attention_matches_single_sequence_per_row_packed() {
let hidden_size = 8;
let num_heads = 2;
let head_dim = 4;
let intermediate_size = hidden_size;
let identity_8x8: Vec<f32> = {
let mut m = vec![0.0f32; hidden_size * hidden_size];
for i in 0..hidden_size {
m[i * hidden_size + i] = 1.0;
}
m
};
let scaled_identity = |scale: f32| -> Vec<f32> {
let mut m = vec![0.0f32; hidden_size * hidden_size];
for i in 0..hidden_size {
m[i * hidden_size + i] = scale;
}
m
};
let zero_bias_8: Vec<f32> = vec![0.0; hidden_size];
let ones_8: Vec<f32> = vec![1.0; hidden_size];
let query_w: Vec<f32> = scaled_identity(1.0);
let key_w: Vec<f32> = scaled_identity(2.0);
let value_w: Vec<f32> = scaled_identity(3.0);
let query_bias_v: Vec<f32> = (0..hidden_size).map(|i| 0.1 * (i as f32 + 1.0)).collect();
let key_bias_v: Vec<f32> = (0..hidden_size).map(|i| 1.0 + i as f32).collect();
let value_bias_v: Vec<f32> = (0..hidden_size).map(|i| 10.0 + i as f32).collect();
let mut fused_qkv_weight: Vec<f32> = Vec::with_capacity(3 * hidden_size * hidden_size);
fused_qkv_weight.extend_from_slice(&query_w);
fused_qkv_weight.extend_from_slice(&key_w);
fused_qkv_weight.extend_from_slice(&value_w);
let mut fused_qkv_bias: Vec<f32> = Vec::with_capacity(3 * hidden_size);
fused_qkv_bias.extend_from_slice(&query_bias_v);
fused_qkv_bias.extend_from_slice(&key_bias_v);
fused_qkv_bias.extend_from_slice(&value_bias_v);
let layer = TransformerLayerWeights {
query_weight: Tensor2D {
data: &query_w,
rows: hidden_size,
cols: hidden_size,
},
query_bias: Tensor1D {
data: &query_bias_v,
len: hidden_size,
},
key_weight: Tensor2D {
data: &key_w,
rows: hidden_size,
cols: hidden_size,
},
key_bias: Tensor1D {
data: &key_bias_v,
len: hidden_size,
},
value_weight: Tensor2D {
data: &value_w,
rows: hidden_size,
cols: hidden_size,
},
value_bias: Tensor1D {
data: &value_bias_v,
len: hidden_size,
},
attn_output_weight: Tensor2D {
data: &identity_8x8,
rows: hidden_size,
cols: hidden_size,
},
attn_output_bias: Tensor1D {
data: &zero_bias_8,
len: hidden_size,
},
attn_layer_norm_weight: Tensor1D {
data: &ones_8,
len: hidden_size,
},
attn_layer_norm_bias: Tensor1D {
data: &zero_bias_8,
len: hidden_size,
},
ffn_intermediate_weight: Tensor2D {
data: &identity_8x8,
rows: intermediate_size,
cols: hidden_size,
},
ffn_intermediate_bias: Tensor1D {
data: &zero_bias_8,
len: intermediate_size,
},
ffn_output_weight: Tensor2D {
data: &identity_8x8,
rows: intermediate_size,
cols: hidden_size,
},
ffn_output_bias: Tensor1D {
data: &zero_bias_8,
len: hidden_size,
},
ffn_layer_norm_weight: Tensor1D {
data: &ones_8,
len: hidden_size,
},
ffn_layer_norm_bias: Tensor1D {
data: &zero_bias_8,
len: hidden_size,
},
};
let seq0_real: Vec<f32> = (0..2 * hidden_size).map(|i| 1.0 + i as f32 * 0.1).collect();
let seq1_real: Vec<f32> = (0..3 * hidden_size)
.map(|i| 100.0 + i as f32 * 0.1)
.collect();
let mut hidden_states_packed = Vec::with_capacity(5 * hidden_size);
hidden_states_packed.extend_from_slice(&seq0_real);
hidden_states_packed.extend_from_slice(&seq1_real);
let cu_seqlens = vec![0usize, 2, 5];
let total = *cu_seqlens.last().unwrap();
let used_hidden = total * hidden_size;
let mut q = vec![0.0f32; used_hidden];
let mut k = vec![0.0f32; used_hidden];
let mut v = vec![0.0f32; used_hidden];
let mut qkv = vec![0.0f32; 3 * used_hidden];
let mut concat = vec![0.0f32; used_hidden];
let mut output = vec![0.0f32; used_hidden];
let lora = RowShapeHook {
row_width: hidden_size,
calls: AtomicUsize::new(0),
};
multi_head_attention_batched(
&hidden_states_packed,
&layer,
&fused_qkv_weight,
&fused_qkv_bias,
&cu_seqlens,
hidden_size,
num_heads,
head_dim,
&mut q,
&mut k,
&mut v,
&mut qkv,
&mut concat,
&mut output,
&lora,
0,
);
assert_eq!(lora.calls.load(Ordering::Relaxed), total * 4);
for out_val in output.iter() {
assert!(
out_val.is_finite(),
"batched output must be finite: {output:?}"
);
}
let mut expected_q = vec![0.0f32; used_hidden];
matmul_bt(
&hidden_states_packed,
&query_w,
&mut expected_q,
total,
hidden_size,
hidden_size,
);
add_bias(&mut expected_q, &query_bias_v, hidden_size);
for (i, (&got, &exp)) in q.iter().zip(expected_q.iter()).enumerate() {
assert!(
(got - exp).abs() <= 1e-5,
"q scratch element {i} mismatch: batched={got} independent={exp}"
);
}
let mut expected_k = vec![0.0f32; used_hidden];
matmul_bt(
&hidden_states_packed,
&key_w,
&mut expected_k,
total,
hidden_size,
hidden_size,
);
add_bias(&mut expected_k, &key_bias_v, hidden_size);
for (i, (&got, &exp)) in k.iter().zip(expected_k.iter()).enumerate() {
assert!(
(got - exp).abs() <= 1e-5,
"k scratch element {i} mismatch: batched={got} independent={exp}"
);
}
let mut expected_v = vec![0.0f32; used_hidden];
matmul_bt(
&hidden_states_packed,
&value_w,
&mut expected_v,
total,
hidden_size,
hidden_size,
);
add_bias(&mut expected_v, &value_bias_v, hidden_size);
for (i, (&got, &exp)) in v.iter().zip(expected_v.iter()).enumerate() {
assert!(
(got - exp).abs() <= 1e-5,
"v scratch element {i} mismatch: batched={got} independent={exp}"
);
}
let mut buf0 = AttentionBuffers::new(2, hidden_size, num_heads, intermediate_size);
let expected0 = multi_head_attention(
&seq0_real,
&layer,
&[1u32, 1],
2,
hidden_size,
num_heads,
head_dim,
&mut buf0,
&NoopLoraHook,
0,
);
let seq0_row_start = 0;
let got0 = &output[seq0_row_start..seq0_row_start + 2 * hidden_size];
for (i, (&g, &e)) in got0.iter().zip(expected0.iter()).enumerate() {
assert!(
(g - e).abs() <= 1e-6,
"seq0 row element {i} mismatch: batched={g} single={e}"
);
}
let mut buf1 = AttentionBuffers::new(3, hidden_size, num_heads, intermediate_size);
let expected1 = multi_head_attention(
&seq1_real,
&layer,
&[1u32, 1, 1],
3,
hidden_size,
num_heads,
head_dim,
&mut buf1,
&NoopLoraHook,
0,
);
let seq1_row_start = 2 * hidden_size; let got1 = &output[seq1_row_start..seq1_row_start + 3 * hidden_size];
for (i, (&g, &e)) in got1.iter().zip(expected1.iter()).enumerate() {
assert!(
(g - e).abs() <= 1e-6,
"seq1 row element {i} mismatch: batched={g} single={e}"
);
}
}
}