use crate::attention::gdn::{
GatedDeltaNetState, GatedDeltaNetWeights, gated_rms_norm, l2_normalize_vec, sigmoid, softplus,
};
use crate::attention::gdn_fused::GatedDeltaNetFusedScratch;
use crate::forward::cpu::{elementwise_mul, silu_inplace};
use crate::forward::neon::{matmul_q8_neon_into, pack_weights_q8};
use crate::model::qwen35::{
AttentionWeights, CommonLayerWeights, FeedForwardWeights, ForwardScratch,
FullAttentionLayerWeights, KvCache, ModelWeights, decode_tokens, qwen35_rms_norm, resize,
sample_token,
};
use crate::model::qwen35_config::{GenerateConfig, GenerateOutput, Qwen35Config};
use crate::rope::RopeTable;
use crate::tokenizer::bpe::BpeTokenizer;
use crate::tokenizer::common::Tokenizer;
pub struct Q8NeonGdnWeights {
pub in_proj_qkv_packed: Vec<u8>,
pub in_proj_qkv_rows: usize,
pub in_proj_qkv_cols: usize,
pub in_proj_z_packed: Vec<u8>,
pub in_proj_z_rows: usize,
pub in_proj_z_cols: usize,
pub in_proj_b_packed: Vec<u8>,
pub in_proj_b_rows: usize,
pub in_proj_b_cols: usize,
pub in_proj_a_packed: Vec<u8>,
pub in_proj_a_rows: usize,
pub in_proj_a_cols: usize,
pub out_proj_packed: Vec<u8>,
pub out_proj_rows: usize,
pub out_proj_cols: usize,
pub a_log: Vec<f32>,
pub dt_bias: Vec<f32>,
pub conv1d_weight: Vec<f32>,
pub conv_dim: usize,
pub kernel_size: usize,
pub norm_weight: Vec<f32>,
}
pub struct Q8NeonFullAttnWeights {
pub q_proj_packed: Vec<u8>,
pub q_proj_rows: usize,
pub q_proj_cols: usize,
pub k_proj_packed: Vec<u8>,
pub k_proj_rows: usize,
pub k_proj_cols: usize,
pub v_proj_packed: Vec<u8>,
pub v_proj_rows: usize,
pub v_proj_cols: usize,
pub o_proj_packed: Vec<u8>,
pub o_proj_rows: usize,
pub o_proj_cols: usize,
pub q_norm: Vec<f32>,
pub k_norm: Vec<f32>,
}
pub struct Q8NeonCommonWeights {
pub input_layernorm: Vec<f32>,
pub post_attention_layernorm: Vec<f32>,
pub gate_proj_packed: Vec<u8>,
pub gate_proj_rows: usize,
pub gate_proj_cols: usize,
pub up_proj_packed: Vec<u8>,
pub up_proj_rows: usize,
pub up_proj_cols: usize,
pub down_proj_packed: Vec<u8>,
pub down_proj_rows: usize,
pub down_proj_cols: usize,
}
pub enum Q8NeonAttentionWeights {
Linear(Q8NeonGdnWeights),
Full(Q8NeonFullAttnWeights),
}
pub struct Q8NeonModel {
pub embed_tokens: Vec<f32>,
pub final_norm: Vec<f32>,
pub lm_head_packed: Vec<u8>,
pub lm_head_rows: usize,
pub lm_head_cols: usize,
pub layers: Vec<(Q8NeonAttentionWeights, Q8NeonCommonWeights)>,
}
fn pack_gdn_weights(w: &GatedDeltaNetWeights) -> Q8NeonGdnWeights {
Q8NeonGdnWeights {
in_proj_qkv_packed: pack_weights_q8(&w.in_proj_qkv, w.in_proj_qkv_rows, w.in_proj_qkv_cols),
in_proj_qkv_rows: w.in_proj_qkv_rows,
in_proj_qkv_cols: w.in_proj_qkv_cols,
in_proj_z_packed: pack_weights_q8(&w.in_proj_z, w.in_proj_z_rows, w.in_proj_z_cols),
in_proj_z_rows: w.in_proj_z_rows,
in_proj_z_cols: w.in_proj_z_cols,
in_proj_b_packed: pack_weights_q8(&w.in_proj_b, w.in_proj_b_rows, w.in_proj_b_cols),
in_proj_b_rows: w.in_proj_b_rows,
in_proj_b_cols: w.in_proj_b_cols,
in_proj_a_packed: pack_weights_q8(&w.in_proj_a, w.in_proj_a_rows, w.in_proj_a_cols),
in_proj_a_rows: w.in_proj_a_rows,
in_proj_a_cols: w.in_proj_a_cols,
out_proj_packed: pack_weights_q8(&w.out_proj, w.out_proj_rows, w.out_proj_cols),
out_proj_rows: w.out_proj_rows,
out_proj_cols: w.out_proj_cols,
a_log: w.a_log.clone(),
dt_bias: w.dt_bias.clone(),
conv1d_weight: w.conv1d_weight.clone(),
conv_dim: w.conv_dim,
kernel_size: w.kernel_size,
norm_weight: w.norm_weight.clone(),
}
}
fn pack_full_attn_weights(
w: &FullAttentionLayerWeights,
cfg: &Qwen35Config,
) -> Q8NeonFullAttnWeights {
let hidden = cfg.hidden_size;
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let q_proj_rows = 2 * q_dim;
Q8NeonFullAttnWeights {
q_proj_packed: pack_weights_q8(&w.q_proj, q_proj_rows, hidden),
q_proj_rows,
q_proj_cols: hidden,
k_proj_packed: pack_weights_q8(&w.k_proj, kv_dim, hidden),
k_proj_rows: kv_dim,
k_proj_cols: hidden,
v_proj_packed: pack_weights_q8(&w.v_proj, kv_dim, hidden),
v_proj_rows: kv_dim,
v_proj_cols: hidden,
o_proj_packed: pack_weights_q8(&w.o_proj, hidden, q_dim),
o_proj_rows: hidden,
o_proj_cols: q_dim,
q_norm: w.q_norm.clone(),
k_norm: w.k_norm.clone(),
}
}
fn pack_common_weights(w: &CommonLayerWeights, cfg: &Qwen35Config) -> Q8NeonCommonWeights {
let hidden = cfg.hidden_size;
let inter = cfg.intermediate_size;
let (gate_proj, up_proj, down_proj) = match &w.ffn {
FeedForwardWeights::Dense(dense) => (&dense.gate_proj, &dense.up_proj, &dense.down_proj),
FeedForwardWeights::Moe(_) => {
panic!("Q8 NEON packing is dense-only; MoE configs are not supported");
}
};
Q8NeonCommonWeights {
input_layernorm: w.input_layernorm.clone(),
post_attention_layernorm: w.post_attention_layernorm.clone(),
gate_proj_packed: pack_weights_q8(gate_proj, inter, hidden),
gate_proj_rows: inter,
gate_proj_cols: hidden,
up_proj_packed: pack_weights_q8(up_proj, inter, hidden),
up_proj_rows: inter,
up_proj_cols: hidden,
down_proj_packed: pack_weights_q8(down_proj, hidden, inter),
down_proj_rows: hidden,
down_proj_cols: inter,
}
}
pub fn quantize_model(weights: &ModelWeights, cfg: &Qwen35Config) -> Q8NeonModel {
let hidden = cfg.hidden_size;
let vocab = cfg.vocab_size;
let layers = weights
.layers
.iter()
.map(|(attn, common)| {
let q8_attn = match attn {
AttentionWeights::Linear(gdn_w) => {
Q8NeonAttentionWeights::Linear(pack_gdn_weights(gdn_w))
}
AttentionWeights::Full(full_w) => {
Q8NeonAttentionWeights::Full(pack_full_attn_weights(full_w, cfg))
}
};
let q8_common = pack_common_weights(common, cfg);
(q8_attn, q8_common)
})
.collect();
let lm_head_packed = pack_weights_q8(weights.logits_weight(), vocab, hidden);
Q8NeonModel {
embed_tokens: weights.embed_tokens.clone(),
final_norm: weights.final_norm.clone(),
lm_head_packed,
lm_head_rows: vocab,
lm_head_cols: hidden,
layers,
}
}
fn gdn_step_q8_neon(
input: &[f32],
state: &mut GatedDeltaNetState,
weights: &Q8NeonGdnWeights,
cfg: &Qwen35Config,
gdn_scratch: &mut GatedDeltaNetFusedScratch,
x_q_scratch: &mut Vec<i8>,
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;
gdn_scratch.ensure_capacity(qkv_dim, output_dim, num_heads, key_dim, value_dim);
matmul_q8_neon_into(
input,
&weights.in_proj_qkv_packed,
qkv_dim,
hidden,
&mut gdn_scratch.qkv_proj[..qkv_dim],
x_q_scratch,
);
matmul_q8_neon_into(
input,
&weights.in_proj_z_packed,
output_dim,
hidden,
&mut gdn_scratch.z_proj[..output_dim],
x_q_scratch,
);
matmul_q8_neon_into(
input,
&weights.in_proj_b_packed,
num_heads,
hidden,
&mut gdn_scratch.beta_proj[..num_heads],
x_q_scratch,
);
matmul_q8_neon_into(
input,
&weights.in_proj_a_packed,
num_heads,
hidden,
&mut gdn_scratch.alpha_proj[..num_heads],
x_q_scratch,
);
for b in &mut gdn_scratch.beta_proj[..num_heads] {
*b = sigmoid(*b);
}
let conv_dim = weights.conv_dim;
let buf_len = kernel_size - 1;
for ch in 0..conv_dim {
let qkv_ch = gdn_scratch.qkv_proj[ch];
let buf_start = ch * buf_len;
for j in 0..buf_len.saturating_sub(1) {
state.conv_buffer[buf_start + j] = state.conv_buffer[buf_start + j + 1];
}
if buf_len > 0 {
state.conv_buffer[buf_start + buf_len - 1] = qkv_ch;
}
let mut acc =
gdn_scratch.qkv_proj[ch] * weights.conv1d_weight[ch * kernel_size + kernel_size - 1];
for k in 0..buf_len {
acc += state.conv_buffer[buf_start + k] * weights.conv1d_weight[ch * kernel_size + k];
}
let sig = 1.0 / (1.0 + (-acc).exp());
gdn_scratch.conv_output[ch] = acc * sig;
}
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;
gdn_scratch.q_head[..key_dim]
.copy_from_slice(&gdn_scratch.conv_output[q_start..q_start + key_dim]);
gdn_scratch.k_head[..key_dim]
.copy_from_slice(&gdn_scratch.conv_output[k_start..k_start + key_dim]);
l2_normalize_vec(&mut gdn_scratch.q_head[..key_dim]);
l2_normalize_vec(&mut gdn_scratch.k_head[..key_dim]);
let a_val = weights.a_log[k_head].exp().min(f32::MAX);
let sp = softplus(gdn_scratch.alpha_proj[k_head] + weights.dt_bias[k_head]);
let g = (-a_val * sp).exp();
let s_off = h * key_dim * value_dim;
let s = &mut state.s_matrices[s_off..s_off + key_dim * value_dim];
for j in 0..value_dim {
let mut dot = 0.0f32;
for i in 0..key_dim {
dot += s[i * value_dim + j] * gdn_scratch.k_head[i];
}
gdn_scratch.kv_mem[j] = dot;
}
let beta_h = gdn_scratch.beta_proj[k_head];
for j in 0..value_dim {
let v_j = gdn_scratch.conv_output[v_start + j];
gdn_scratch.delta[j] = (v_j - gdn_scratch.kv_mem[j] * g) * beta_h;
}
for i in 0..key_dim {
for j in 0..value_dim {
s[i * value_dim + j] =
g * s[i * value_dim + j] + gdn_scratch.k_head[i] * gdn_scratch.delta[j];
}
}
let out_start = h * value_dim;
for j in 0..value_dim {
let mut dot = 0.0f32;
for i in 0..key_dim {
dot += s[i * value_dim + j] * gdn_scratch.q_head[i];
}
gdn_scratch.output_heads[out_start + j] = dot * scale;
}
}
let gamma = &weights.norm_weight[..value_dim];
for h in 0..value_heads {
let start = h * value_dim;
let end = start + value_dim;
gated_rms_norm(
&gdn_scratch.output_heads[start..end],
&gdn_scratch.z_proj[start..end],
gamma,
&mut gdn_scratch.gated_norm_buf[start..end],
cfg.rms_norm_eps,
);
}
matmul_q8_neon_into(
&gdn_scratch.gated_norm_buf[..output_dim],
&weights.out_proj_packed,
hidden,
output_dim,
&mut output[..hidden],
x_q_scratch,
);
}
fn full_attention_step_q8_neon(
weights: &Q8NeonFullAttnWeights,
cache_idx: usize,
position: usize,
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
cfg: &Qwen35Config,
rope: &RopeTable,
hidden: usize,
) {
scratch.input_tmp[..hidden].copy_from_slice(&scratch.attn_out[..hidden]);
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;
matmul_q8_neon_into(
&scratch.input_tmp[..hidden],
&weights.q_proj_packed,
q_proj_dim,
hidden,
&mut scratch.q_and_gate[..q_proj_dim],
&mut scratch.x_q_scratch,
);
for h in 0..num_q_heads {
let src = h * head_dim * 2;
let dst = h * head_dim;
scratch.q_buf[dst..dst + head_dim]
.copy_from_slice(&scratch.q_and_gate[src..src + head_dim]);
scratch.gate_z[dst..dst + head_dim]
.copy_from_slice(&scratch.q_and_gate[src + head_dim..src + head_dim * 2]);
}
matmul_q8_neon_into(
&scratch.input_tmp[..hidden],
&weights.k_proj_packed,
kv_dim,
hidden,
&mut scratch.k_buf[..kv_dim],
&mut scratch.x_q_scratch,
);
matmul_q8_neon_into(
&scratch.input_tmp[..hidden],
&weights.v_proj_packed,
kv_dim,
hidden,
&mut scratch.v_buf[..kv_dim],
&mut scratch.x_q_scratch,
);
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;
}
if sum_exp > 0.0 {
let inv_sum = 1.0 / sum_exp;
for t in 0..cur_seq_len {
scratch.scores[scores_start + t] *= inv_sum;
}
} else {
scratch.scores[scores_start..scores_start + cur_seq_len].fill(0.0);
}
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;
}
}
for d in 0..q_dim {
let sig = 1.0 / (1.0 + (-scratch.gate_z[d]).exp());
scratch.context[d] *= sig;
}
matmul_q8_neon_into(
&scratch.context[..q_dim],
&weights.o_proj_packed,
hidden,
q_dim,
&mut scratch.attn_out[..hidden],
&mut scratch.x_q_scratch,
);
}
#[inline]
fn ffn_step_q8_neon(common: &Q8NeonCommonWeights, scratch: &mut ForwardScratch, hidden: usize) {
let inter = common.gate_proj_rows;
matmul_q8_neon_into(
&scratch.ffn_out[..hidden],
&common.gate_proj_packed,
inter,
hidden,
&mut scratch.gate_buf[..inter],
&mut scratch.x_q_scratch,
);
matmul_q8_neon_into(
&scratch.ffn_out[..hidden],
&common.up_proj_packed,
inter,
hidden,
&mut scratch.up_buf[..inter],
&mut scratch.x_q_scratch,
);
silu_inplace(&mut scratch.gate_buf[..inter]);
elementwise_mul(&mut scratch.gate_buf[..inter], &scratch.up_buf[..inter]);
matmul_q8_neon_into(
&scratch.gate_buf[..inter],
&common.down_proj_packed,
hidden,
inter,
&mut scratch.ffn_out[..hidden],
&mut scratch.x_q_scratch,
);
}
pub(crate) fn forward_step_q8_neon(
model: &Q8NeonModel,
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(&model.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) = &model.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 {
Q8NeonAttentionWeights::Linear(gdn_w) => {
gdn_step_q8_neon(
&scratch.hidden[..hidden],
&mut gdn_states[linear_idx],
gdn_w,
cfg,
&mut scratch.gdn_scratch,
&mut scratch.x_q_scratch,
&mut scratch.attn_out[..hidden],
);
linear_idx += 1;
}
Q8NeonAttentionWeights::Full(full_w) => {
scratch.attn_out[..hidden].copy_from_slice(&scratch.hidden[..hidden]);
full_attention_step_q8_neon(
full_w, 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_neon(common, scratch, hidden);
for i in 0..hidden {
scratch.hidden[i] = scratch.residual[i] + scratch.ffn_out[i];
}
}
qwen35_rms_norm(
&mut scratch.hidden[..hidden],
&model.final_norm,
hidden,
cfg.rms_norm_eps,
);
resize(&mut scratch.logits, cfg.vocab_size);
matmul_q8_neon_into(
&scratch.hidden[..hidden],
&model.lm_head_packed,
model.lm_head_rows,
model.lm_head_cols,
&mut scratch.logits[..cfg.vocab_size],
&mut scratch.x_q_scratch,
);
}
pub fn generate_q8_neon(
model: &Q8NeonModel,
cfg: &Qwen35Config,
tokenizer: &BpeTokenizer,
rope: &RopeTable,
prompt: &str,
gen_cfg: &GenerateConfig,
) -> Result<GenerateOutput, crate::error::InferenceError> {
let mut rng_state = match gen_cfg.seed {
Some(s) => {
if s == 0 {
1
} else {
s
}
}
None => {
use std::time::SystemTime;
let t = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0x12345678_9abcdef0);
if t == 0 { 1 } else { t }
}
};
let input = tokenizer.tokenize(prompt);
let prompt_ids: Vec<u32> = input.input_ids[..input.real_length].to_vec();
let prompt_len = prompt_ids.len();
if prompt_len == 0 {
return Err(crate::error::InferenceError::Inference(
"empty prompt".into(),
));
}
let max_context = rope.max_positions();
if prompt_len.saturating_add(gen_cfg.max_new_tokens) > max_context {
return Err(crate::error::InferenceError::Inference(format!(
"prompt ({prompt_len} tokens) plus max_new_tokens ({}) exceeds \
model context window ({max_context})",
gen_cfg.max_new_tokens
)));
}
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 max_seq_len = prompt_len
.saturating_add(gen_cfg.max_new_tokens)
.saturating_add(1);
kv_cache.reserve(max_seq_len, cfg.full_kv_dim());
scratch.ensure_capacity(cfg, max_seq_len);
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_neon(
model,
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 next_id == cfg.eos_token_id {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: true,
});
}
generated_ids.push(next_id);
all_ids.push(next_id);
let mut stopped = false;
for _ in 1..gen_cfg.max_new_tokens {
let pos = kv_cache.seq_len;
let last_token = *all_ids
.last()
.expect("invariant: prompt or previous sample populated all_ids");
forward_step_q8_neon(
model,
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 next_id == cfg.eos_token_id {
stopped = true;
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,
})
}
#[cfg(feature = "bench-internals")]
pub mod bench_support {
use super::*;
use crate::model::qwen35_config::LayerType;
use crate::rope::RopeTable;
pub struct Q8ForwardBenchFixture {
model: Q8NeonModel,
cfg: Qwen35Config,
rope: RopeTable,
}
pub struct Q8ForwardBenchState {
gdn_states: Vec<GatedDeltaNetState>,
kv_cache: KvCache,
scratch: ForwardScratch,
}
fn xorshift64(state: &mut u64) -> f32 {
let mut x = *state;
x ^= x << 13;
x ^= x >> 7;
x ^= x << 17;
*state = x;
(x & 0xFFFF) as f32 / 0x10000_u64 as f32 * 0.04 - 0.02
}
fn gen_weights_q8(n: usize, k: usize, rng: &mut u64) -> Vec<u8> {
let floats: Vec<f32> = (0..n * k).map(|_| xorshift64(rng)).collect();
pack_weights_q8(&floats, n, k)
}
impl Q8ForwardBenchFixture {
pub fn synthetic_2layer() -> Self {
let mut rng: u64 = 0xdeadbeef_cafebabe;
let hidden: usize = 256;
let vocab: usize = 8192;
let inter: usize = 768;
let num_attn_heads: usize = 4;
let num_kv_heads: usize = 2;
let head_dim: usize = 64;
let q_dim = num_attn_heads * head_dim; let kv_dim = num_kv_heads * head_dim;
let lin_key_heads: usize = 4;
let lin_val_heads: usize = 4;
let lin_key_dim: usize = 64;
let lin_val_dim: usize = 64;
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: 8191,
max_position_embeddings: 512,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: 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 gdn_weights = Q8NeonGdnWeights {
in_proj_qkv_packed: gen_weights_q8(lin_qkv_dim, hidden, &mut rng),
in_proj_qkv_rows: lin_qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z_packed: gen_weights_q8(lin_output_dim, hidden, &mut rng),
in_proj_z_rows: lin_output_dim,
in_proj_z_cols: hidden,
in_proj_b_packed: gen_weights_q8(lin_key_heads, hidden, &mut rng),
in_proj_b_rows: lin_key_heads,
in_proj_b_cols: hidden,
in_proj_a_packed: gen_weights_q8(lin_key_heads, hidden, &mut rng),
in_proj_a_rows: lin_key_heads,
in_proj_a_cols: hidden,
out_proj_packed: gen_weights_q8(hidden, lin_output_dim, &mut rng),
out_proj_rows: hidden,
out_proj_cols: lin_output_dim,
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],
};
let full_weights = Q8NeonFullAttnWeights {
q_proj_packed: gen_weights_q8(2 * q_dim, hidden, &mut rng),
q_proj_rows: 2 * q_dim,
q_proj_cols: hidden,
k_proj_packed: gen_weights_q8(kv_dim, hidden, &mut rng),
k_proj_rows: kv_dim,
k_proj_cols: hidden,
v_proj_packed: gen_weights_q8(kv_dim, hidden, &mut rng),
v_proj_rows: kv_dim,
v_proj_cols: hidden,
o_proj_packed: gen_weights_q8(hidden, q_dim, &mut rng),
o_proj_rows: hidden,
o_proj_cols: q_dim,
q_norm: vec![0.0f32; head_dim],
k_norm: vec![0.0f32; head_dim],
};
let make_common = |rng: &mut u64| Q8NeonCommonWeights {
input_layernorm: vec![0.0f32; hidden],
post_attention_layernorm: vec![0.0f32; hidden],
gate_proj_packed: gen_weights_q8(inter, hidden, rng),
gate_proj_rows: inter,
gate_proj_cols: hidden,
up_proj_packed: gen_weights_q8(inter, hidden, rng),
up_proj_rows: inter,
up_proj_cols: hidden,
down_proj_packed: gen_weights_q8(hidden, inter, rng),
down_proj_rows: hidden,
down_proj_cols: inter,
};
let common0 = make_common(&mut rng);
let common1 = make_common(&mut rng);
let embed_tokens: Vec<f32> = (0..vocab * hidden)
.map(|i| {
let mut s = (i as u64).wrapping_mul(0x9e3779b9_7f4a7c15);
s ^= s >> 33;
s &= 0xFFFF;
s as f32 / 0x10000_u64 as f32 * 0.04 - 0.02
})
.collect();
let lm_head_packed = pack_weights_q8(&embed_tokens, vocab, hidden);
let model = Q8NeonModel {
embed_tokens,
final_norm: vec![0.0f32; hidden],
lm_head_packed,
lm_head_rows: vocab,
lm_head_cols: hidden,
layers: vec![
(Q8NeonAttentionWeights::Linear(gdn_weights), common0),
(Q8NeonAttentionWeights::Full(full_weights), common1),
],
};
Self { model, cfg, rope }
}
pub fn qwen35_24layer_shape() -> Self {
let mut rng: u64 = 0x0123_4567_89ab_cdef;
let mut cfg = Qwen35Config::qwen35_2b();
cfg.vocab_size = 256;
let hidden = cfg.hidden_size;
let vocab = cfg.vocab_size;
let inter = cfg.intermediate_size;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let num_heads_lin = cfg.linear_num_key_heads;
let lin_val_dim = cfg.linear_value_head_dim;
let kernel_size = cfg.linear_conv_kernel_dim;
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let make_gdn = |rng: &mut u64| Q8NeonGdnWeights {
in_proj_qkv_packed: gen_weights_q8(qkv_dim, hidden, rng),
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z_packed: gen_weights_q8(output_dim, hidden, rng),
in_proj_z_rows: output_dim,
in_proj_z_cols: hidden,
in_proj_b_packed: gen_weights_q8(num_heads_lin, hidden, rng),
in_proj_b_rows: num_heads_lin,
in_proj_b_cols: hidden,
in_proj_a_packed: gen_weights_q8(num_heads_lin, hidden, rng),
in_proj_a_rows: num_heads_lin,
in_proj_a_cols: hidden,
out_proj_packed: gen_weights_q8(hidden, output_dim, rng),
out_proj_rows: hidden,
out_proj_cols: output_dim,
a_log: vec![0.0f32; num_heads_lin],
dt_bias: vec![0.0f32; num_heads_lin],
conv1d_weight: vec![0.01f32; qkv_dim * kernel_size],
conv_dim: qkv_dim,
kernel_size,
norm_weight: vec![0.0f32; lin_val_dim],
};
let make_full = |rng: &mut u64| Q8NeonFullAttnWeights {
q_proj_packed: gen_weights_q8(2 * q_dim, hidden, rng),
q_proj_rows: 2 * q_dim,
q_proj_cols: hidden,
k_proj_packed: gen_weights_q8(kv_dim, hidden, rng),
k_proj_rows: kv_dim,
k_proj_cols: hidden,
v_proj_packed: gen_weights_q8(kv_dim, hidden, rng),
v_proj_rows: kv_dim,
v_proj_cols: hidden,
o_proj_packed: gen_weights_q8(hidden, q_dim, rng),
o_proj_rows: hidden,
o_proj_cols: q_dim,
q_norm: vec![0.0f32; cfg.head_dim],
k_norm: vec![0.0f32; cfg.head_dim],
};
let make_common = |rng: &mut u64| Q8NeonCommonWeights {
input_layernorm: vec![0.0f32; hidden],
post_attention_layernorm: vec![0.0f32; hidden],
gate_proj_packed: gen_weights_q8(inter, hidden, rng),
gate_proj_rows: inter,
gate_proj_cols: hidden,
up_proj_packed: gen_weights_q8(inter, hidden, rng),
up_proj_rows: inter,
up_proj_cols: hidden,
down_proj_packed: gen_weights_q8(hidden, inter, rng),
down_proj_rows: hidden,
down_proj_cols: inter,
};
let layers: Vec<(Q8NeonAttentionWeights, Q8NeonCommonWeights)> = cfg
.layer_types
.iter()
.map(|lt| match lt {
LayerType::LinearAttention => (
Q8NeonAttentionWeights::Linear(make_gdn(&mut rng)),
make_common(&mut rng),
),
LayerType::FullAttention => (
Q8NeonAttentionWeights::Full(make_full(&mut rng)),
make_common(&mut rng),
),
})
.collect();
let embed_tokens: Vec<f32> = (0..vocab * hidden)
.map(|i| {
let mut s = (i as u64).wrapping_mul(0x9e3779b9_7f4a7c15);
s ^= s >> 33;
s &= 0xFFFF;
s as f32 / 0x10000_u64 as f32 * 0.04 - 0.02
})
.collect();
let lm_head_packed = pack_weights_q8(&embed_tokens, vocab, hidden);
let model = Q8NeonModel {
embed_tokens,
final_norm: vec![0.0f32; hidden],
lm_head_packed,
lm_head_rows: vocab,
lm_head_cols: hidden,
layers,
};
let rope_dim = cfg.rope_dim();
let rope = RopeTable::new(rope_dim, cfg.max_position_embeddings, cfg.rope_theta);
Self { model, cfg, rope }
}
pub fn state(&self, warm_len: usize) -> Q8ForwardBenchState {
self.state_with_capacity(warm_len, 1)
}
pub fn state_with_capacity(
&self,
warm_len: usize,
measured_tokens: usize,
) -> Q8ForwardBenchState {
let num_linear = self.cfg.num_linear_attention_layers();
let num_full = self.cfg.num_full_attention_layers();
let mut gdn_states: Vec<GatedDeltaNetState> = (0..num_linear)
.map(|_| GatedDeltaNetState::new(&self.cfg))
.collect();
let mut kv_cache = KvCache::new(num_full);
let mut scratch = ForwardScratch::new();
let max_seq_len = warm_len.saturating_add(measured_tokens).saturating_add(1);
kv_cache.reserve(max_seq_len, self.cfg.full_kv_dim());
scratch.ensure_capacity(&self.cfg, max_seq_len);
for pos in 0..warm_len {
let token_id = (pos as u32) % (self.cfg.vocab_size as u32);
forward_step_q8_neon(
&self.model,
&self.cfg,
&self.rope,
token_id,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len += 1;
}
Q8ForwardBenchState {
gdn_states,
kv_cache,
scratch,
}
}
pub fn step(&self, state: &mut Q8ForwardBenchState, token_id: u32) -> f32 {
let pos = state.kv_cache.seq_len;
forward_step_q8_neon(
&self.model,
&self.cfg,
&self.rope,
token_id % (self.cfg.vocab_size as u32),
pos,
&mut state.gdn_states,
&mut state.kv_cache,
&mut state.scratch,
);
state.kv_cache.seq_len += 1;
state.scratch.logits[0]
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::qwen35_config::LayerType;
fn zero_packed(n: usize, k: usize) -> Vec<u8> {
pack_weights_q8(&vec![0.0f32; n * k], n, k)
}
#[test]
fn test_quantize_model_produces_valid_packed_sizes() {
let cfg = Qwen35Config::qwen35_2b();
let hidden = cfg.hidden_size;
let vocab = cfg.vocab_size;
let inter = cfg.intermediate_size;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let q8_packed_size = |n: usize, k: usize| -> usize {
assert_eq!(k % 32, 0, "k={k} must be multiple of 32");
n * (k / 32) * 36
};
let gdn_w = GatedDeltaNetWeights {
in_proj_qkv: vec![0.0; qkv_dim * hidden],
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z: vec![0.0; output_dim * hidden],
in_proj_z_rows: output_dim,
in_proj_z_cols: hidden,
in_proj_b: vec![0.0; cfg.linear_num_key_heads * hidden],
in_proj_b_rows: cfg.linear_num_key_heads,
in_proj_b_cols: hidden,
in_proj_a: vec![0.0; cfg.linear_num_key_heads * hidden],
in_proj_a_rows: cfg.linear_num_key_heads,
in_proj_a_cols: hidden,
a_log: vec![0.0; cfg.linear_num_key_heads],
dt_bias: vec![0.0; cfg.linear_num_key_heads],
conv1d_weight: vec![0.0; qkv_dim * cfg.linear_conv_kernel_dim],
conv_dim: qkv_dim,
kernel_size: cfg.linear_conv_kernel_dim,
norm_weight: vec![0.0; cfg.linear_value_head_dim],
out_proj: vec![0.0; hidden * output_dim],
out_proj_rows: hidden,
out_proj_cols: output_dim,
};
let full_w = FullAttentionLayerWeights {
q_proj: vec![0.0; 2 * q_dim * hidden],
k_proj: vec![0.0; kv_dim * hidden],
v_proj: vec![0.0; kv_dim * hidden],
o_proj: vec![0.0; hidden * q_dim],
q_norm: vec![0.0; cfg.head_dim],
k_norm: vec![0.0; cfg.head_dim],
};
let make_common_w = || CommonLayerWeights {
input_layernorm: vec![0.0; hidden],
post_attention_layernorm: vec![0.0; hidden],
ffn: crate::model::qwen35::FeedForwardWeights::Dense(
crate::model::qwen35::DenseFfnWeights {
gate_proj: vec![0.0; inter * hidden],
up_proj: vec![0.0; inter * hidden],
down_proj: vec![0.0; hidden * inter],
},
),
};
let weights = ModelWeights {
embed_tokens: vec![0.0; vocab * hidden],
lm_head: None,
final_norm: vec![0.0; hidden],
layers: vec![
(AttentionWeights::Linear(gdn_w), make_common_w()),
(AttentionWeights::Full(full_w), make_common_w()),
],
};
let q8 = quantize_model(&weights, &cfg);
assert_eq!(q8.lm_head_packed.len(), q8_packed_size(vocab, hidden));
assert_eq!(q8.lm_head_rows, vocab);
assert_eq!(q8.lm_head_cols, hidden);
assert_eq!(q8.embed_tokens.len(), vocab * hidden);
match &q8.layers[0].0 {
Q8NeonAttentionWeights::Linear(gdn) => {
assert_eq!(
gdn.in_proj_qkv_packed.len(),
q8_packed_size(qkv_dim, hidden)
);
assert_eq!(
gdn.in_proj_z_packed.len(),
q8_packed_size(output_dim, hidden)
);
assert_eq!(
gdn.out_proj_packed.len(),
q8_packed_size(hidden, output_dim)
);
assert_eq!(gdn.a_log.len(), cfg.linear_num_key_heads);
}
Q8NeonAttentionWeights::Full(_) => panic!("expected Linear layer"),
}
match &q8.layers[1].0 {
Q8NeonAttentionWeights::Full(full) => {
assert_eq!(full.q_proj_packed.len(), q8_packed_size(2 * q_dim, hidden));
assert_eq!(full.k_proj_packed.len(), q8_packed_size(kv_dim, hidden));
assert_eq!(full.v_proj_packed.len(), q8_packed_size(kv_dim, hidden));
assert_eq!(full.o_proj_packed.len(), q8_packed_size(hidden, q_dim));
}
Q8NeonAttentionWeights::Linear(_) => panic!("expected Full layer"),
}
let c = &q8.layers[0].1;
assert_eq!(c.gate_proj_packed.len(), q8_packed_size(inter, hidden));
assert_eq!(c.up_proj_packed.len(), q8_packed_size(inter, hidden));
assert_eq!(c.down_proj_packed.len(), q8_packed_size(hidden, inter));
assert_eq!(c.input_layernorm.len(), hidden);
}
#[test]
fn test_forward_step_q8_neon_zero_weights_produces_zero_logits() {
let cfg = Qwen35Config::qwen35_2b();
let hidden = cfg.hidden_size;
let vocab = cfg.vocab_size;
let inter = cfg.intermediate_size;
let qkv_dim = cfg.linear_qkv_dim();
let output_dim = cfg.linear_output_dim();
let num_heads = cfg.linear_num_key_heads;
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let make_linear = || Q8NeonGdnWeights {
in_proj_qkv_packed: zero_packed(qkv_dim, hidden),
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z_packed: zero_packed(output_dim, hidden),
in_proj_z_rows: output_dim,
in_proj_z_cols: hidden,
in_proj_b_packed: zero_packed(num_heads, hidden),
in_proj_b_rows: num_heads,
in_proj_b_cols: hidden,
in_proj_a_packed: zero_packed(num_heads, hidden),
in_proj_a_rows: num_heads,
in_proj_a_cols: hidden,
out_proj_packed: zero_packed(hidden, output_dim),
out_proj_rows: hidden,
out_proj_cols: output_dim,
a_log: vec![0.0; num_heads],
dt_bias: vec![0.0; num_heads],
conv1d_weight: vec![0.0; qkv_dim * cfg.linear_conv_kernel_dim],
conv_dim: qkv_dim,
kernel_size: cfg.linear_conv_kernel_dim,
norm_weight: vec![0.0; cfg.linear_value_head_dim],
};
let make_full = || Q8NeonFullAttnWeights {
q_proj_packed: zero_packed(2 * q_dim, hidden),
q_proj_rows: 2 * q_dim,
q_proj_cols: hidden,
k_proj_packed: zero_packed(kv_dim, hidden),
k_proj_rows: kv_dim,
k_proj_cols: hidden,
v_proj_packed: zero_packed(kv_dim, hidden),
v_proj_rows: kv_dim,
v_proj_cols: hidden,
o_proj_packed: zero_packed(hidden, q_dim),
o_proj_rows: hidden,
o_proj_cols: q_dim,
q_norm: vec![0.0; cfg.head_dim],
k_norm: vec![0.0; cfg.head_dim],
};
let make_common = || Q8NeonCommonWeights {
input_layernorm: vec![0.0; hidden],
post_attention_layernorm: vec![0.0; hidden],
gate_proj_packed: zero_packed(inter, hidden),
gate_proj_rows: inter,
gate_proj_cols: hidden,
up_proj_packed: zero_packed(inter, hidden),
up_proj_rows: inter,
up_proj_cols: hidden,
down_proj_packed: zero_packed(hidden, inter),
down_proj_rows: hidden,
down_proj_cols: inter,
};
let mut test_cfg = cfg.clone();
test_cfg.num_hidden_layers = 2;
test_cfg.layer_types = vec![
crate::model::qwen35_config::LayerType::LinearAttention,
crate::model::qwen35_config::LayerType::FullAttention,
];
let model = Q8NeonModel {
embed_tokens: vec![0.0; vocab * hidden],
final_norm: vec![0.0; hidden],
lm_head_packed: zero_packed(vocab, hidden),
lm_head_rows: vocab,
lm_head_cols: hidden,
layers: vec![
(Q8NeonAttentionWeights::Linear(make_linear()), make_common()),
(Q8NeonAttentionWeights::Full(make_full()), make_common()),
],
};
let rope_dim = test_cfg.rope_dim();
let rope_max = test_cfg.max_position_embeddings.min(8192);
let rope = RopeTable::new(rope_dim, rope_max, test_cfg.rope_theta);
let num_linear = test_cfg.num_linear_attention_layers();
let num_full = test_cfg.num_full_attention_layers();
let mut gdn_states: Vec<GatedDeltaNetState> = (0..num_linear)
.map(|_| GatedDeltaNetState::new(&test_cfg))
.collect();
let mut kv_cache = KvCache::new(num_full);
let mut scratch = ForwardScratch::new();
forward_step_q8_neon(
&model,
&test_cfg,
&rope,
0, 0, &mut gdn_states,
&mut kv_cache,
&mut scratch,
);
for (i, &v) in scratch.logits[..test_cfg.vocab_size].iter().enumerate() {
assert!(
v.abs() < 1e-6,
"logit[{i}] = {v}, expected ~0.0 with zero weights"
);
}
}
#[test]
fn test_gdn_step_q8_neon_zero_weights_produces_zero_output() {
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 weights = Q8NeonGdnWeights {
in_proj_qkv_packed: zero_packed(qkv_dim, hidden),
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z_packed: zero_packed(output_dim, hidden),
in_proj_z_rows: output_dim,
in_proj_z_cols: hidden,
in_proj_b_packed: zero_packed(num_heads, hidden),
in_proj_b_rows: num_heads,
in_proj_b_cols: hidden,
in_proj_a_packed: zero_packed(num_heads, hidden),
in_proj_a_rows: num_heads,
in_proj_a_cols: hidden,
out_proj_packed: zero_packed(hidden, output_dim),
out_proj_rows: hidden,
out_proj_cols: output_dim,
a_log: vec![0.0; num_heads],
dt_bias: vec![0.0; num_heads],
conv1d_weight: vec![0.0; qkv_dim * cfg.linear_conv_kernel_dim],
conv_dim: qkv_dim,
kernel_size: cfg.linear_conv_kernel_dim,
norm_weight: vec![0.0; cfg.linear_value_head_dim],
};
let mut state = GatedDeltaNetState::new(&cfg);
let input = vec![0.0f32; hidden];
let mut output = vec![0.0f32; hidden];
let mut gdn_scratch = GatedDeltaNetFusedScratch::default();
let mut x_q_scratch = Vec::new();
gdn_step_q8_neon(
&input,
&mut state,
&weights,
&cfg,
&mut gdn_scratch,
&mut x_q_scratch,
&mut output,
);
for (i, &v) in output[..hidden].iter().enumerate() {
assert!(
v.abs() < 1e-6,
"output[{i}] = {v}, expected 0.0 with zero weights + zero input"
);
}
}
#[test]
fn test_ffn_step_q8_neon_zero_weights() {
let cfg = Qwen35Config::qwen35_2b();
let hidden = cfg.hidden_size;
let inter = cfg.intermediate_size;
let common = Q8NeonCommonWeights {
input_layernorm: vec![0.0; hidden],
post_attention_layernorm: vec![0.0; hidden],
gate_proj_packed: zero_packed(inter, hidden),
gate_proj_rows: inter,
gate_proj_cols: hidden,
up_proj_packed: zero_packed(inter, hidden),
up_proj_rows: inter,
up_proj_cols: hidden,
down_proj_packed: zero_packed(hidden, inter),
down_proj_rows: hidden,
down_proj_cols: inter,
};
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(&cfg, 1);
scratch.ffn_out[..hidden].fill(0.0);
ffn_step_q8_neon(&common, &mut scratch, hidden);
for (i, &v) in scratch.ffn_out[..hidden].iter().enumerate() {
assert!(
v.abs() < 1e-6,
"ffn_out[{i}] = {v}, expected 0.0 with zero weights"
);
}
}
#[test]
fn test_full_attn_step_q8_neon_rope_stride_half_parity() {
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,
};
let rope_dim = cfg.rope_dim(); let half = rope_dim / 2; let rope = RopeTable::new(rope_dim, 512, cfg.rope_theta);
let identity_packed = {
let mut mat = vec![0.0f32; kv_dim * hidden];
for j in 0..kv_dim {
mat[j * hidden + j] = 1.0;
}
pack_weights_q8(&mat, kv_dim, hidden)
};
let q_identity_packed = {
let mut mat = vec![0.0f32; 2 * q_dim * hidden];
for j in 0..q_dim {
mat[j * hidden + j] = 1.0;
}
pack_weights_q8(&mat, 2 * q_dim, hidden)
};
let weights = Q8NeonFullAttnWeights {
q_proj_packed: q_identity_packed,
q_proj_rows: 2 * q_dim,
q_proj_cols: hidden,
k_proj_packed: identity_packed,
k_proj_rows: kv_dim,
k_proj_cols: hidden,
v_proj_packed: zero_packed(kv_dim, hidden),
v_proj_rows: kv_dim,
v_proj_cols: hidden,
o_proj_packed: zero_packed(hidden, q_dim),
o_proj_rows: hidden,
o_proj_cols: 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_neon(
&weights,
0,
position,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
);
let mut k_ref = vec![0.0f32; kv_dim];
matmul_q8_neon_into(
&input,
&weights.k_proj_packed,
kv_dim,
hidden,
&mut k_ref,
&mut scratch.x_q_scratch,
);
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,
"NEON 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 q_proj_dim = 2 * q_dim;
let mut q_and_gate_ref = vec![0.0f32; q_proj_dim];
matmul_q8_neon_into(
&input,
&weights.q_proj_packed,
q_proj_dim,
hidden,
&mut q_and_gate_ref,
&mut scratch.x_q_scratch,
);
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,
"NEON 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."
);
}
fn make_nonzero_q8_neon_test_model() -> (Qwen35Config, Q8NeonModel, RopeTable) {
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,
};
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| -> Vec<u8> {
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();
pack_weights_q8(&floats, n, k)
};
let gdn_w = Q8NeonGdnWeights {
in_proj_qkv_packed: next_weight(lin_qkv_dim, hidden),
in_proj_qkv_rows: lin_qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z_packed: next_weight(lin_output_dim, hidden),
in_proj_z_rows: lin_output_dim,
in_proj_z_cols: hidden,
in_proj_b_packed: next_weight(lin_key_heads, hidden),
in_proj_b_rows: lin_key_heads,
in_proj_b_cols: hidden,
in_proj_a_packed: next_weight(lin_key_heads, hidden),
in_proj_a_rows: lin_key_heads,
in_proj_a_cols: hidden,
out_proj_packed: next_weight(hidden, lin_output_dim),
out_proj_rows: hidden,
out_proj_cols: lin_output_dim,
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],
};
let full_w = Q8NeonFullAttnWeights {
q_proj_packed: next_weight(2 * q_dim, hidden),
q_proj_rows: 2 * q_dim,
q_proj_cols: hidden,
k_proj_packed: next_weight(kv_dim, hidden),
k_proj_rows: kv_dim,
k_proj_cols: hidden,
v_proj_packed: next_weight(kv_dim, hidden),
v_proj_rows: kv_dim,
v_proj_cols: hidden,
o_proj_packed: next_weight(hidden, q_dim),
o_proj_rows: hidden,
o_proj_cols: q_dim,
q_norm: vec![0.0f32; head_dim],
k_norm: vec![0.0f32; head_dim],
};
let common_w = |seed: &mut u64| {
let mut nw = |n: usize, k: usize| -> Vec<u8> {
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();
pack_weights_q8(&floats, n, k)
};
Q8NeonCommonWeights {
input_layernorm: vec![0.0f32; hidden],
post_attention_layernorm: vec![0.0f32; hidden],
gate_proj_packed: nw(inter, hidden),
gate_proj_rows: inter,
gate_proj_cols: hidden,
up_proj_packed: nw(inter, hidden),
up_proj_rows: inter,
up_proj_cols: hidden,
down_proj_packed: nw(hidden, inter),
down_proj_rows: hidden,
down_proj_cols: inter,
}
};
let common0 = common_w(&mut seed);
let common1 = common_w(&mut seed);
let embed_tokens: Vec<f32> = (0..vocab * hidden)
.map(|i| {
let mut s = (i as u64).wrapping_mul(0x9e3779b9_7f4a7c15);
s ^= s >> 33;
s &= 0xFFFF;
s as f32 / 0x10000_u64 as f32 * 0.04 - 0.02
})
.collect();
let lm_head_packed = pack_weights_q8(&embed_tokens, vocab, hidden);
let model = Q8NeonModel {
embed_tokens,
final_norm: vec![0.0f32; hidden],
lm_head_packed,
lm_head_rows: vocab,
lm_head_cols: hidden,
layers: vec![
(Q8NeonAttentionWeights::Linear(gdn_w), common0),
(Q8NeonAttentionWeights::Full(full_w), common1),
],
};
(cfg, model, rope)
}
#[test]
fn test_forward_step_q8_neon_into_migration_preserves_nonzero_logits() {
let (cfg, model, rope) = make_nonzero_q8_neon_test_model();
let num_linear = cfg.num_linear_attention_layers();
let num_full = cfg.num_full_attention_layers();
let run_two_steps = || -> Vec<f32> {
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();
forward_step_q8_neon(
&model,
&cfg,
&rope,
7,
0,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
kv_cache.seq_len += 1;
forward_step_q8_neon(
&model,
&cfg,
&rope,
11,
1,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
);
scratch.logits[..16].to_vec()
};
let logits_a = run_two_steps();
let logits_b = run_two_steps();
assert_eq!(
logits_a, logits_b,
"forward_step_q8_neon is non-deterministic"
);
assert!(
logits_a.iter().any(|&v| v.abs() > 1e-9),
"all 16 logits are zero — check weight generation"
);
let expected: [f32; 16] = [
-0.036519982,
0.04271806,
0.0027702842,
0.007089529,
-0.074217916,
0.001136966,
0.07316325,
0.027880985,
-0.027898913,
-0.10935478,
0.04180234,
0.08315472,
0.008195344,
-0.07999688,
0.014749594,
0.028362377,
];
for (i, (&actual, &exp)) in logits_a.iter().zip(expected.iter()).enumerate() {
assert!(
(actual - exp).abs() <= 1e-6,
"logit[{i}] mismatch: actual={actual:.8}, expected={exp:.8}"
);
}
}
#[test]
fn test_all_projection_dims_are_multiples_of_32() {
let cfg = Qwen35Config::qwen35_2b();
let dims_to_check = [
("hidden_size", cfg.hidden_size),
("intermediate_size", cfg.intermediate_size),
("full_q_dim", cfg.full_q_dim()),
("full_kv_dim", cfg.full_kv_dim()),
("2*full_q_dim", 2 * cfg.full_q_dim()),
("linear_qkv_dim", cfg.linear_qkv_dim()),
("linear_output_dim", cfg.linear_output_dim()),
];
for (name, dim) in &dims_to_check {
assert_eq!(
dim % 32,
0,
"{name} = {dim} is not a multiple of 32 (required for Q8_0)"
);
}
}
#[test]
fn test_full_attn_step_q8_neon_nan_score_fails_closed() {
let head_dim: usize = 32;
let hidden: usize = 64;
let q_dim = head_dim; let kv_dim = 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: 1,
num_key_value_heads: 1,
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,
};
let rope = RopeTable::new(cfg.rope_dim(), 512, cfg.rope_theta);
let identity = |rows: usize| {
let mut mat = vec![0.0f32; rows * hidden];
for j in 0..rows.min(hidden) {
mat[j * hidden + j] = 1.0;
}
pack_weights_q8(&mat, rows, hidden)
};
let weights = Q8NeonFullAttnWeights {
q_proj_packed: identity(2 * q_dim),
q_proj_rows: 2 * q_dim,
q_proj_cols: hidden,
k_proj_packed: identity(kv_dim),
k_proj_rows: kv_dim,
k_proj_cols: hidden,
v_proj_packed: identity(kv_dim),
v_proj_rows: kv_dim,
v_proj_cols: hidden,
o_proj_packed: zero_packed(hidden, q_dim),
o_proj_rows: hidden,
o_proj_cols: 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.05).collect();
let mut scratch = ForwardScratch::new();
scratch.ensure_capacity(&cfg, 4);
let mut kv_cache = KvCache::new(1);
kv_cache.reserve(4, kv_dim);
scratch.attn_out[..hidden].copy_from_slice(&input);
full_attention_step_q8_neon(
&weights,
0,
0,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
);
kv_cache.seq_len = 1;
kv_cache.k[0][0] = f32::NAN;
scratch.attn_out[..hidden].copy_from_slice(&input);
full_attention_step_q8_neon(
&weights,
0,
1,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
);
for (d, &v) in scratch.context[..q_dim].iter().enumerate() {
assert!(
v.is_finite(),
"context[{d}] = {v} is non-finite; NEON Q8 softmax must fail \
closed on a NaN score row"
);
}
}
#[test]
fn test_gdn_step_q8_neon_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 mut weights = Q8NeonGdnWeights {
in_proj_qkv_packed: zero_packed(qkv_dim, hidden),
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z_packed: zero_packed(output_dim, hidden),
in_proj_z_rows: output_dim,
in_proj_z_cols: hidden,
in_proj_b_packed: zero_packed(num_heads, hidden),
in_proj_b_rows: num_heads,
in_proj_b_cols: hidden,
in_proj_a_packed: zero_packed(num_heads, hidden),
in_proj_a_rows: num_heads,
in_proj_a_cols: hidden,
out_proj_packed: zero_packed(hidden, output_dim),
out_proj_rows: hidden,
out_proj_cols: output_dim,
a_log: vec![0.0; num_heads],
dt_bias: vec![0.0; num_heads],
conv1d_weight: vec![0.0; qkv_dim * cfg.linear_conv_kernel_dim],
conv_dim: qkv_dim,
kernel_size: cfg.linear_conv_kernel_dim,
norm_weight: vec![0.0; cfg.linear_value_head_dim],
};
weights.a_log[0] = 100.0;
weights.dt_bias[0] = -100.0;
let mut state = GatedDeltaNetState::new(&cfg);
let input = vec![0.05f32; hidden];
let mut output = vec![0.0f32; hidden];
let mut gdn_scratch = GatedDeltaNetFusedScratch::default();
let mut x_q_scratch = Vec::new();
gdn_step_q8_neon(
&input,
&mut state,
&weights,
&cfg,
&mut gdn_scratch,
&mut x_q_scratch,
&mut output,
);
for (i, &v) in output[..hidden].iter().enumerate() {
assert!(
v.is_finite(),
"output[{i}] = {v} non-finite; decay gate must not overflow to NaN"
);
}
for (i, &v) in state.s_matrices.iter().enumerate() {
assert!(
v.is_finite(),
"state.s_matrices[{i}] = {v} non-finite; decay gate overflow poisoned state"
);
}
}
#[test]
fn test_generate_q8_neon_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 model = Q8NeonModel {
embed_tokens: vec![],
final_norm: vec![],
lm_head_packed: vec![],
lm_head_rows: 0,
lm_head_cols: 0,
layers: vec![],
};
let gen_cfg = GenerateConfig {
max_new_tokens: usize::MAX,
..Default::default()
};
let err = generate_q8_neon(&model, &cfg, &tokenizer, &rope, "hello", &gen_cfg)
.expect_err("expected context-window rejection");
let msg = format!("{err}");
assert!(
msg.contains("context window"),
"error should mention context window, got: {msg}"
);
}
}