use crate::attention::gqa::GqaConfig;
use crate::error::InferenceError;
use crate::forward::cpu::{elementwise_mul, matmul_bt, rms_norm, silu_inplace};
use crate::grammar::GrammarEngine;
use crate::kv_cache::{FlatKVCache, FlatKVCacheConfig};
use crate::model::qwen::{QwenConfig, QwenModel};
use crate::sampling::{Sampler, SamplingConfig};
use std::sync::Arc;
#[derive(Clone)]
pub struct GenerateConfig {
pub max_new_tokens: usize,
pub sampling: SamplingConfig,
pub eos_token_id: Option<u32>,
pub include_prompt: bool,
pub grammar: Option<Arc<GrammarEngine>>,
}
impl std::fmt::Debug for GenerateConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("GenerateConfig")
.field("max_new_tokens", &self.max_new_tokens)
.field("sampling", &self.sampling)
.field("eos_token_id", &self.eos_token_id)
.field("include_prompt", &self.include_prompt)
.field("grammar", &self.grammar.as_ref().map(|_| "<GrammarEngine>"))
.finish()
}
}
impl Default for GenerateConfig {
fn default() -> Self {
Self {
max_new_tokens: 256,
sampling: SamplingConfig::default(),
eos_token_id: None,
include_prompt: false,
grammar: None,
}
}
}
#[derive(Debug)]
pub struct GenerateOutput {
pub text: String,
pub token_ids: Vec<u32>,
pub prompt_tokens: usize,
pub generated_tokens: usize,
pub stopped_by_eos: bool,
}
struct ForwardScratch {
hidden: Vec<f32>, residual: Vec<f32>, qkv_buf: Vec<f32>, q_buf: Vec<f32>, k_buf: Vec<f32>, v_buf: Vec<f32>, attn_out: Vec<f32>, gate_up_buf: Vec<f32>, gate_buf: Vec<f32>, up_buf: Vec<f32>, ffn_out: Vec<f32>, scores: Vec<f32>, logits: Vec<f32>, cached_k_f32: Vec<f32>, cached_v_f32: Vec<f32>, }
impl ForwardScratch {
fn new() -> Self {
Self {
hidden: Vec::new(),
residual: Vec::new(),
qkv_buf: Vec::new(),
q_buf: Vec::new(),
k_buf: Vec::new(),
v_buf: Vec::new(),
attn_out: Vec::new(),
gate_up_buf: Vec::new(),
gate_buf: Vec::new(),
up_buf: Vec::new(),
ffn_out: Vec::new(),
scores: Vec::new(),
logits: Vec::new(),
cached_k_f32: Vec::new(),
cached_v_f32: Vec::new(),
}
}
fn ensure_capacity(&mut self, cfg: &QwenConfig, seq_len_cap: usize, max_seq_len: usize) {
let h = cfg.hidden_size;
let q_dim = cfg.q_dim();
let kv_dim = cfg.kv_dim();
let qkv_dim = q_dim + 2 * kv_dim;
let inter = cfg.intermediate_size;
grow(&mut self.hidden, seq_len_cap * h);
grow(&mut self.residual, seq_len_cap * h);
grow(&mut self.qkv_buf, seq_len_cap * qkv_dim);
grow(&mut self.q_buf, seq_len_cap * q_dim);
grow(&mut self.k_buf, seq_len_cap * kv_dim);
grow(&mut self.v_buf, seq_len_cap * kv_dim);
grow(&mut self.attn_out, seq_len_cap * q_dim);
grow(&mut self.gate_up_buf, seq_len_cap * 2 * inter);
grow(&mut self.gate_buf, seq_len_cap * inter);
grow(&mut self.up_buf, seq_len_cap * inter);
grow(&mut self.ffn_out, seq_len_cap * h);
grow(&mut self.scores, cfg.num_attention_heads * max_seq_len);
grow(&mut self.logits, cfg.vocab_size);
let kv_dim = cfg.kv_dim();
grow(&mut self.cached_k_f32, max_seq_len * kv_dim);
grow(&mut self.cached_v_f32, max_seq_len * kv_dim);
}
}
fn grow(buf: &mut Vec<f32>, n: usize) {
if buf.len() < n {
buf.resize(n, 0.0);
}
}
pub fn generate(
model: &QwenModel,
prompt: &str,
config: &GenerateConfig,
) -> Result<GenerateOutput, InferenceError> {
let cfg = model.config();
let tokenizer = model.tokenizer();
let tokenized = tokenizer.tokenize(prompt);
let prompt_ids: Vec<u32> = tokenized.input_ids.clone();
let prompt_len = tokenized.real_length;
if prompt_len == 0 {
return Err(InferenceError::InvalidInput("Empty prompt".into()));
}
let max_seq = prompt_len + config.max_new_tokens;
let cache_cfg = FlatKVCacheConfig::for_qwen3(
cfg.num_hidden_layers,
cfg.num_key_value_heads,
cfg.head_dim,
max_seq,
);
let mut cache = FlatKVCache::new(cache_cfg);
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(cfg, prompt_len.max(1), max_seq);
let mut sampler = Sampler::new(config.sampling.clone());
let mut grammar_state = config.grammar.as_ref().map(|g| g.initial_state());
forward_with_cache(
model,
&prompt_ids[..prompt_len],
&mut cache,
0,
&mut scratch,
max_seq,
)?;
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
}
let mut generated_ids: Vec<u32> = Vec::with_capacity(config.max_new_tokens);
let first_token = sampler.sample(&scratch.logits[..cfg.vocab_size]);
generated_ids.push(first_token);
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state) {
if !engine.advance(gs, first_token) {
let text = tokenizer.decode(&generated_ids).unwrap_or_default();
return Ok(GenerateOutput {
text,
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
token_ids: generated_ids,
stopped_by_eos: false,
});
}
}
let mut stopped_by_eos = false;
if config.eos_token_id == Some(first_token) {
stopped_by_eos = true;
}
if !stopped_by_eos {
for step in 0..config.max_new_tokens.saturating_sub(1) {
let pos = prompt_len + step + 1; let input = [*generated_ids
.last()
.expect("invariant: first generated token exists before decode loop")];
forward_with_cache(model, &input, &mut cache, pos - 1, &mut scratch, max_seq)?;
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
}
let token = sampler.sample(&scratch.logits[..cfg.vocab_size]);
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state) {
if !engine.advance(gs, token) {
break;
}
}
generated_ids.push(token);
if config.eos_token_id == Some(token) {
stopped_by_eos = true;
break;
}
}
}
let decoded = tokenizer.decode(&generated_ids).unwrap_or_default();
let full_text = if config.include_prompt {
format!("{prompt}{decoded}")
} else {
decoded
};
Ok(GenerateOutput {
text: full_text,
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
token_ids: generated_ids,
stopped_by_eos,
})
}
fn forward_with_cache<'a>(
model: &QwenModel,
input_ids: &[u32],
cache: &mut FlatKVCache,
start_pos: usize,
scratch: &'a mut ForwardScratch,
max_seq_len: usize,
) -> Result<&'a [f32], InferenceError> {
let cfg = model.config();
let seq_len = input_ids.len();
let hidden_size = cfg.hidden_size;
let q_dim = cfg.q_dim();
let kv_dim = cfg.kv_dim();
let qkv_dim = q_dim + 2 * kv_dim;
let inter = cfg.intermediate_size;
let head_dim = cfg.head_dim;
scratch.ensure_capacity(cfg, seq_len, max_seq_len);
let weights = model.weights();
let rope = model.rope();
for (i, &tok) in input_ids.iter().enumerate() {
let tok = tok as usize;
if tok >= cfg.vocab_size {
return Err(InferenceError::InvalidInput(format!(
"Token ID {tok} exceeds vocab size {}",
cfg.vocab_size
)));
}
let row = &weights.embed_tokens.data[tok * hidden_size..(tok + 1) * hidden_size];
scratch.hidden[i * hidden_size..(i + 1) * hidden_size].copy_from_slice(row);
}
let gqa_cfg = GqaConfig {
num_heads: cfg.num_attention_heads,
num_kv_heads: cfg.num_key_value_heads,
head_dim,
};
for layer_idx in 0..cfg.num_hidden_layers {
let lw = &weights.layers[layer_idx];
scratch.residual[..seq_len * hidden_size]
.copy_from_slice(&scratch.hidden[..seq_len * hidden_size]);
rms_norm(
&mut scratch.hidden[..seq_len * hidden_size],
lw.input_layernorm_weight.data,
hidden_size,
cfg.rms_norm_eps,
);
matmul_bt(
&scratch.hidden[..seq_len * hidden_size],
&lw.fused_qkv,
&mut scratch.qkv_buf[..seq_len * qkv_dim],
seq_len,
hidden_size,
qkv_dim,
);
for i in 0..seq_len {
let qkv_row = i * qkv_dim;
scratch.q_buf[i * q_dim..(i + 1) * q_dim]
.copy_from_slice(&scratch.qkv_buf[qkv_row..qkv_row + q_dim]);
scratch.k_buf[i * kv_dim..(i + 1) * kv_dim]
.copy_from_slice(&scratch.qkv_buf[qkv_row + q_dim..qkv_row + q_dim + kv_dim]);
scratch.v_buf[i * kv_dim..(i + 1) * kv_dim]
.copy_from_slice(&scratch.qkv_buf[qkv_row + q_dim + kv_dim..qkv_row + qkv_dim]);
}
for i in 0..seq_len {
for h in 0..cfg.num_attention_heads {
let off = i * q_dim + h * head_dim;
rms_norm(
&mut scratch.q_buf[off..off + head_dim],
lw.q_norm_weight.data,
head_dim,
cfg.rms_norm_eps,
);
}
for h in 0..cfg.num_key_value_heads {
let off = i * kv_dim + h * head_dim;
rms_norm(
&mut scratch.k_buf[off..off + head_dim],
lw.k_norm_weight.data,
head_dim,
cfg.rms_norm_eps,
);
}
}
for i in 0..seq_len {
let pos = start_pos + i;
for h in 0..cfg.num_attention_heads {
let off = i * q_dim + h * head_dim;
rope.apply(&mut scratch.q_buf[off..off + head_dim], pos);
}
for h in 0..cfg.num_key_value_heads {
let off = i * kv_dim + h * head_dim;
rope.apply(&mut scratch.k_buf[off..off + head_dim], pos);
}
}
{
let base_pos = cache.seq_len();
let k_layer = cache.k_buffer_mut(layer_idx);
for i in 0..seq_len {
let dst_off = (base_pos + i) * kv_dim;
for (j, &val) in scratch.k_buf[i * kv_dim..(i + 1) * kv_dim]
.iter()
.enumerate()
{
k_layer[dst_off + j] = half::f16::from_f32(val);
}
}
let v_layer = cache.v_buffer_mut(layer_idx);
for i in 0..seq_len {
let dst_off = (base_pos + i) * kv_dim;
for (j, &val) in scratch.v_buf[i * kv_dim..(i + 1) * kv_dim]
.iter()
.enumerate()
{
v_layer[dst_off + j] = half::f16::from_f32(val);
}
}
}
let cached_seq_len = cache.seq_len() + seq_len; let k_end = cached_seq_len * kv_dim;
for (i, &h) in cache.k_buffer(layer_idx)[..k_end].iter().enumerate() {
scratch.cached_k_f32[i] = h.to_f32();
}
for (i, &h) in cache.v_buffer(layer_idx)[..k_end].iter().enumerate() {
scratch.cached_v_f32[i] = h.to_f32();
}
{
let q_ptr = scratch.q_buf.as_ptr();
let ck_ptr = scratch.cached_k_f32.as_ptr();
let cv_ptr = scratch.cached_v_f32.as_ptr();
let attn_ptr = scratch.attn_out.as_mut_ptr();
let scores_ptr = scratch.scores.as_mut_ptr();
let attn_len = seq_len * q_dim;
let scores_len = scratch.scores.len();
let (q_slice, cached_k, cached_v, attn_slice, scores_slice) = unsafe {
(
std::slice::from_raw_parts(q_ptr, seq_len * q_dim),
std::slice::from_raw_parts(ck_ptr, k_end),
std::slice::from_raw_parts(cv_ptr, k_end),
std::slice::from_raw_parts_mut(attn_ptr, attn_len),
std::slice::from_raw_parts_mut(scores_ptr, scores_len),
)
};
compute_attention(
attn_slice,
q_slice,
cached_k,
cached_v,
seq_len,
cached_seq_len,
start_pos,
&gqa_cfg,
scores_slice,
max_seq_len,
);
}
matmul_bt(
&scratch.attn_out[..seq_len * q_dim],
lw.o_proj_weight.data,
&mut scratch.hidden[..seq_len * hidden_size],
seq_len,
q_dim,
hidden_size,
);
for i in 0..seq_len * hidden_size {
scratch.hidden[i] += scratch.residual[i];
}
scratch.residual[..seq_len * hidden_size]
.copy_from_slice(&scratch.hidden[..seq_len * hidden_size]);
rms_norm(
&mut scratch.hidden[..seq_len * hidden_size],
lw.post_attention_layernorm_weight.data,
hidden_size,
cfg.rms_norm_eps,
);
matmul_bt(
&scratch.hidden[..seq_len * hidden_size],
&lw.fused_gate_up,
&mut scratch.gate_up_buf[..seq_len * 2 * inter],
seq_len,
hidden_size,
2 * inter,
);
for i in 0..seq_len {
let gu_row = i * 2 * inter;
scratch.gate_buf[i * inter..(i + 1) * inter]
.copy_from_slice(&scratch.gate_up_buf[gu_row..gu_row + inter]);
scratch.up_buf[i * inter..(i + 1) * inter]
.copy_from_slice(&scratch.gate_up_buf[gu_row + inter..gu_row + 2 * inter]);
}
silu_inplace(&mut scratch.gate_buf[..seq_len * inter]);
elementwise_mul(
&mut scratch.gate_buf[..seq_len * inter],
&scratch.up_buf[..seq_len * inter],
);
matmul_bt(
&scratch.gate_buf[..seq_len * inter],
lw.down_proj_weight.data,
&mut scratch.ffn_out[..seq_len * hidden_size],
seq_len,
inter,
hidden_size,
);
for i in 0..seq_len * hidden_size {
scratch.hidden[i] = scratch.residual[i] + scratch.ffn_out[i];
}
}
cache.advance_by(seq_len);
let last_start = (seq_len - 1) * hidden_size;
rms_norm(
&mut scratch.hidden[last_start..last_start + hidden_size],
weights.norm_weight.data,
hidden_size,
cfg.rms_norm_eps,
);
matmul_bt(
&scratch.hidden[last_start..last_start + hidden_size],
weights.embed_tokens.data,
&mut scratch.logits[..cfg.vocab_size],
1,
hidden_size,
cfg.vocab_size,
);
Ok(&scratch.logits[..cfg.vocab_size])
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn dot_f32_neon_128(q: *const f32, k: *const f32) -> f32 {
use std::arch::aarch64::*;
let mut acc0 = vdupq_n_f32(0.0);
let mut acc1 = vdupq_n_f32(0.0);
let mut acc2 = vdupq_n_f32(0.0);
let mut acc3 = vdupq_n_f32(0.0);
let mut d = 0usize;
while d < 128 {
let q0 = vld1q_f32(q.add(d));
let k0 = vld1q_f32(k.add(d));
let q1 = vld1q_f32(q.add(d + 4));
let k1 = vld1q_f32(k.add(d + 4));
let q2 = vld1q_f32(q.add(d + 8));
let k2 = vld1q_f32(k.add(d + 8));
let q3 = vld1q_f32(q.add(d + 12));
let k3 = vld1q_f32(k.add(d + 12));
acc0 = vfmaq_f32(acc0, q0, k0);
acc1 = vfmaq_f32(acc1, q1, k1);
acc2 = vfmaq_f32(acc2, q2, k2);
acc3 = vfmaq_f32(acc3, q3, k3);
d += 16;
}
let acc = vaddq_f32(vaddq_f32(acc0, acc1), vaddq_f32(acc2, acc3));
vaddvq_f32(acc)
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn accum_v_neon_128(
out: *mut f32,
scores: *const f32,
v_base: *const f32,
kv_len: usize,
kv_row_stride: usize,
) {
use std::arch::aarch64::*;
let mut d = 0usize;
while d < 128 {
let mut acc0 = vdupq_n_f32(0.0);
let mut acc1 = vdupq_n_f32(0.0);
let mut acc2 = vdupq_n_f32(0.0);
let mut acc3 = vdupq_n_f32(0.0);
for ki in 0..kv_len {
let w = vdupq_n_f32(*scores.add(ki));
let v_row = v_base.add(ki * kv_row_stride + d);
acc0 = vfmaq_f32(acc0, w, vld1q_f32(v_row));
acc1 = vfmaq_f32(acc1, w, vld1q_f32(v_row.add(4)));
acc2 = vfmaq_f32(acc2, w, vld1q_f32(v_row.add(8)));
acc3 = vfmaq_f32(acc3, w, vld1q_f32(v_row.add(12)));
}
vst1q_f32(out.add(d), acc0);
vst1q_f32(out.add(d + 4), acc1);
vst1q_f32(out.add(d + 8), acc2);
vst1q_f32(out.add(d + 12), acc3);
d += 16;
}
}
#[inline(always)]
fn dot_f32_dispatch(q: &[f32], q_off: usize, k: &[f32], k_off: usize, head_dim: usize) -> f32 {
#[cfg(target_arch = "aarch64")]
if head_dim == 128 {
return unsafe { dot_f32_neon_128(q.as_ptr().add(q_off), k.as_ptr().add(k_off)) };
}
let mut dot = 0.0f32;
for d in 0..head_dim {
dot += q[q_off + d] * k[k_off + d];
}
dot
}
#[inline(always)]
fn accum_v_dispatch(
out: &mut [f32],
out_off: usize,
scores: &[f32],
score_off: usize,
v: &[f32],
v_base_off: usize,
kv_len: usize,
kv_row_stride: usize,
head_dim: usize,
) {
#[cfg(target_arch = "aarch64")]
if head_dim == 128 {
unsafe {
accum_v_neon_128(
out.as_mut_ptr().add(out_off),
scores.as_ptr().add(score_off),
v.as_ptr().add(v_base_off),
kv_len,
kv_row_stride,
);
}
return;
}
for ki in 0..kv_len {
let w = scores[score_off + ki];
let v_off = v_base_off + ki * kv_row_stride;
for d in 0..head_dim {
out[out_off + d] += w * v[v_off + d];
}
}
}
fn compute_attention(
output: &mut [f32],
q: &[f32],
k: &[f32],
v: &[f32],
q_seq_len: usize,
kv_seq_len: usize,
start_pos: usize,
cfg: &GqaConfig,
scores_scratch: &mut [f32],
score_stride: usize,
) {
let head_dim = cfg.head_dim;
let num_heads = cfg.num_heads;
let num_kv_heads = cfg.num_kv_heads;
let groups = num_heads / num_kv_heads;
let scale = 1.0 / (head_dim as f32).sqrt();
debug_assert!(kv_seq_len <= score_stride);
debug_assert!(scores_scratch.len() >= num_heads * score_stride);
if q_seq_len == 1 {
output[..num_heads * head_dim].fill(0.0);
for kv_h in 0..num_kv_heads {
let group_start = kv_h * groups;
for ki in 0..kv_seq_len {
let k_off = ki * (num_kv_heads * head_dim) + kv_h * head_dim;
for gi in 0..groups {
let h = group_start + gi;
let q_off = h * head_dim; scores_scratch[h * score_stride + ki] =
dot_f32_dispatch(q, q_off, k, k_off, head_dim) * scale;
}
}
}
for h in 0..num_heads {
let scores = &mut scores_scratch[h * score_stride..h * score_stride + kv_seq_len];
let mut m = f32::NEG_INFINITY;
let mut l = 0.0f32;
for &s in scores.iter() {
let m_new = m.max(s);
let alpha = if m == f32::NEG_INFINITY {
0.0
} else {
(m - m_new).exp()
};
l = l * alpha + (s - m_new).exp();
m = m_new;
}
if l > 0.0 {
let inv_l = 1.0 / l;
for s in scores.iter_mut() {
*s = (*s - m).exp() * inv_l;
}
} else {
scores.fill(0.0);
}
}
let kv_row_stride = num_kv_heads * head_dim;
for h in 0..num_heads {
let kv_h = h / groups;
accum_v_dispatch(
output,
h * head_dim,
scores_scratch,
h * score_stride,
v,
kv_h * head_dim,
kv_seq_len,
kv_row_stride,
head_dim,
);
}
return;
}
for h in 0..num_heads {
let kv_h = h / groups;
let head_score_base = h * score_stride;
for qi in 0..q_seq_len {
let q_off = qi * (num_heads * head_dim) + h * head_dim;
let scores = &mut scores_scratch[head_score_base..head_score_base + kv_seq_len];
let max_attend = start_pos + qi;
for ki in 0..kv_seq_len {
let k_off = ki * (num_kv_heads * head_dim) + kv_h * head_dim;
if ki > max_attend {
scores[ki] = f32::NEG_INFINITY;
} else {
let mut dot = 0.0f32;
for d in 0..head_dim {
dot += q[q_off + d] * k[k_off + d];
}
scores[ki] = dot * scale;
}
}
let max_score = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for s in scores.iter_mut() {
*s = (*s - max_score).exp();
sum += *s;
}
if sum > 0.0 {
let inv = 1.0 / sum;
for s in scores.iter_mut() {
*s *= inv;
}
}
let out_off = qi * (num_heads * head_dim) + h * head_dim;
for d in 0..head_dim {
let mut val = 0.0f32;
for ki in 0..kv_seq_len {
let v_off = ki * (num_kv_heads * head_dim) + kv_h * head_dim;
val += scores[ki] * v[v_off + d];
}
output[out_off + d] = val;
}
}
}
}
#[cfg(feature = "bench-internals")]
pub mod bench_support {
use super::compute_attention;
use crate::attention::gqa::GqaConfig;
#[inline]
pub fn compute_attention_for_bench(
output: &mut [f32],
q: &[f32],
k: &[f32],
v: &[f32],
q_seq_len: usize,
kv_seq_len: usize,
start_pos: usize,
cfg: &GqaConfig,
scores_scratch: &mut [f32],
score_stride: usize,
) {
compute_attention(
output,
q,
k,
v,
q_seq_len,
kv_seq_len,
start_pos,
cfg,
scores_scratch,
score_stride,
);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generate_config_default() {
let cfg = GenerateConfig::default();
assert_eq!(cfg.max_new_tokens, 256);
assert!(!cfg.include_prompt);
}
#[test]
fn test_generate_output_struct() {
let output = GenerateOutput {
text: "Hello world".into(),
token_ids: vec![1, 2, 3],
prompt_tokens: 5,
generated_tokens: 3,
stopped_by_eos: false,
};
assert_eq!(output.generated_tokens, 3);
assert!(!output.stopped_by_eos);
}
#[test]
fn test_compute_attention_single_token() {
let cfg = GqaConfig {
num_heads: 1,
num_kv_heads: 1,
head_dim: 4,
};
let q = vec![1.0f32, 0.0, 0.0, 0.0];
let k = vec![1.0f32, 0.0, 0.0, 0.0];
let v = vec![0.5f32, 0.5, 0.5, 0.5];
let mut output = vec![0.0f32; 4];
let mut scores = vec![0.0f32; 4];
compute_attention(&mut output, &q, &k, &v, 1, 1, 0, &cfg, &mut scores, 1);
for i in 0..4 {
assert!(
(output[i] - 0.5).abs() < 1e-5,
"output[{i}] = {}",
output[i]
);
}
}
#[test]
fn test_compute_attention_causal_mask() {
let cfg = GqaConfig {
num_heads: 1,
num_kv_heads: 1,
head_dim: 2,
};
let q = vec![1.0f32, 0.0, 0.0, 1.0]; let k = vec![1.0f32, 0.0, 0.0, 1.0]; let v = vec![1.0f32, 0.0, 0.0, 1.0]; let mut output = vec![0.0f32; 4];
let mut scores = vec![0.0f32; 4];
compute_attention(&mut output, &q, &k, &v, 2, 2, 0, &cfg, &mut scores, 2);
assert!((output[0] - 1.0).abs() < 1e-3, "output[0] = {}", output[0]);
assert!((output[1] - 0.0).abs() < 1e-3, "output[1] = {}", output[1]);
}
fn compute_attention_ref(
output: &mut [f32],
q: &[f32],
k: &[f32],
v: &[f32],
q_seq_len: usize,
kv_seq_len: usize,
start_pos: usize,
cfg: &GqaConfig,
) {
let head_dim = cfg.head_dim;
let num_heads = cfg.num_heads;
let num_kv_heads = cfg.num_kv_heads;
let groups = num_heads / num_kv_heads;
let scale = 1.0 / (head_dim as f32).sqrt();
let mut scores_scratch = vec![0.0f32; num_heads * kv_seq_len];
for h in 0..num_heads {
let kv_h = h / groups;
let head_score_base = h * kv_seq_len;
for qi in 0..q_seq_len {
let q_off = qi * (num_heads * head_dim) + h * head_dim;
let scores = &mut scores_scratch[head_score_base..head_score_base + kv_seq_len];
let max_attend = start_pos + qi;
for ki in 0..kv_seq_len {
let k_off = ki * (num_kv_heads * head_dim) + kv_h * head_dim;
if ki > max_attend {
scores[ki] = f32::NEG_INFINITY;
} else {
let mut dot = 0.0f32;
for d in 0..head_dim {
dot += q[q_off + d] * k[k_off + d];
}
scores[ki] = dot * scale;
}
}
let max_score = scores.iter().copied().fold(f32::NEG_INFINITY, f32::max);
let mut sum = 0.0f32;
for s in scores.iter_mut() {
*s = (*s - max_score).exp();
sum += *s;
}
if sum > 0.0 {
let inv = 1.0 / sum;
for s in scores.iter_mut() {
*s *= inv;
}
}
let out_off = qi * (num_heads * head_dim) + h * head_dim;
for d in 0..head_dim {
let mut val = 0.0f32;
for ki in 0..kv_seq_len {
let v_off = ki * (num_kv_heads * head_dim) + kv_h * head_dim;
val += scores[ki] * v[v_off + d];
}
output[out_off + d] = val;
}
}
}
}
#[test]
fn test_compute_attention_gqa_decode_parity() {
let cfg = GqaConfig {
num_heads: 4,
num_kv_heads: 2,
head_dim: 4,
};
let kv_seq_len = 5usize;
let num_heads = cfg.num_heads;
let head_dim = cfg.head_dim;
let q: Vec<f32> = (0..num_heads * head_dim)
.map(|i| (i as f32 + 1.0) * 0.15 - 0.5)
.collect();
let k: Vec<f32> = (0..kv_seq_len * cfg.num_kv_heads * head_dim)
.map(|i| (i as f32) * 0.1 - 0.3)
.collect();
let v: Vec<f32> = (0..kv_seq_len * cfg.num_kv_heads * head_dim)
.map(|i| (i as f32) * 0.07 + 0.05)
.collect();
let start_pos = kv_seq_len - 1;
let mut ref_output = vec![0.0f32; num_heads * head_dim];
compute_attention_ref(&mut ref_output, &q, &k, &v, 1, kv_seq_len, start_pos, &cfg);
let mut fast_output = vec![0.0f32; num_heads * head_dim];
let mut scores = vec![0.0f32; num_heads * kv_seq_len];
compute_attention(
&mut fast_output,
&q,
&k,
&v,
1,
kv_seq_len,
start_pos,
&cfg,
&mut scores,
kv_seq_len,
);
for i in 0..fast_output.len() {
assert!(
(fast_output[i] - ref_output[i]).abs() < 1e-5,
"output[{i}]: fast={} ref={}",
fast_output[i],
ref_output[i]
);
}
}
#[test]
fn test_forward_scratch_capacity() {
let cfg = QwenConfig {
vocab_size: 100,
hidden_size: 64,
num_hidden_layers: 2,
num_attention_heads: 4,
num_key_value_heads: 2,
head_dim: 16,
intermediate_size: 128,
max_position_embeddings: 512,
rms_norm_eps: 1e-6,
rope_theta: 10_000.0,
};
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(&cfg, 8, 64);
assert!(scratch.hidden.len() >= 8 * 64);
assert!(scratch.logits.len() >= 100);
assert!(scratch.scores.len() >= 4 * 64);
let old_len = scratch.hidden.len();
scratch.ensure_capacity(&cfg, 1, 64);
assert_eq!(scratch.hidden.len(), old_len);
}
}