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 crate::stop_reason::StopReason;
use std::sync::Arc;
#[deprecated(
since = "0.5.1",
note = "crate::generate is repository-dead and duplicates the canonical Qwen3.5 decode loop; \
text-generation users should migrate to Qwen35Model::generate / generate_streaming \
with model::GenerateConfig (embedding users of QwenModel are unaffected). \
See the module-level migration guide in crates/inference/src/generate.rs. \
Scheduled for removal in 0.6.0."
)]
#[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>>,
pub kv_cache_capacity: Option<usize>,
}
#[allow(deprecated)] 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>"))
.field("kv_cache_capacity", &self.kv_cache_capacity)
.finish()
}
}
#[allow(deprecated)] impl Default for GenerateConfig {
fn default() -> Self {
Self {
max_new_tokens: 256,
sampling: SamplingConfig::default(),
eos_token_id: None,
include_prompt: false,
grammar: None,
kv_cache_capacity: None,
}
}
}
#[deprecated(
since = "0.5.1",
note = "crate::generate is repository-dead and duplicates the canonical Qwen3.5 decode loop; \
text-generation users should migrate to Qwen35Model::generate / generate_streaming \
with model::qwen35_config::GenerateOutput (embedding users of QwenModel are unaffected). \
See the module-level migration guide in crates/inference/src/generate.rs. \
Scheduled for removal in 0.6.0."
)]
#[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,
pub stop_reason: Option<StopReason>,
}
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);
}
}
fn check_kv_cache_capacity(
requested: Option<usize>,
effective_cap: usize,
prompt_len: usize,
) -> Result<(), InferenceError> {
if requested.is_some() && effective_cap < prompt_len {
return Err(InferenceError::InvalidInput(format!(
"kv_cache_capacity ({effective_cap}) is smaller than the prompt length \
({prompt_len}); the cache must hold at least the prompt"
)));
}
Ok(())
}
fn compute_max_seq(prompt_len: usize, max_new_tokens: usize) -> Result<usize, InferenceError> {
prompt_len.checked_add(max_new_tokens).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"prompt_len ({prompt_len}) + max_new_tokens ({max_new_tokens}) overflows usize"
))
})
}
fn compute_effective_cap(max_seq: usize, kv_cache_capacity: Option<usize>) -> usize {
kv_cache_capacity
.map(|c| c.max(1).min(max_seq))
.unwrap_or(max_seq)
}
fn check_context_window(effective_cap: usize, max_context: usize) -> Result<(), InferenceError> {
if effective_cap > max_context {
return Err(InferenceError::InvalidInput(format!(
"effective context length ({effective_cap}) exceeds \
the model context window ({max_context})"
)));
}
Ok(())
}
fn check_alloc_capacity(
cfg: &QwenConfig,
num_layers: usize,
effective_cap: usize,
) -> Result<(), InferenceError> {
let kv_dim = cfg
.num_key_value_heads
.checked_mul(cfg.head_dim)
.ok_or_else(|| {
InferenceError::InvalidInput("num_key_value_heads * head_dim overflows usize".into())
})?;
let q_dim = cfg
.num_attention_heads
.checked_mul(cfg.head_dim)
.ok_or_else(|| {
InferenceError::InvalidInput("num_attention_heads * head_dim overflows usize".into())
})?;
let qkv_dim = kv_dim
.checked_mul(2)
.and_then(|two_kv| q_dim.checked_add(two_kv))
.ok_or_else(|| {
InferenceError::InvalidInput("q_dim + 2*kv_dim (qkv_dim) overflows usize".into())
})?;
let inter = cfg.intermediate_size;
effective_cap.checked_mul(cfg.hidden_size).ok_or_else(|| {
InferenceError::InvalidInput("effective_cap * hidden_size overflows usize".into())
})?;
effective_cap.checked_mul(qkv_dim).ok_or_else(|| {
InferenceError::InvalidInput("effective_cap * qkv_dim overflows usize".into())
})?;
effective_cap.checked_mul(q_dim).ok_or_else(|| {
InferenceError::InvalidInput("effective_cap * q_dim overflows usize".into())
})?;
effective_cap.checked_mul(kv_dim).ok_or_else(|| {
InferenceError::InvalidInput("effective_cap * kv_dim overflows usize".into())
})?;
inter
.checked_mul(2)
.and_then(|two_inter| effective_cap.checked_mul(two_inter))
.ok_or_else(|| {
InferenceError::InvalidInput(
"effective_cap * 2 * intermediate_size overflows usize".into(),
)
})?;
effective_cap.checked_mul(inter).ok_or_else(|| {
InferenceError::InvalidInput("effective_cap * intermediate_size overflows usize".into())
})?;
let coeff = (|| -> Option<usize> {
let cache = 2usize.checked_mul(num_layers)?.checked_mul(kv_dim)?;
let dequant = 2usize.checked_mul(kv_dim)?;
cache
.checked_add(dequant)?
.checked_add(cfg.num_attention_heads)
})()
.ok_or_else(|| InferenceError::InvalidInput("model dimensions overflow usize".into()))?;
effective_cap.checked_mul(coeff).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"KV cache + scratch for max_seq ({effective_cap}) overflows usize"
))
})?;
Ok(())
}
fn has_finite_logit(logits: &[f32]) -> bool {
logits.iter().any(|&l| l > f32::NEG_INFINITY)
}
#[deprecated(
since = "0.5.1",
note = "crate::generate::generate is repository-dead and duplicates the canonical Qwen3.5 \
decode loop; text-generation users should migrate to Qwen35Model::generate / \
generate_streaming (embedding users of QwenModel are unaffected). \
See the module-level migration guide in crates/inference/src/generate.rs. \
Scheduled for removal in 0.6.0."
)]
#[allow(deprecated)] 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()));
}
if config.max_new_tokens == 0 {
return Ok(GenerateOutput {
text: String::new(),
token_ids: Vec::new(),
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped_by_eos: false,
stop_reason: Some(StopReason::Length),
});
}
let max_seq = compute_max_seq(prompt_len, config.max_new_tokens)?;
let effective_cap = compute_effective_cap(max_seq, config.kv_cache_capacity);
check_kv_cache_capacity(config.kv_cache_capacity, effective_cap, prompt_len)?;
check_context_window(effective_cap, model.rope().max_positions())?;
check_alloc_capacity(cfg, cfg.num_hidden_layers, effective_cap)?;
let cache_cfg = FlatKVCacheConfig::for_qwen3(
cfg.num_hidden_layers,
cfg.num_key_value_heads,
cfg.head_dim,
effective_cap,
);
let mut cache = FlatKVCache::try_new(cache_cfg)?;
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(cfg, prompt_len.max(1), effective_cap);
let mut sampler = Sampler::new(config.sampling.clone());
sampler.seed_history(&prompt_ids[..prompt_len]);
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,
effective_cap,
)?;
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token at step 0; \
no legal first token exists in the current grammar state"
.into(),
));
}
}
let mut generated_ids: Vec<u32> = Vec::with_capacity(config.max_new_tokens.min(effective_cap));
let first_token = sampler.sample(&scratch.logits[..cfg.vocab_size]);
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state)
&& !engine.advance(gs, first_token)
{
return Ok(GenerateOutput {
text: String::new(),
prompt_tokens: prompt_len,
generated_tokens: 0,
token_ids: vec![],
stopped_by_eos: false,
stop_reason: Some(StopReason::Grammar),
});
}
let mut stopped_by_eos = false;
let mut stop_reason = StopReason::Length;
if config.eos_token_id == Some(first_token) {
stopped_by_eos = true;
stop_reason = StopReason::Eos;
} else {
generated_ids.push(first_token);
}
if !stopped_by_eos {
for step in 0..config.max_new_tokens.saturating_sub(1) {
if cache.is_full() {
stop_reason = StopReason::KvFull;
break;
}
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,
effective_cap,
)?;
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state) {
engine.mask_logits(gs, &mut scratch.logits[..cfg.vocab_size]);
if !has_finite_logit(&scratch.logits[..cfg.vocab_size]) {
return Err(InferenceError::InvalidInput(
"grammar constraint blocked every token; \
no legal continuation exists in the current grammar state"
.into(),
));
}
}
let token = sampler.sample(&scratch.logits[..cfg.vocab_size]);
if let (Some(engine), Some(gs)) = (&config.grammar, &mut grammar_state)
&& !engine.advance(gs, token)
{
stop_reason = StopReason::Grammar;
break;
}
if config.eos_token_id == Some(token) {
stopped_by_eos = true;
stop_reason = StopReason::Eos;
break;
}
generated_ids.push(token);
}
}
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,
stop_reason: Some(stop_reason),
})
}
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;
}
for s in scores.iter_mut() {
*s = (*s - m).exp();
}
crate::attention::softmax_row::finalize_row(scores, l);
}
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;
}
crate::attention::softmax_row::finalize_row(scores, sum);
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)]
#[allow(deprecated)] mod tests {
use super::*;
#[test]
fn has_finite_logit_detects_all_neg_infinity() {
assert!(
!has_finite_logit(&[f32::NEG_INFINITY; 8]),
"all-NEG_INFINITY must fail the guard (empty grammar mask)"
);
let mut mixed = vec![f32::NEG_INFINITY; 4];
mixed[2] = 1.0_f32;
assert!(
has_finite_logit(&mixed),
"a single finite logit must pass the guard"
);
assert!(has_finite_logit(&[0.0_f32; 4]));
let mut with_inf = vec![f32::NEG_INFINITY; 4];
with_inf[1] = f32::INFINITY;
assert!(has_finite_logit(&with_inf));
}
#[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,
stop_reason: Some(StopReason::Length),
};
assert_eq!(output.generated_tokens, 3);
assert!(!output.stopped_by_eos);
}
#[test]
fn test_compute_max_seq_normal() {
assert_eq!(compute_max_seq(10, 100).unwrap(), 110);
assert_eq!(compute_max_seq(0, 0).unwrap(), 0);
assert_eq!(compute_max_seq(1, usize::MAX - 1).unwrap(), usize::MAX);
}
#[test]
fn test_compute_max_seq_overflow_is_error_not_panic() {
let err = compute_max_seq(10, usize::MAX).unwrap_err();
assert!(matches!(err, InferenceError::InvalidInput(_)));
let err = compute_max_seq(usize::MAX, 1).unwrap_err();
assert!(matches!(err, InferenceError::InvalidInput(_)));
}
#[test]
fn test_check_context_window_over_capacity_is_error_not_panic() {
let err = check_context_window(usize::MAX, 4096)
.expect_err("over-capacity request must be rejected, not admitted to a panic");
let msg = format!("{err}");
assert!(
msg.contains("context window"),
"error should mention context window, got: {msg}"
);
assert!(matches!(err, InferenceError::InvalidInput(_)));
}
#[test]
fn test_check_context_window_boundary() {
assert!(
check_context_window(4096, 4096).is_ok(),
"effective length == max_context must be accepted"
);
assert!(
check_context_window(4097, 4096).is_err(),
"effective length == max_context + 1 must be rejected"
);
}
#[test]
fn test_check_context_window_respects_kv_cache_cap() {
let prompt_len = 4000usize;
let max_new_tokens = 10_000usize;
let max_context = 4096usize;
let kv_cache_capacity = Some(4096usize);
let max_seq = compute_max_seq(prompt_len, max_new_tokens).expect("no overflow");
let effective_cap = compute_effective_cap(max_seq, kv_cache_capacity);
assert!(
check_context_window(effective_cap, max_context).is_ok(),
"kv_cache_capacity-capped request within the window must be accepted"
);
assert!(
check_context_window(max_seq, max_context).is_err(),
"the raw uncapped sum exceeds the window (guarding it would over-reject)"
);
}
#[test]
fn test_check_alloc_capacity_normal() {
let cfg = QwenConfig::qwen3_embedding_0_6b();
assert!(check_alloc_capacity(&cfg, cfg.num_hidden_layers, 4096).is_ok());
assert!(check_alloc_capacity(&cfg, cfg.num_hidden_layers, 262_144).is_ok());
assert!(check_alloc_capacity(&cfg, cfg.num_hidden_layers, 0).is_ok());
}
#[test]
fn check_alloc_capacity_rejects_kv_dim_overflow() {
let overflow_kv_heads = usize::MAX / 4 + 1;
let cfg = QwenConfig {
vocab_size: 1,
hidden_size: 1,
num_hidden_layers: 1,
num_attention_heads: 1,
num_key_value_heads: overflow_kv_heads,
head_dim: 4,
intermediate_size: 1,
max_position_embeddings: 1,
rms_norm_eps: 1e-6,
rope_theta: 10_000.0,
};
let err = check_alloc_capacity(&cfg, 1, 1).unwrap_err();
assert!(
matches!(err, InferenceError::InvalidInput(_)),
"expected InvalidInput on kv_dim overflow, got {err:?}"
);
}
#[test]
fn check_alloc_capacity_rejects_q_dim_overflow() {
let overflow_q_heads = usize::MAX / 4 + 1;
let cfg = QwenConfig {
vocab_size: 1,
hidden_size: 1,
num_hidden_layers: 1,
num_attention_heads: overflow_q_heads,
num_key_value_heads: 1,
head_dim: 4,
intermediate_size: 1,
max_position_embeddings: 1,
rms_norm_eps: 1e-6,
rope_theta: 10_000.0,
};
let err = check_alloc_capacity(&cfg, 1, 1).unwrap_err();
assert!(
matches!(err, InferenceError::InvalidInput(_)),
"expected InvalidInput on q_dim overflow, got {err:?}"
);
}
#[test]
fn test_check_alloc_capacity_multiplication_overflow_is_error() {
let cfg = QwenConfig::qwen3_embedding_0_6b();
let effective_cap = usize::MAX / 1024;
let err = check_alloc_capacity(&cfg, cfg.num_hidden_layers, effective_cap).unwrap_err();
assert!(matches!(err, InferenceError::InvalidInput(_)));
}
#[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]);
}
#[test]
fn test_compute_attention_prefill_nan_score_fails_closed() {
let cfg = GqaConfig {
num_heads: 1,
num_kv_heads: 1,
head_dim: 1,
};
let q = vec![f32::NAN, 0.0]; let k = vec![1.0f32, 1.0];
let v = vec![7.0f32, 11.0];
let mut output = vec![0.0f32; 2];
let mut scores = vec![0.0f32; 2];
compute_attention(&mut output, &q, &k, &v, 2, 2, 0, &cfg, &mut scores, 2);
assert!(
output[0].is_finite(),
"prefill attention must fail closed on NaN score, got {}",
output[0]
);
assert_eq!(output[0], 0.0, "failed-closed row must be zeroed");
assert!(
(output[1] - 9.0).abs() < 1e-5,
"clean sibling query must normalize to the exact weighted average, got {}",
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);
}
#[test]
fn kv_cache_capacity_below_prompt_len_is_rejected() {
assert!(check_kv_cache_capacity(Some(8), 8, 32).is_err());
assert!(check_kv_cache_capacity(Some(32), 32, 32).is_ok());
assert!(check_kv_cache_capacity(Some(64), 64, 16).is_ok());
assert!(check_kv_cache_capacity(None, 100, 32).is_ok());
}
#[test]
fn kv_cache_capacity_allocates_fewer_bytes() {
use crate::kv_cache::{FlatKVCache, FlatKVCacheConfig};
let num_layers = 28usize;
let num_kv_heads = 8usize;
let head_dim = 128usize;
let cap_small = 128usize;
let cap_large = 4096usize;
let bytes_small =
FlatKVCacheConfig::for_qwen3(num_layers, num_kv_heads, head_dim, cap_small)
.total_bytes();
let bytes_large =
FlatKVCacheConfig::for_qwen3(num_layers, num_kv_heads, head_dim, cap_large)
.total_bytes();
assert!(
bytes_small < bytes_large,
"cap=128 bytes ({bytes_small}) should be less than cap=4096 bytes ({bytes_large})"
);
assert_eq!(
bytes_large / bytes_small,
cap_large / cap_small,
"allocation should scale linearly with cap"
);
let mb_small = bytes_small as f64 / (1024.0 * 1024.0);
let mb_large = bytes_large as f64 / (1024.0 * 1024.0);
assert!(
mb_small < 20.0,
"cap=128 should be under 20 MB, got {mb_small:.1} MB"
);
assert!(
mb_large > 400.0,
"cap=4096 should be over 400 MB, got {mb_large:.1} MB"
);
let cache_small = FlatKVCache::new(FlatKVCacheConfig::for_qwen3(
num_layers,
num_kv_heads,
head_dim,
cap_small,
));
let cache_large = FlatKVCache::new(FlatKVCacheConfig::for_qwen3(
num_layers,
num_kv_heads,
head_dim,
cap_large,
));
assert_eq!(cache_small.memory_bytes(), bytes_small);
assert_eq!(cache_large.memory_bytes(), bytes_large);
}
#[test]
fn kv_cache_capacity_default_is_none() {
let cfg = GenerateConfig::default();
assert!(
cfg.kv_cache_capacity.is_none(),
"default must be None to preserve backward-compatible allocation"
);
}
#[test]
fn kv_cache_capacity_clamp_above_max_seq() {
let prompt_len = 10usize;
let max_new_tokens = 50usize;
let max_seq = compute_max_seq(prompt_len, max_new_tokens).expect("no overflow");
let effective = |cap: Option<usize>| compute_effective_cap(max_seq, cap);
assert_eq!(effective(None), 60, "None -> full max_seq");
assert_eq!(
effective(Some(100)),
60,
"cap > max_seq -> clamped to max_seq"
);
assert_eq!(effective(Some(60)), 60, "cap == max_seq -> unchanged");
assert_eq!(effective(Some(30)), 30, "cap < max_seq -> respected");
assert_eq!(effective(Some(0)), 1, "cap=0 -> clamped to 1");
assert_eq!(effective(Some(1)), 1, "cap=1 -> 1 (minimum)");
}
#[test]
fn check_alloc_capacity_rejects_seq_scaled_qkv_overflow() {
let cfg = QwenConfig {
vocab_size: 1,
hidden_size: 1,
num_hidden_layers: 1,
num_attention_heads: 1000,
num_key_value_heads: 1,
head_dim: 4,
intermediate_size: 1,
max_position_embeddings: 1,
rms_norm_eps: 1e-6,
rope_theta: 10_000.0,
};
let effective_cap = 4_602_481_056_314_759usize;
let err = check_alloc_capacity(&cfg, 1, effective_cap).unwrap_err();
assert!(
matches!(err, InferenceError::InvalidInput(_)),
"expected InvalidInput on seq-scaled qkv overflow, got {err:?}"
);
}
#[test]
fn check_alloc_capacity_accepts_realistic_scratch_dims() {
let cfg = QwenConfig::qwen3_embedding_0_6b();
assert!(
check_alloc_capacity(&cfg, cfg.num_hidden_layers, 4_096).is_ok(),
"4K-token cap rejected"
);
assert!(
check_alloc_capacity(&cfg, cfg.num_hidden_layers, 32_768).is_ok(),
"32K-token cap rejected"
);
}
}