use crate::attention::gdn::{GatedDeltaNetState, sigmoid, softplus};
use crate::attention::gdn_fused::{
GatedDeltaNetFusedScratch, conv1d_silu_fused, simd_decay_and_rank1_update, simd_gated_rms_norm,
simd_l2_normalize, simd_matvec_transpose,
};
use crate::forward::cpu::{elementwise_mul, matmul_bt, silu_inplace};
use crate::model::qwen35::Qwen35Model;
use crate::model::qwen35::{
ForwardScratch, GenerationEntryContract, GenerationPlan, GenerationPreparation, KvCache,
decode_tokens, prepare_generation, qwen35_rms_norm, resize, sample_token, should_stop_token,
};
use crate::model::qwen35_config::{GenerateConfig, GenerateOutput, Qwen35Config};
use crate::rope::RopeTable;
use crate::stop_reason::StopReason;
use crate::tokenizer::bpe::BpeTokenizer;
use crate::weights::q8_weights::{
Q8AttentionWeights, Q8CommonLayerWeights, Q8FullAttentionLayerWeights, Q8GatedDeltaNetWeights,
Q8ModelWeights, matmul_bt_q8,
};
pub fn quantize_from_model(
model: &Qwen35Model,
) -> Result<Q8ModelWeights, crate::error::InferenceError> {
let cfg = model.config.clone();
crate::weights::q8_weights::quantize_model_weights(&model.weights, &cfg)
}
#[inline]
pub fn gated_delta_net_step_fused_q8(
input: &[f32],
state: &mut GatedDeltaNetState,
weights: &Q8GatedDeltaNetWeights,
cfg: &Qwen35Config,
scratch: &mut GatedDeltaNetFusedScratch,
output: &mut [f32],
) {
let hidden = cfg.hidden_size;
let num_heads = cfg.linear_num_key_heads;
let value_heads = cfg.linear_num_value_heads();
let ratio = value_heads / num_heads;
let key_dim = cfg.linear_key_head_dim;
let value_dim = cfg.linear_value_head_dim;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let kernel_size = cfg.linear_conv_kernel_dim;
debug_assert_eq!(
value_heads % num_heads,
0,
"value_heads must be divisible by key_heads"
);
debug_assert!(input.len() >= hidden);
debug_assert!(output.len() >= hidden);
scratch.ensure_capacity(qkv_dim, output_dim, value_heads, key_dim, value_dim);
matmul_bt_q8(
input,
&weights.in_proj_qkv,
&mut scratch.qkv_proj[..qkv_dim],
1,
hidden,
qkv_dim,
);
matmul_bt_q8(
input,
&weights.in_proj_z,
&mut scratch.z_proj[..output_dim],
1,
hidden,
output_dim,
);
matmul_bt_q8(
input,
&weights.in_proj_b,
&mut scratch.beta_proj[..value_heads],
1,
hidden,
value_heads,
);
matmul_bt_q8(
input,
&weights.in_proj_a,
&mut scratch.alpha_proj[..value_heads],
1,
hidden,
value_heads,
);
for b in &mut scratch.beta_proj[..value_heads] {
*b = sigmoid(*b);
}
conv1d_silu_fused(
&scratch.qkv_proj[..qkv_dim],
&mut state.conv_buffer,
&weights.conv1d_weight,
&mut scratch.conv_output[..qkv_dim],
qkv_dim,
kernel_size,
);
let q_total = num_heads * key_dim;
let k_total = num_heads * key_dim;
let v_offset = q_total + k_total;
let scale = 1.0 / (key_dim as f32).sqrt();
for h in 0..value_heads {
let k_head = h / ratio;
let q_start = k_head * key_dim;
let k_start = q_total + k_head * key_dim;
let v_start = v_offset + h * value_dim;
scratch.q_head[..key_dim].copy_from_slice(&scratch.conv_output[q_start..q_start + key_dim]);
scratch.k_head[..key_dim].copy_from_slice(&scratch.conv_output[k_start..k_start + key_dim]);
let v = &scratch.conv_output[v_start..v_start + value_dim];
simd_l2_normalize(&mut scratch.q_head[..key_dim]);
simd_l2_normalize(&mut scratch.k_head[..key_dim]);
let a = weights.a_log[h].exp().min(f32::MAX);
let sp = softplus(scratch.alpha_proj[h] + weights.dt_bias[h]);
let g = (-a * sp).exp();
let s_offset = h * key_dim * value_dim;
let s = &mut state.s_matrices[s_offset..s_offset + key_dim * value_dim];
simd_matvec_transpose(
s,
&scratch.k_head[..key_dim],
&mut scratch.kv_mem[..value_dim],
key_dim,
value_dim,
);
let beta_h = scratch.beta_proj[h];
for ((d, &vj), &mem) in scratch.delta[..value_dim]
.iter_mut()
.zip(&v[..value_dim])
.zip(&scratch.kv_mem[..value_dim])
{
*d = (vj - mem * g) * beta_h;
}
simd_decay_and_rank1_update(
s,
&scratch.k_head[..key_dim],
&scratch.delta[..value_dim],
g,
key_dim,
value_dim,
);
let out_start = h * value_dim;
let out_head = &mut scratch.output_heads[out_start..out_start + value_dim];
simd_matvec_transpose(s, &scratch.q_head[..key_dim], out_head, key_dim, value_dim);
for val in out_head.iter_mut() {
*val *= scale;
}
}
let gamma = &weights.norm_weight[..value_dim];
debug_assert_eq!(gamma.len(), value_dim);
for h in 0..value_heads {
let start = h * value_dim;
let end = start + value_dim;
simd_gated_rms_norm(
&scratch.output_heads[start..end],
&scratch.z_proj[start..end],
gamma,
&mut scratch.gated_norm_buf[start..end],
cfg.rms_norm_eps,
);
}
matmul_bt_q8(
&scratch.gated_norm_buf[..output_dim],
&weights.out_proj,
&mut output[..hidden],
1,
output_dim,
hidden,
);
}
fn full_attention_step_q8(
weights: &Q8FullAttentionLayerWeights,
cache_idx: usize,
position: usize,
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
cfg: &Qwen35Config,
rope: &RopeTable,
hidden: usize,
) {
{
let ForwardScratch {
attn_out,
input_tmp,
..
} = scratch;
let src = &attn_out[..hidden];
input_tmp[..hidden].copy_from_slice(src);
}
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let head_dim = cfg.head_dim;
let num_q_heads = cfg.num_attention_heads;
let num_kv_heads = cfg.num_key_value_heads;
let rope_dim = cfg.rope_dim();
let q_proj_dim = 2 * q_dim;
{
let ForwardScratch {
input_tmp,
q_and_gate,
..
} = scratch;
matmul_bt_q8(
&input_tmp[..hidden],
&weights.q_proj,
&mut q_and_gate[..q_proj_dim],
1,
hidden,
q_proj_dim,
);
}
scratch.split_q_and_gate(num_q_heads, head_dim);
{
let ForwardScratch {
input_tmp, k_buf, ..
} = scratch;
matmul_bt_q8(
&input_tmp[..hidden],
&weights.k_proj,
&mut k_buf[..kv_dim],
1,
hidden,
kv_dim,
);
}
{
let ForwardScratch {
input_tmp, v_buf, ..
} = scratch;
matmul_bt_q8(
&input_tmp[..hidden],
&weights.v_proj,
&mut v_buf[..kv_dim],
1,
hidden,
kv_dim,
);
}
for h in 0..num_q_heads {
let start = h * head_dim;
qwen35_rms_norm(
&mut scratch.q_buf[start..start + head_dim],
&weights.q_norm,
head_dim,
cfg.rms_norm_eps,
);
}
for h in 0..num_kv_heads {
let start = h * head_dim;
qwen35_rms_norm(
&mut scratch.k_buf[start..start + head_dim],
&weights.k_norm,
head_dim,
cfg.rms_norm_eps,
);
}
let half = rope_dim / 2;
for h in 0..num_q_heads {
let start = h * head_dim;
let base = position * half;
for i in 0..half {
let cos_val = rope.cos_at(base + i);
let sin_val = rope.sin_at(base + i);
let x0 = scratch.q_buf[start + i];
let x1 = scratch.q_buf[start + half + i];
scratch.q_buf[start + i] = x0 * cos_val - x1 * sin_val;
scratch.q_buf[start + half + i] = x0 * sin_val + x1 * cos_val;
}
}
for h in 0..num_kv_heads {
let start = h * head_dim;
let base = position * half;
for i in 0..half {
let cos_val = rope.cos_at(base + i);
let sin_val = rope.sin_at(base + i);
let x0 = scratch.k_buf[start + i];
let x1 = scratch.k_buf[start + half + i];
scratch.k_buf[start + i] = x0 * cos_val - x1 * sin_val;
scratch.k_buf[start + half + i] = x0 * sin_val + x1 * cos_val;
}
}
kv_cache.append_kv(
cache_idx,
&scratch.k_buf[..kv_dim],
&scratch.v_buf[..kv_dim],
);
let cur_seq_len = kv_cache.seq_len + 1;
let groups = num_q_heads / num_kv_heads;
let scale = 1.0 / (head_dim as f32).sqrt();
let k_cache = &kv_cache.k[cache_idx];
let v_cache = &kv_cache.v[cache_idx];
for qh in 0..num_q_heads {
let kvh = qh / groups;
let q_off = qh * head_dim;
let q = &scratch.q_buf[q_off..q_off + head_dim];
let scores_start = qh * cur_seq_len;
let mut max_score = f32::NEG_INFINITY;
for t in 0..cur_seq_len {
let k_off = t * kv_dim + kvh * head_dim;
let mut dot = 0.0f32;
for d in 0..head_dim {
dot += q[d] * k_cache[k_off + d];
}
let s = dot * scale;
scratch.scores[scores_start + t] = s;
if s > max_score {
max_score = s;
}
}
let mut sum_exp = 0.0f32;
for t in 0..cur_seq_len {
let e = (scratch.scores[scores_start + t] - max_score).exp();
scratch.scores[scores_start + t] = e;
sum_exp += e;
}
crate::attention::softmax_row::finalize_row(
&mut scratch.scores[scores_start..scores_start + cur_seq_len],
sum_exp,
);
let ctx_off = qh * head_dim;
for d in 0..head_dim {
let mut sum = 0.0f32;
for t in 0..cur_seq_len {
let v_off = t * kv_dim + kvh * head_dim;
sum += scratch.scores[scores_start + t] * v_cache[v_off + d];
}
scratch.context[ctx_off + d] = sum;
}
}
{
let ForwardScratch {
context, gate_z, ..
} = scratch;
for (ctx, &gz) in context[..q_dim].iter_mut().zip(&gate_z[..q_dim]) {
let sig = 1.0 / (1.0 + (-gz).exp());
*ctx *= sig;
}
}
matmul_bt_q8(
&scratch.context[..q_dim],
&weights.o_proj,
&mut scratch.attn_out[..hidden],
1,
q_dim,
hidden,
);
}
#[inline]
fn ffn_step_q8(
common: &Q8CommonLayerWeights,
scratch: &mut ForwardScratch,
cfg: &Qwen35Config,
hidden: usize,
) {
let inter = cfg.intermediate_size;
{
let ForwardScratch {
ffn_out, input_tmp, ..
} = scratch;
let src = &ffn_out[..hidden];
input_tmp[..hidden].copy_from_slice(src);
}
{
let ForwardScratch {
input_tmp,
gate_buf,
..
} = scratch;
matmul_bt_q8(
&input_tmp[..hidden],
&common.gate_proj,
&mut gate_buf[..inter],
1,
hidden,
inter,
);
}
{
let ForwardScratch {
input_tmp, up_buf, ..
} = scratch;
matmul_bt_q8(
&input_tmp[..hidden],
&common.up_proj,
&mut up_buf[..inter],
1,
hidden,
inter,
);
}
silu_inplace(&mut scratch.gate_buf[..inter]);
elementwise_mul(&mut scratch.gate_buf[..inter], &scratch.up_buf[..inter]);
matmul_bt_q8(
&scratch.gate_buf[..inter],
&common.down_proj,
&mut scratch.ffn_out[..hidden],
1,
inter,
hidden,
);
}
pub(crate) fn forward_step_q8(
weights: &Q8ModelWeights,
cfg: &Qwen35Config,
rope: &RopeTable,
token_id: u32,
position: usize,
gdn_states: &mut [GatedDeltaNetState],
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
) {
let hidden = cfg.hidden_size;
scratch.ensure_capacity(cfg, kv_cache.seq_len + 1);
let embed_start = token_id as usize * hidden;
scratch.hidden[..hidden]
.copy_from_slice(&weights.embed_tokens[embed_start..embed_start + hidden]);
let mut linear_idx = 0usize;
let mut full_idx = 0usize;
for layer_i in 0..cfg.num_hidden_layers {
let (attn_weights, common) = &weights.layers[layer_i];
scratch.residual[..hidden].copy_from_slice(&scratch.hidden[..hidden]);
qwen35_rms_norm(
&mut scratch.hidden[..hidden],
&common.input_layernorm,
hidden,
cfg.rms_norm_eps,
);
match attn_weights {
Q8AttentionWeights::Linear(gdn_w) => {
gated_delta_net_step_fused_q8(
&scratch.hidden[..hidden],
&mut gdn_states[linear_idx],
gdn_w,
cfg,
&mut scratch.gdn_scratch,
&mut scratch.attn_out[..hidden],
);
linear_idx += 1;
}
Q8AttentionWeights::Full(full_w) => {
scratch.attn_out[..hidden].copy_from_slice(&scratch.hidden[..hidden]);
full_attention_step_q8(
full_w,
cache_idx_of(full_idx),
position,
kv_cache,
scratch,
cfg,
rope,
hidden,
);
full_idx += 1;
}
}
for i in 0..hidden {
scratch.hidden[i] = scratch.residual[i] + scratch.attn_out[i];
}
scratch.residual[..hidden].copy_from_slice(&scratch.hidden[..hidden]);
qwen35_rms_norm(
&mut scratch.hidden[..hidden],
&common.post_attention_layernorm,
hidden,
cfg.rms_norm_eps,
);
scratch.ffn_out[..hidden].copy_from_slice(&scratch.hidden[..hidden]);
ffn_step_q8(common, scratch, cfg, hidden);
for i in 0..hidden {
scratch.hidden[i] = scratch.residual[i] + scratch.ffn_out[i];
}
}
qwen35_rms_norm(
&mut scratch.hidden[..hidden],
&weights.final_norm,
hidden,
cfg.rms_norm_eps,
);
resize(&mut scratch.logits, cfg.vocab_size);
matmul_bt(
&scratch.hidden[..hidden],
&weights.embed_tokens,
&mut scratch.logits[..cfg.vocab_size],
1,
hidden,
cfg.vocab_size,
);
}
#[inline(always)]
fn cache_idx_of(full_idx: usize) -> usize {
full_idx
}
pub fn generate_q8(
weights: &Q8ModelWeights,
cfg: &Qwen35Config,
tokenizer: &BpeTokenizer,
rope: &RopeTable,
prompt: &str,
gen_cfg: &GenerateConfig,
) -> Result<GenerateOutput, crate::error::InferenceError> {
let plan = match prepare_generation(
tokenizer,
prompt,
gen_cfg,
cfg.vocab_size,
rope.max_positions(),
GenerationEntryContract::StandaloneCpu,
)? {
GenerationPreparation::Ready(plan) => plan,
GenerationPreparation::Complete(output) => return Ok(output),
};
let GenerationPlan {
mut rng_state,
prompt_ids,
prompt_len,
..
} = plan;
let num_linear = cfg.num_linear_attention_layers();
let num_full = cfg.num_full_attention_layers();
let mut gdn_states: Vec<GatedDeltaNetState> = (0..num_linear)
.map(|_| GatedDeltaNetState::new(cfg))
.collect();
let mut kv_cache = KvCache::new(num_full);
let mut scratch = ForwardScratch::new();
let mut generated_ids: Vec<u32> = Vec::with_capacity(gen_cfg.max_new_tokens);
let mut all_ids = prompt_ids.clone();
for (pos, &token_id) in prompt_ids.iter().enumerate() {
forward_step_q8(
weights,
cfg,
rope,
token_id,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
if pos < prompt_len - 1 {
kv_cache.seq_len += 1;
}
}
kv_cache.seq_len = prompt_len;
let next_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
&all_ids,
&mut rng_state,
);
if should_stop_token(cfg, gen_cfg, next_id) {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: true,
stop_reason: Some(StopReason::Eos),
token_logprobs: vec![],
});
}
generated_ids.push(next_id);
all_ids.push(next_id);
let mut stopped = false;
let mut stop_reason = StopReason::Length;
for _ in 1..gen_cfg.max_new_tokens {
let pos = kv_cache.seq_len;
let Some(&last_token) = all_ids.last() else {
return Err(crate::error::InferenceError::Inference(
"empty generation state".into(),
));
};
forward_step_q8(
weights,
cfg,
rope,
last_token,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len += 1;
let next_id = sample_token(
&scratch.logits[..cfg.vocab_size],
gen_cfg,
&all_ids,
&mut rng_state,
);
if should_stop_token(cfg, gen_cfg, next_id) {
stopped = true;
stop_reason = StopReason::Eos;
break;
}
generated_ids.push(next_id);
all_ids.push(next_id);
}
let text = decode_tokens(tokenizer, &generated_ids);
Ok(GenerateOutput {
text,
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
token_logprobs: vec![],
stop_reason: Some(stop_reason),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_full_attn_step_q8_rope_stride_half_parity() {
use crate::model::qwen35_config::LayerType;
use crate::weights::q8_weights::quantize_matrix;
let head_dim: usize = 32;
let num_q_heads: usize = 1;
let num_kv_heads: usize = 1;
let hidden: usize = 64;
let q_dim = num_q_heads * head_dim;
let kv_dim = num_kv_heads * head_dim;
let position: usize = 3;
let cfg = Qwen35Config {
hidden_size: hidden,
num_hidden_layers: 2,
vocab_size: 128,
intermediate_size: 128,
rms_norm_eps: 1e-6,
num_attention_heads: num_q_heads,
num_key_value_heads: num_kv_heads,
head_dim,
rope_theta: 10_000.0,
partial_rotary_factor: 0.5, rope_parameters: None,
linear_num_key_heads: 2,
linear_num_value_heads: Some(2),
linear_key_head_dim: 32,
linear_value_head_dim: 32,
linear_conv_kernel_dim: 4,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval: 2,
layer_types: vec![LayerType::LinearAttention, LayerType::FullAttention],
layer_mask: vec![true; 2],
eos_token_id: 127,
max_position_embeddings: 512,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
vision_config: None,
image_token_id: None,
video_token_id: None,
vision_start_token_id: None,
vision_end_token_id: None,
};
let rope_dim = cfg.rope_dim(); let half = rope_dim / 2; let rope = RopeTable::new(rope_dim, 512, cfg.rope_theta);
let mut k_proj_f32 = vec![0.0f32; kv_dim * hidden];
for j in 0..kv_dim {
k_proj_f32[j * hidden + j] = 1.0;
}
let zero_q8 = |rows: usize, cols: usize| -> crate::weights::q8_weights::Q8Matrix {
quantize_matrix(&vec![0.0f32; rows * cols], rows, cols).unwrap()
};
let mut q_proj_f32 = vec![0.0f32; 2 * q_dim * hidden];
for j in 0..q_dim {
q_proj_f32[j * hidden + j] = 1.0;
}
let weights = Q8FullAttentionLayerWeights {
q_proj: quantize_matrix(&q_proj_f32, 2 * q_dim, hidden).unwrap(),
k_proj: quantize_matrix(&k_proj_f32, kv_dim, hidden).unwrap(),
v_proj: zero_q8(kv_dim, hidden),
o_proj: zero_q8(hidden, q_dim),
q_norm: vec![0.0f32; head_dim],
k_norm: vec![0.0f32; head_dim],
};
let input: Vec<f32> = (0..hidden).map(|i| (i as f32 + 1.0) * 0.07).collect();
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(&cfg, 2);
scratch.attn_out[..hidden].copy_from_slice(&input);
let mut kv_cache = KvCache::new(1);
full_attention_step_q8(
&weights,
0,
position,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
);
let mut k_ref = vec![0.0f32; kv_dim];
matmul_bt_q8(&input, &weights.k_proj, &mut k_ref, 1, hidden, kv_dim);
qwen35_rms_norm(&mut k_ref, &weights.k_norm, head_dim, cfg.rms_norm_eps);
let base = position * half;
for i in 0..half {
let cos_val = rope.cos_at(base + i);
let sin_val = rope.sin_at(base + i);
let x0 = k_ref[i];
let x1 = k_ref[half + i];
k_ref[i] = x0 * cos_val - x1 * sin_val;
k_ref[half + i] = x0 * sin_val + x1 * cos_val;
}
let k_cached = &kv_cache.k[0][..kv_dim];
let max_k_diff = k_cached
.iter()
.zip(k_ref.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_k_diff < 1e-4,
"cpu Q8 K-loop stride-half RoPE diverges from reference: max_k_diff = {max_k_diff:.6}. \
With interleaved pairing the diff is O(0.1-1). Bug: #392."
);
let mut q_and_gate_ref = vec![0.0f32; 2 * q_dim];
matmul_bt_q8(
&input,
&weights.q_proj,
&mut q_and_gate_ref,
1,
hidden,
2 * q_dim,
);
let mut q_ref = q_and_gate_ref[..q_dim].to_vec();
qwen35_rms_norm(&mut q_ref, &weights.q_norm, head_dim, cfg.rms_norm_eps);
for i in 0..half {
let cos_val = rope.cos_at(base + i);
let sin_val = rope.sin_at(base + i);
let x0 = q_ref[i];
let x1 = q_ref[half + i];
q_ref[i] = x0 * cos_val - x1 * sin_val;
q_ref[half + i] = x0 * sin_val + x1 * cos_val;
}
let max_q_diff = scratch.q_buf[..q_dim]
.iter()
.zip(q_ref.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_q_diff < 1e-4,
"cpu Q8 Q-loop stride-half RoPE diverges from reference: max_q_diff = {max_q_diff:.6}. \
With interleaved pairing the diff is O(0.1-1). Bug: #392."
);
}
#[test]
#[allow(clippy::type_complexity)]
fn test_q8_forward_compiles() {
let cfg = Qwen35Config::qwen35_2b();
let _fn_ptr: fn(
&Q8ModelWeights,
&Qwen35Config,
&RopeTable,
u32,
usize,
&mut [GatedDeltaNetState],
&mut KvCache,
&mut ForwardScratch,
) = forward_step_q8;
let _gdn_fn_ptr: fn(
&[f32],
&mut GatedDeltaNetState,
&Q8GatedDeltaNetWeights,
&Qwen35Config,
&mut GatedDeltaNetFusedScratch,
&mut [f32],
) = gated_delta_net_step_fused_q8;
let _gen_fn_ptr: fn(
&Q8ModelWeights,
&Qwen35Config,
&BpeTokenizer,
&RopeTable,
&str,
&GenerateConfig,
) -> Result<GenerateOutput, crate::error::InferenceError> = generate_q8;
assert!(cfg.num_full_attention_layers() > 0);
assert!(cfg.num_linear_attention_layers() > 0);
assert_eq!(
cfg.num_full_attention_layers() + cfg.num_linear_attention_layers(),
cfg.num_hidden_layers
);
}
#[test]
fn test_gdn_q8_step_with_zeros() {
let cfg = Qwen35Config::qwen35_2b();
let hidden = cfg.hidden_size;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let num_heads = cfg.linear_num_key_heads;
let kernel_size = cfg.linear_conv_kernel_dim;
let make_zero_q8 = |rows: usize, cols: usize| -> crate::weights::q8_weights::Q8Matrix {
crate::weights::q8_weights::Q8Matrix {
data: vec![0i8; rows * cols],
scales: vec![1.0f32; rows],
rows,
cols,
}
};
let weights = Q8GatedDeltaNetWeights {
in_proj_qkv: make_zero_q8(qkv_dim, hidden),
in_proj_z: make_zero_q8(output_dim, hidden),
in_proj_b: make_zero_q8(num_heads, hidden),
in_proj_a: make_zero_q8(num_heads, hidden),
a_log: vec![0.0f32; num_heads],
dt_bias: vec![0.0f32; num_heads],
conv1d_weight: vec![0.0f32; qkv_dim * kernel_size],
conv_dim: qkv_dim,
kernel_size,
norm_weight: vec![0.0f32; cfg.linear_value_head_dim],
out_proj: make_zero_q8(hidden, output_dim),
};
let mut state = GatedDeltaNetState::new(&cfg);
let mut scratch = GatedDeltaNetFusedScratch::default();
let input = vec![0.0f32; hidden];
let mut output = vec![0.0f32; hidden];
gated_delta_net_step_fused_q8(
&input,
&mut state,
&weights,
&cfg,
&mut scratch,
&mut output,
);
for &v in &output[..hidden] {
assert_eq!(
v, 0.0,
"zero weights + zero input should produce zero output"
);
}
}
fn tiny_full_attn_fixture() -> (Qwen35Config, RopeTable, Q8FullAttentionLayerWeights, usize) {
use crate::model::qwen35_config::LayerType;
use crate::weights::q8_weights::quantize_matrix;
let head_dim: usize = 32;
let num_q_heads: usize = 1;
let num_kv_heads: usize = 1;
let hidden: usize = 64;
let q_dim = num_q_heads * head_dim;
let kv_dim = num_kv_heads * head_dim;
let cfg = Qwen35Config {
hidden_size: hidden,
num_hidden_layers: 2,
vocab_size: 128,
intermediate_size: 128,
rms_norm_eps: 1e-6,
num_attention_heads: num_q_heads,
num_key_value_heads: num_kv_heads,
head_dim,
rope_theta: 10_000.0,
partial_rotary_factor: 0.5,
rope_parameters: None,
linear_num_key_heads: 2,
linear_num_value_heads: Some(2),
linear_key_head_dim: 32,
linear_value_head_dim: 32,
linear_conv_kernel_dim: 4,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval: 2,
layer_types: vec![LayerType::LinearAttention, LayerType::FullAttention],
layer_mask: vec![true; 2],
eos_token_id: 127,
max_position_embeddings: 512,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
vision_config: None,
image_token_id: None,
video_token_id: None,
vision_start_token_id: None,
vision_end_token_id: None,
};
let rope = RopeTable::new(cfg.rope_dim(), 512, cfg.rope_theta);
let mut k_proj_f32 = vec![0.0f32; kv_dim * hidden];
for j in 0..kv_dim {
k_proj_f32[j * hidden + j] = 1.0;
}
let zero_q8 = |rows: usize, cols: usize| -> crate::weights::q8_weights::Q8Matrix {
quantize_matrix(&vec![0.0f32; rows * cols], rows, cols).unwrap()
};
let mut q_proj_f32 = vec![0.0f32; 2 * q_dim * hidden];
for j in 0..q_dim {
q_proj_f32[j * hidden + j] = 1.0;
}
let weights = Q8FullAttentionLayerWeights {
q_proj: quantize_matrix(&q_proj_f32, 2 * q_dim, hidden).unwrap(),
k_proj: quantize_matrix(&k_proj_f32, kv_dim, hidden).unwrap(),
v_proj: zero_q8(kv_dim, hidden),
o_proj: zero_q8(hidden, q_dim),
q_norm: vec![0.0f32; head_dim],
k_norm: vec![0.0f32; head_dim],
};
(cfg, rope, weights, hidden)
}
#[test]
fn test_full_attn_step_q8_nan_score_fails_closed() {
let (cfg, rope, mut weights, hidden) = tiny_full_attn_fixture();
let q_dim = cfg.full_q_dim();
for d in weights.q_proj.data[0..hidden].iter_mut() {
*d = 0;
}
weights.q_proj.scales[0] = f32::INFINITY;
let input: Vec<f32> = (0..hidden).map(|i| (i as f32 + 1.0) * 0.07).collect();
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(&cfg, 2);
scratch.attn_out[..hidden].copy_from_slice(&input);
let mut kv_cache = KvCache::new(1);
full_attention_step_q8(
&weights,
0,
3,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
);
assert!(
scratch.context[..q_dim].iter().all(|v| v.is_finite()),
"Q8 attention must fail closed (zero context), not propagate NaN"
);
}
#[test]
fn test_full_attn_step_q8_clean_row_still_normalizes() {
let (cfg, rope, weights, hidden) = tiny_full_attn_fixture();
let num_heads = cfg.num_attention_heads;
let input: Vec<f32> = (0..hidden).map(|i| (i as f32 + 1.0) * 0.07).collect();
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(&cfg, 2);
scratch.attn_out[..hidden].copy_from_slice(&input);
let mut kv_cache = KvCache::new(1);
full_attention_step_q8(
&weights,
0,
3,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
);
for h in 0..num_heads {
assert_eq!(
scratch.scores[h], 1.0,
"head {h}: clean single-position row must normalize to 1.0"
);
}
}
#[test]
fn test_gdn_q8_decay_gate_overflow_fails_closed() {
let cfg = Qwen35Config::qwen35_2b();
let hidden = cfg.hidden_size;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let num_heads = cfg.linear_num_key_heads;
let kernel_size = cfg.linear_conv_kernel_dim;
let make_zero_q8 = |rows: usize, cols: usize| -> crate::weights::q8_weights::Q8Matrix {
crate::weights::q8_weights::Q8Matrix {
data: vec![0i8; rows * cols],
scales: vec![1.0f32; rows],
rows,
cols,
}
};
let mut a_log = vec![0.0f32; num_heads];
let mut dt_bias = vec![0.0f32; num_heads];
a_log[0] = 100.0; dt_bias[0] = -100.0;
let weights = Q8GatedDeltaNetWeights {
in_proj_qkv: make_zero_q8(qkv_dim, hidden),
in_proj_z: make_zero_q8(output_dim, hidden),
in_proj_b: make_zero_q8(num_heads, hidden),
in_proj_a: make_zero_q8(num_heads, hidden),
a_log,
dt_bias,
conv1d_weight: vec![0.0f32; qkv_dim * kernel_size],
conv_dim: qkv_dim,
kernel_size,
norm_weight: vec![0.0f32; cfg.linear_value_head_dim],
out_proj: make_zero_q8(hidden, output_dim),
};
let mut state = GatedDeltaNetState::new(&cfg);
let mut scratch = GatedDeltaNetFusedScratch::default();
let input = vec![0.0f32; hidden];
let mut output = vec![0.0f32; hidden];
gated_delta_net_step_fused_q8(
&input,
&mut state,
&weights,
&cfg,
&mut scratch,
&mut output,
);
assert!(
state.s_matrices.iter().all(|v| v.is_finite()),
"GDN recurrent state must stay finite when the decay gate overflows"
);
assert!(
output[..hidden].iter().all(|v| v.is_finite()),
"GDN output must stay finite when the decay gate overflows"
);
}
fn zero_layer_q8_fixture() -> (Qwen35Config, Q8ModelWeights, RopeTable, BpeTokenizer) {
use std::collections::HashMap;
let hidden = 4usize;
let vocab = 8usize;
let cfg = Qwen35Config {
hidden_size: hidden,
num_hidden_layers: 0,
vocab_size: vocab,
intermediate_size: 4,
rms_norm_eps: 1e-6,
num_attention_heads: 1,
num_key_value_heads: 1,
head_dim: 4,
rope_theta: 10_000.0,
partial_rotary_factor: 0.5,
rope_parameters: None,
linear_num_key_heads: 1,
linear_num_value_heads: Some(1),
linear_key_head_dim: 4,
linear_value_head_dim: 4,
linear_conv_kernel_dim: 4,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval: 2,
layer_types: vec![],
layer_mask: vec![],
eos_token_id: 5,
max_position_embeddings: 512,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
vision_config: None,
image_token_id: None,
video_token_id: None,
vision_start_token_id: None,
vision_end_token_id: None,
};
let weights = Q8ModelWeights {
embed_tokens: vec![0.0f32; vocab * hidden],
final_norm: vec![0.0f32; hidden],
layers: vec![],
};
let rope = RopeTable::new(2, 64, 10_000.0);
let mut vocab_map: HashMap<String, u32> = HashMap::new();
for (i, c) in ["h", "e", "l", "o", "w", "r", "d", "!"].iter().enumerate() {
vocab_map.insert((*c).to_string(), i as u32);
}
let merges = vec![
("h".to_string(), "e".to_string()),
("he".to_string(), "l".to_string()),
];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab_map, merges).unwrap();
(cfg, weights, rope, tokenizer)
}
#[test]
fn test_generate_q8_max_new_tokens_zero_returns_empty() {
let (cfg, weights, rope, tokenizer) = zero_layer_q8_fixture();
let gen_cfg = GenerateConfig {
max_new_tokens: 0,
..Default::default()
};
let out = generate_q8(&weights, &cfg, &tokenizer, &rope, "h", &gen_cfg)
.expect("max_new_tokens=0 must succeed, not error");
assert_eq!(
out.generated_tokens, 0,
"max_new_tokens=0 must produce zero generated tokens"
);
assert!(
out.token_ids.is_empty(),
"max_new_tokens=0 must produce an empty token list"
);
assert_eq!(out.prompt_tokens, 1, "prompt 'h' tokenizes to one token");
}
#[test]
fn test_generate_q8_honors_stop_token_ids() {
let (cfg, weights, rope, tokenizer) = zero_layer_q8_fixture();
let gen_cfg = GenerateConfig {
max_new_tokens: 4,
stop_token_ids: vec![0], temperature: 0.0, ..Default::default()
};
let out = generate_q8(&weights, &cfg, &tokenizer, &rope, "h", &gen_cfg)
.expect("generate_q8 must succeed with valid stop_token_ids");
assert_eq!(
out.generated_tokens, 0,
"stop token 0 must halt generation before any token is emitted"
);
assert!(
out.stopped,
"stopped flag must be true when a stop token fires"
);
}
#[test]
fn test_generate_q8_rejects_context_overflow() {
use std::collections::HashMap;
let mut vocab: HashMap<String, u32> = HashMap::new();
for (i, c) in ["h", "e", "l", "o"].iter().enumerate() {
vocab.insert((*c).to_string(), i as u32);
}
let merges = vec![
("h".to_string(), "e".to_string()),
("he".to_string(), "l".to_string()),
];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
let cfg = Qwen35Config::qwen35_2b();
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let weights = Q8ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig {
max_new_tokens: usize::MAX,
..Default::default()
};
let err = generate_q8(&weights, &cfg, &tokenizer, &rope, "hello", &gen_cfg)
.expect_err("request beyond context window must error, not panic");
let msg = format!("{err}");
assert!(
msg.contains("context window"),
"error must name the context window; got: {msg}"
);
}
#[test]
fn generate_q8_rejects_empty_prompt() {
use crate::error::InferenceError;
use std::collections::HashMap;
let mut vocab: HashMap<String, u32> = HashMap::new();
for (i, c) in ["h", "e", "l", "o"].iter().enumerate() {
vocab.insert((*c).to_string(), i as u32);
}
let merges = vec![
("h".to_string(), "e".to_string()),
("he".to_string(), "l".to_string()),
];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
let cfg = Qwen35Config::qwen35_2b();
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let weights = Q8ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig::default();
let result = generate_q8(&weights, &cfg, &tokenizer, &rope, "", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::Inference(ref msg)) if msg.contains("empty prompt")),
"generate_q8 must reject an empty prompt with Err(Inference(\"empty \
prompt\")) (#856); got {result:?}"
);
}
#[test]
fn generate_q8_rejects_out_of_vocab_prompt_id() {
use std::collections::HashMap;
let (cfg, weights, rope, _tokenizer) = zero_layer_q8_fixture();
let mut vocab_map: HashMap<String, u32> = HashMap::new();
for (i, c) in ["h", "e", "l", "o", "w", "r", "d", "!"].iter().enumerate() {
vocab_map.insert((*c).to_string(), i as u32);
}
vocab_map.insert("z".to_string(), cfg.vocab_size as u32);
let mismatched_tokenizer = BpeTokenizer::from_vocab_and_merges(vocab_map, vec![])
.expect("tokenizer with an OOV vocab entry still constructs");
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
..Default::default()
};
let err = generate_q8(&weights, &cfg, &mismatched_tokenizer, &rope, "z", &gen_cfg)
.expect_err("an out-of-vocabulary prompt token id must be rejected, not panic");
assert!(
matches!(err, crate::error::InferenceError::InvalidInput(_)),
"expected InvalidInput, got {err:?}"
);
}
#[test]
fn generate_q8_rejects_grammar_config_before_sampling() {
use crate::error::InferenceError;
use crate::grammar::{GrammarEngine, GrammarSpec};
use std::collections::HashMap;
use std::sync::Arc;
let mut vocab: HashMap<String, u32> = HashMap::new();
for (i, c) in ["h", "e", "l", "o"].iter().enumerate() {
vocab.insert((*c).to_string(), i as u32);
}
let merges = vec![
("h".to_string(), "e".to_string()),
("he".to_string(), "l".to_string()),
];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
let cfg = Qwen35Config::qwen35_2b();
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let weights = Q8ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let spec = GrammarSpec::Gbnf("root ::= \"t\" | \"f\"\n".to_string());
let grammar_vocab = vec![b"t".to_vec(), b"f".to_vec()];
let engine =
GrammarEngine::new(&spec, grammar_vocab).expect("trivial grammar must compile");
let gen_cfg = GenerateConfig {
grammar: Some(Arc::new(engine)),
..Default::default()
};
let result = generate_q8(&weights, &cfg, &tokenizer, &rope, "hello", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"generate_q8 must fail closed with InvalidInput when grammar is set (#397/#398); \
got {result:?}"
);
}
#[test]
fn generate_q8_rejects_stop_strings_config_before_sampling() {
use crate::error::InferenceError;
use std::collections::HashMap;
let mut vocab: HashMap<String, u32> = HashMap::new();
for (i, c) in ["h", "e", "l", "o"].iter().enumerate() {
vocab.insert((*c).to_string(), i as u32);
}
let merges = vec![
("h".to_string(), "e".to_string()),
("he".to_string(), "l".to_string()),
];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
let cfg = Qwen35Config::qwen35_2b();
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let weights = Q8ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig {
stop_strings: vec!["</s>".to_string()],
..Default::default()
};
let result = generate_q8(&weights, &cfg, &tokenizer, &rope, "hello", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"generate_q8 must fail closed with InvalidInput when stop_strings is set \
(ADR-080 C3, #783); got {result:?}"
);
}
#[test]
fn generate_q8_rejects_reasoning_budget_config_before_sampling() {
use crate::error::InferenceError;
use std::collections::HashMap;
let mut vocab: HashMap<String, u32> = HashMap::new();
for (i, c) in ["h", "e", "l", "o"].iter().enumerate() {
vocab.insert((*c).to_string(), i as u32);
}
let merges = vec![
("h".to_string(), "e".to_string()),
("he".to_string(), "l".to_string()),
];
let tokenizer = BpeTokenizer::from_vocab_and_merges(vocab, merges).unwrap();
let cfg = Qwen35Config::qwen35_2b();
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let weights = Q8ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig {
reasoning_budget: Some(16),
..Default::default()
};
let result = generate_q8(&weights, &cfg, &tokenizer, &rope, "hello", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"generate_q8 must fail closed with InvalidInput when reasoning_budget is set \
(ADR-080 C3, #783); got {result:?}"
);
}
fn make_nonzero_q8_cpu_test_model() -> (Qwen35Config, Q8ModelWeights, RopeTable) {
use crate::model::qwen35_config::LayerType;
use crate::weights::q8_weights::quantize_matrix;
let hidden: usize = 64;
let vocab: usize = 128;
let inter: usize = 128;
let num_attn_heads: usize = 2;
let num_kv_heads: usize = 1;
let head_dim: usize = 32;
let q_dim = num_attn_heads * head_dim; let kv_dim = num_kv_heads * head_dim; let lin_key_heads: usize = 2;
let lin_val_heads: usize = 2;
let lin_key_dim: usize = 32;
let lin_val_dim: usize = 32;
let lin_qkv_dim = lin_key_heads * lin_key_dim * 2 + lin_val_heads * lin_val_dim; let lin_output_dim = lin_val_heads * lin_val_dim; let kernel_size: usize = 4;
let cfg = Qwen35Config {
hidden_size: hidden,
num_hidden_layers: 2,
vocab_size: vocab,
intermediate_size: inter,
rms_norm_eps: 1e-6,
num_attention_heads: num_attn_heads,
num_key_value_heads: num_kv_heads,
head_dim,
rope_theta: 10_000.0,
partial_rotary_factor: 0.5,
rope_parameters: None,
linear_num_key_heads: lin_key_heads,
linear_num_value_heads: Some(lin_val_heads),
linear_key_head_dim: lin_key_dim,
linear_value_head_dim: lin_val_dim,
linear_conv_kernel_dim: kernel_size,
num_experts: None,
num_experts_per_tok: None,
moe_intermediate_size: None,
shared_expert_intermediate_size: None,
output_router_logits: false,
router_aux_loss_coef: None,
tie_word_embeddings: true,
full_attention_interval: 2,
layer_types: vec![LayerType::LinearAttention, LayerType::FullAttention],
layer_mask: vec![true; 2],
eos_token_id: 127,
max_position_embeddings: 512,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
vision_config: None,
image_token_id: None,
video_token_id: None,
vision_start_token_id: None,
vision_end_token_id: None,
};
let rope_dim = (head_dim as f32 * cfg.partial_rotary_factor) as usize; let rope = RopeTable::new(rope_dim, cfg.max_position_embeddings, cfg.rope_theta);
let mut seed: u64 = 0xdead_beef_cafe_babe;
let mut next_weight = |n: usize, k: usize| -> crate::weights::q8_weights::Q8Matrix {
let floats: Vec<f32> = (0..n * k)
.map(|_| {
seed = seed
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((seed >> 33) as f32 / u32::MAX as f32) * 0.04 - 0.02
})
.collect();
quantize_matrix(&floats, n, k).unwrap()
};
let gdn_w = Q8GatedDeltaNetWeights {
in_proj_qkv: next_weight(lin_qkv_dim, hidden),
in_proj_z: next_weight(lin_output_dim, hidden),
in_proj_b: next_weight(lin_key_heads, hidden),
in_proj_a: next_weight(lin_key_heads, hidden),
a_log: vec![0.0f32; lin_key_heads],
dt_bias: vec![0.0f32; lin_key_heads],
conv1d_weight: vec![0.01f32; lin_qkv_dim * kernel_size],
conv_dim: lin_qkv_dim,
kernel_size,
norm_weight: vec![0.0f32; lin_val_dim],
out_proj: next_weight(hidden, lin_output_dim),
};
let full_w = Q8FullAttentionLayerWeights {
q_proj: next_weight(2 * q_dim, hidden),
k_proj: next_weight(kv_dim, hidden),
v_proj: next_weight(kv_dim, hidden),
o_proj: next_weight(hidden, q_dim),
q_norm: vec![0.0f32; head_dim],
k_norm: vec![0.0f32; head_dim],
};
let make_common = |seed: &mut u64| -> Q8CommonLayerWeights {
let mut nw = |n: usize, k: usize| -> crate::weights::q8_weights::Q8Matrix {
let floats: Vec<f32> = (0..n * k)
.map(|_| {
*seed = seed
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((*seed >> 33) as f32 / u32::MAX as f32) * 0.04 - 0.02
})
.collect();
quantize_matrix(&floats, n, k).unwrap()
};
Q8CommonLayerWeights {
input_layernorm: vec![0.0f32; hidden],
post_attention_layernorm: vec![0.0f32; hidden],
gate_proj: nw(inter, hidden),
up_proj: nw(inter, hidden),
down_proj: nw(hidden, inter),
}
};
let layers = vec![
(Q8AttentionWeights::Linear(gdn_w), make_common(&mut seed)),
(Q8AttentionWeights::Full(full_w), make_common(&mut seed)),
];
let embed_tokens: Vec<f32> = (0..vocab * hidden)
.map(|_| {
seed = seed
.wrapping_mul(6_364_136_223_846_793_005)
.wrapping_add(1_442_695_040_888_963_407);
((seed >> 33) as f32 / u32::MAX as f32) * 0.04 - 0.02
})
.collect();
let weights = Q8ModelWeights {
embed_tokens,
final_norm: vec![0.0f32; hidden],
layers,
};
(cfg, weights, rope)
}
#[test]
fn test_forward_step_q8_decode_survives_dirty_scratch_reuse() {
let (cfg, weights, rope) = make_nonzero_q8_cpu_test_model();
let num_linear = cfg.num_linear_attention_layers();
let num_full = cfg.num_full_attention_layers();
let decode_two_steps = |scratch: &mut ForwardScratch| -> Vec<f32> {
let mut gdn_states: Vec<GatedDeltaNetState> = (0..num_linear)
.map(|_| GatedDeltaNetState::new(&cfg))
.collect();
let mut kv_cache = KvCache::new(num_full);
forward_step_q8(
&weights,
&cfg,
&rope,
7,
0,
&mut gdn_states,
&mut kv_cache,
scratch,
);
kv_cache.seq_len += 1;
forward_step_q8(
&weights,
&cfg,
&rope,
11,
1,
&mut gdn_states,
&mut kv_cache,
scratch,
);
scratch.logits[..16].to_vec()
};
let mut pristine_scratch = ForwardScratch::new();
let pristine = decode_two_steps(&mut pristine_scratch);
let mut dirtied_scratch = ForwardScratch::new();
let _ = decode_two_steps(&mut dirtied_scratch); let dirty = decode_two_steps(&mut dirtied_scratch);
assert_eq!(
pristine, dirty,
"decoding through a previously-used ForwardScratch must produce identical logits \
to a pristine scratch -- a buffer-reuse bug in the Q8 decode alloc sites (#416) \
would leak stale data here"
);
assert!(
pristine.iter().any(|&v| v.abs() > 1e-9),
"all 16 logits are zero -- check weight generation"
);
for (i, &v) in pristine.iter().enumerate() {
assert!(v.is_finite(), "logit[{i}] is not finite: {v}");
}
let expected: [f32; 16] = [
0.6030709, 0.6551996, 0.60194814, 0.64469767, 0.69641805, 0.5887781, 0.62342393,
0.6259914, 0.65434885, 0.74827236, 0.69311845, 0.7374152, 0.6237912, 0.65976304,
0.6478549, 0.68329644,
];
for (i, (&actual, &exp)) in pristine.iter().zip(expected.iter()).enumerate() {
assert!(
(actual - exp).abs() <= 1e-6,
"logit[{i}] mismatch: actual={actual:.8}, expected={exp:.8}"
);
}
}
}