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, silu_inplace};
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::tokenizer::common::Tokenizer;
use crate::vision::multimodal::Qwen35VisionRequest;
use crate::weights::f16_weights::{
F16AttentionWeights, F16FeedForwardWeights, F16FullAttentionLayerWeights,
F16GatedDeltaNetWeights, F16ModelWeights, F16MoeLayerWeights, f16_to_f32_slice, matmul_bt_f16,
};
#[inline]
pub fn gated_delta_net_step_fused_f16(
input: &[f32],
state: &mut GatedDeltaNetState,
weights: &F16GatedDeltaNetWeights,
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!(input.len() >= hidden);
debug_assert!(output.len() >= hidden);
scratch.ensure_capacity(qkv_dim, output_dim, value_heads, key_dim, value_dim);
matmul_bt_f16(
input,
&weights.in_proj_qkv,
&mut scratch.qkv_proj[..qkv_dim],
1,
hidden,
qkv_dim,
);
matmul_bt_f16(
input,
&weights.in_proj_z,
&mut scratch.z_proj[..output_dim],
1,
hidden,
output_dim,
);
matmul_bt_f16(
input,
&weights.in_proj_b,
&mut scratch.beta_proj[..value_heads],
1,
hidden,
value_heads,
);
matmul_bt_f16(
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_f16(
&scratch.gated_norm_buf[..output_dim],
&weights.out_proj,
&mut output[..hidden],
1,
output_dim,
hidden,
);
}
fn full_attention_step_f16(
weights: &F16FullAttentionLayerWeights,
cache_idx: usize,
position: usize,
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
cfg: &Qwen35Config,
rope: &RopeTable,
hidden: usize,
mrope_cos_sin: Option<(&[f32], &[f32])>,
) {
let input: Vec<f32> = scratch.attn_out[..hidden].to_vec();
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 mut q_and_gate = vec![0.0f32; q_proj_dim];
matmul_bt_f16(
&input,
&weights.q_proj,
&mut q_and_gate,
1,
hidden,
q_proj_dim,
);
let mut gate_z = vec![0.0f32; q_dim];
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(&q_and_gate[src..src + head_dim]);
gate_z[dst..dst + head_dim]
.copy_from_slice(&q_and_gate[src + head_dim..src + head_dim * 2]);
}
matmul_bt_f16(
&input,
&weights.k_proj,
&mut scratch.k_buf[..kv_dim],
1,
hidden,
kv_dim,
);
matmul_bt_f16(
&input,
&weights.v_proj,
&mut scratch.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;
if let Some((cos_row, sin_row)) = mrope_cos_sin {
for i in 0..half {
let cos_val = cos_row[i];
let sin_val = sin_row[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;
}
} else {
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;
if let Some((cos_row, sin_row)) = mrope_cos_sin {
for i in 0..half {
let cos_val = cos_row[i];
let sin_val = sin_row[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;
}
} else {
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;
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];
}
scratch.scores[scores_start + t] = dot * scale;
}
let row = &mut scratch.scores[scores_start..scores_start + cur_seq_len];
let (max_score, any_nan) = crate::attention::softmax_row::row_max_and_any_nan(row);
if crate::attention::softmax_row::row_fails_closed_pre_exp(max_score, any_nan) {
row.fill(0.0);
} else {
let mut sum_exp = 0.0f32;
for v in row.iter_mut() {
*v = (*v - max_score).exp();
sum_exp += *v;
}
crate::attention::softmax_row::finalize_row(row, 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;
}
}
for (ctx, &gz) in scratch.context[..q_dim].iter_mut().zip(&gate_z[..q_dim]) {
let sig = 1.0 / (1.0 + (-gz).exp());
*ctx *= sig;
}
matmul_bt_f16(
&scratch.context[..q_dim],
&weights.o_proj,
&mut scratch.attn_out[..hidden],
1,
q_dim,
hidden,
);
}
#[inline]
fn ffn_step_f16(
gate_proj: &[u16],
up_proj: &[u16],
down_proj: &[u16],
scratch: &mut ForwardScratch,
inter: usize,
hidden: usize,
) {
scratch.input_tmp[..hidden].copy_from_slice(&scratch.ffn_out[..hidden]);
matmul_bt_f16(
&scratch.input_tmp[..hidden],
gate_proj,
&mut scratch.gate_buf[..inter],
1,
hidden,
inter,
);
matmul_bt_f16(
&scratch.input_tmp[..hidden],
up_proj,
&mut scratch.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_f16(
&scratch.gate_buf[..inter],
down_proj,
&mut scratch.ffn_out[..hidden],
1,
inter,
hidden,
);
}
#[inline]
fn moe_ffn_step_f16(moe: &F16MoeLayerWeights, scratch: &mut ForwardScratch, hidden: usize) {
let inter = moe.experts.intermediate_size;
let shared_inter = moe.shared_expert.intermediate_size;
let num_experts = moe.router.num_experts;
let top_k = moe.router.num_experts_per_tok;
debug_assert_eq!(moe.router.hidden_size, hidden);
debug_assert_eq!(moe.experts.num_experts, num_experts);
debug_assert_eq!(moe.experts.hidden_size, hidden);
debug_assert_eq!(moe.shared_expert.hidden_size, hidden);
scratch.input_tmp[..hidden].copy_from_slice(&scratch.ffn_out[..hidden]);
if scratch.router_logits.len() < num_experts {
scratch.router_logits.resize(num_experts, 0.0);
}
if scratch.router_selected.len() < top_k {
scratch.router_selected.resize(top_k, (usize::MAX, 0.0));
}
matmul_bt_f16(
&scratch.input_tmp[..hidden],
&moe.router.gate,
&mut scratch.router_logits[..num_experts],
1,
hidden,
num_experts,
);
let max_logit = scratch.router_logits[..num_experts]
.iter()
.copied()
.fold(f32::NEG_INFINITY, f32::max);
let mut denom = 0.0f32;
for v in &mut scratch.router_logits[..num_experts] {
*v = (*v - max_logit).exp();
denom += *v;
}
if denom > 0.0 {
for v in &mut scratch.router_logits[..num_experts] {
*v /= denom;
}
} else {
scratch.router_logits[..num_experts].fill(0.0);
}
for slot in &mut scratch.router_selected[..top_k] {
*slot = (usize::MAX, f32::NEG_INFINITY);
}
for (expert_id, prob) in scratch.router_logits[..num_experts]
.iter()
.copied()
.enumerate()
{
for rank in 0..top_k {
if prob > scratch.router_selected[rank].1 {
for shift in (rank + 1..top_k).rev() {
scratch.router_selected[shift] = scratch.router_selected[shift - 1];
}
scratch.router_selected[rank] = (expert_id, prob);
break;
}
}
}
let top_sum: f32 = scratch.router_selected[..top_k]
.iter()
.map(|(_, p)| *p)
.sum();
if top_sum > 0.0 {
for (_, prob) in &mut scratch.router_selected[..top_k] {
*prob /= top_sum;
}
}
scratch.expert_out[..hidden].fill(0.0);
for idx in 0..top_k {
let (expert_id, weight) = scratch.router_selected[idx];
if expert_id >= moe.experts.num_experts {
continue;
}
debug_assert_ne!(expert_id, usize::MAX);
let gate_up_stride = 2 * inter * hidden;
let gate_up_start = expert_id * gate_up_stride;
let down_start = expert_id * hidden * inter;
let gate_w = &moe.experts.gate_up_proj[gate_up_start..gate_up_start + inter * hidden];
let up_w = &moe.experts.gate_up_proj
[gate_up_start + inter * hidden..gate_up_start + 2 * inter * hidden];
let down_w = &moe.experts.down_proj[down_start..down_start + hidden * inter];
matmul_bt_f16(
&scratch.input_tmp[..hidden],
gate_w,
&mut scratch.gate_buf[..inter],
1,
hidden,
inter,
);
matmul_bt_f16(
&scratch.input_tmp[..hidden],
up_w,
&mut scratch.up_buf[..inter],
1,
hidden,
inter,
);
silu_inplace(&mut scratch.gate_buf[..inter]);
elementwise_mul(&mut scratch.gate_buf[..inter], &scratch.up_buf[..inter]);
scratch.down_input[..inter].copy_from_slice(&scratch.gate_buf[..inter]);
matmul_bt_f16(
&scratch.down_input[..inter],
down_w,
&mut scratch.ffn_out[..hidden],
1,
inter,
hidden,
);
for i in 0..hidden {
scratch.expert_out[i] += weight * scratch.ffn_out[i];
}
}
let shared = &moe.shared_expert;
let mut shared_gate_logit = [0.0f32; 1];
matmul_bt_f16(
&scratch.input_tmp[..hidden],
&shared.shared_expert_gate,
&mut shared_gate_logit,
1,
hidden,
1,
);
let shared_gate = sigmoid(shared_gate_logit[0]);
matmul_bt_f16(
&scratch.input_tmp[..hidden],
&shared.gate_proj,
&mut scratch.gate_buf[..shared_inter],
1,
hidden,
shared_inter,
);
matmul_bt_f16(
&scratch.input_tmp[..hidden],
&shared.up_proj,
&mut scratch.up_buf[..shared_inter],
1,
hidden,
shared_inter,
);
silu_inplace(&mut scratch.gate_buf[..shared_inter]);
elementwise_mul(
&mut scratch.gate_buf[..shared_inter],
&scratch.up_buf[..shared_inter],
);
scratch.down_input[..shared_inter].copy_from_slice(&scratch.gate_buf[..shared_inter]);
matmul_bt_f16(
&scratch.down_input[..shared_inter],
&shared.down_proj,
&mut scratch.ffn_out[..hidden],
1,
shared_inter,
hidden,
);
for i in 0..hidden {
scratch.expert_out[i] += shared_gate * scratch.ffn_out[i];
}
scratch.ffn_out[..hidden].copy_from_slice(&scratch.expert_out[..hidden]);
}
pub(crate) fn forward_step_f16(
weights: &F16ModelWeights,
cfg: &Qwen35Config,
rope: &RopeTable,
token_id: u32,
position: usize,
gdn_states: &mut [GatedDeltaNetState],
kv_cache: &mut KvCache,
scratch: &mut ForwardScratch,
injected_embedding: Option<&[f32]>,
mrope_cos_sin: Option<(&[f32], &[f32])>,
) -> Result<(), crate::error::InferenceError> {
let hidden = cfg.hidden_size;
scratch.ensure_capacity(cfg, kv_cache.seq_len + 1);
match injected_embedding {
Some(row) => {
if row.len() != hidden {
return Err(crate::error::InferenceError::InvalidInput(format!(
"injected_embedding length {} does not match hidden_size {hidden}",
row.len()
)));
}
if let Some(bad) = row.iter().find(|v| !v.is_finite()) {
return Err(crate::error::InferenceError::InvalidInput(format!(
"injected_embedding contains a non-finite value: {bad}"
)));
}
scratch.hidden[..hidden].copy_from_slice(row);
}
None => {
let embed_start = token_id as usize * hidden;
f16_to_f32_slice(
&weights.embed_tokens[embed_start..embed_start + hidden],
&mut scratch.hidden[..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 {
F16AttentionWeights::Linear(gdn_w) => {
gated_delta_net_step_fused_f16(
&scratch.hidden[..hidden],
&mut gdn_states[linear_idx],
gdn_w,
cfg,
&mut scratch.gdn_scratch,
&mut scratch.attn_out[..hidden],
);
linear_idx += 1;
}
F16AttentionWeights::Full(full_w) => {
scratch.attn_out[..hidden].copy_from_slice(&scratch.hidden[..hidden]);
full_attention_step_f16(
full_w,
cache_idx_of(full_idx),
position,
kv_cache,
scratch,
cfg,
rope,
hidden,
mrope_cos_sin,
);
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]);
match &common.ffn {
F16FeedForwardWeights::Dense {
gate_proj,
up_proj,
down_proj,
} => {
ffn_step_f16(
gate_proj,
up_proj,
down_proj,
scratch,
cfg.intermediate_size,
hidden,
);
}
F16FeedForwardWeights::Moe(moe) => {
moe_ffn_step_f16(moe, scratch, 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_f16(
&scratch.hidden[..hidden],
&weights.embed_tokens,
&mut scratch.logits[..cfg.vocab_size],
1,
hidden,
cfg.vocab_size,
);
Ok(())
}
#[inline(always)]
fn cache_idx_of(full_idx: usize) -> usize {
full_idx
}
pub fn generate_f16(
weights: &F16ModelWeights,
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_f16(
weights,
cfg,
rope,
token_id,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
None,
None,
)?;
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 last_token = *all_ids
.last()
.expect("invariant: prompt or previous sample populated all_ids");
forward_step_f16(
weights,
cfg,
rope,
last_token,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
None,
None,
)?;
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,
stop_reason: Some(stop_reason),
token_logprobs: vec![],
})
}
pub fn generate_multimodal_f16(
weights: &F16ModelWeights,
cfg: &Qwen35Config,
request: &crate::vision::multimodal::Qwen35VisionRequest,
gen_cfg: &GenerateConfig,
) -> Result<GenerateOutput, crate::error::InferenceError> {
request.validate().map_err(|e| {
crate::error::InferenceError::InvalidInput(format!(
"multimodal request failed validation: {e}"
))
})?;
if let Some(&bad_id) = request
.input_ids
.iter()
.find(|&&id| id as usize >= cfg.vocab_size)
{
return Err(crate::error::InferenceError::InvalidInput(format!(
"input_ids contains out-of-vocabulary token id {bad_id} (vocab_size={})",
cfg.vocab_size
)));
}
let has_image = !request.image_grids.is_empty();
if has_image {
let cfg_image_token_id = cfg.image_token_id.ok_or_else(|| {
crate::error::InferenceError::InvalidInput(
"multimodal request supplied but checkpoint has no image_token_id".to_string(),
)
})?;
if cfg_image_token_id != request.image_token_id {
return Err(crate::error::InferenceError::InvalidInput(format!(
"request image_token_id {} does not match checkpoint image_token_id {cfg_image_token_id}",
request.image_token_id
)));
}
let vision_cfg = cfg.vision_config.as_ref().ok_or_else(|| {
crate::error::InferenceError::InvalidInput(
"multimodal request supplied but checkpoint has no vision_config".to_string(),
)
})?;
if vision_cfg.spatial_merge_size != request.spatial_merge_size {
return Err(crate::error::InferenceError::InvalidInput(format!(
"request spatial_merge_size {} does not match checkpoint \
vision_config.spatial_merge_size {}",
request.spatial_merge_size, vision_cfg.spatial_merge_size
)));
}
if request.decoder_hidden_size != cfg.hidden_size {
return Err(crate::error::InferenceError::InvalidInput(format!(
"request decoder_hidden_size {} does not match checkpoint hidden_size {}",
request.decoder_hidden_size, cfg.hidden_size
)));
}
if vision_cfg.out_hidden_size != cfg.hidden_size {
return Err(crate::error::InferenceError::InvalidInput(format!(
"checkpoint vision_config.out_hidden_size {} does not match decoder \
hidden_size {}",
vision_cfg.out_hidden_size, cfg.hidden_size
)));
}
}
let (positions, tables) = request.build_mrope_tables(cfg)?;
let expected_rope_half = cfg.rope_dim() / 2;
if tables.cos.iter().any(|row| row.len() != expected_rope_half)
|| tables.sin.iter().any(|row| row.len() != expected_rope_half)
{
return Err(crate::error::InferenceError::InvalidInput(format!(
"M-RoPE table row width does not match decoder rotary half-width: expected \
{expected_rope_half}"
)));
}
let prompt_ids = &request.input_ids;
let prompt_len = prompt_ids.len();
crate::model::qwen35::check_prompt_not_empty(prompt_len)?;
if gen_cfg.max_new_tokens == 0 {
return Ok(GenerateOutput {
text: String::new(),
token_ids: vec![],
prompt_tokens: prompt_len,
generated_tokens: 0,
stopped: false,
stop_reason: Some(StopReason::Length),
token_logprobs: vec![],
});
}
crate::model::qwen35::check_grammar_not_set(gen_cfg)?;
crate::model::qwen35::check_logprobs_not_set(gen_cfg)?;
crate::model::qwen35::check_stop_strings_not_set(gen_cfg)?;
crate::model::qwen35::check_reasoning_budget_not_set(gen_cfg)?;
let max_context = cfg.max_position_embeddings;
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 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 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 rope = RopeTable::new(cfg.rope_dim(), max_context, cfg.rope_theta);
let mut generated_ids: Vec<u32> = Vec::with_capacity(gen_cfg.max_new_tokens);
let mut all_ids = prompt_ids.clone();
let mut visual_row = 0usize;
for (pos, &token_id) in prompt_ids.iter().enumerate() {
let injected = if token_id == request.image_token_id {
let start = visual_row * request.decoder_hidden_size;
let end = start + request.decoder_hidden_size;
visual_row += 1;
Some(&request.post_merger_rows[start..end])
} else {
None
};
let cos_sin = if has_image {
Some((tables.cos[pos].as_slice(), tables.sin[pos].as_slice()))
} else {
None
};
forward_step_f16(
weights,
cfg,
&rope,
token_id,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
injected,
cos_sin,
)?;
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 physical_pos = kv_cache.seq_len;
let last_token = *all_ids
.last()
.expect("invariant: prompt or previous sample populated all_ids");
let decode_cos_sin;
let mrope_cos_sin = if has_image {
let decode_axis =
crate::vision::qwen35_mrope::decode_position(physical_pos, positions.rope_delta)?;
decode_cos_sin = request.build_decode_cos_sin(cfg, decode_axis)?;
if decode_cos_sin.0.len() != expected_rope_half
|| decode_cos_sin.1.len() != expected_rope_half
{
return Err(crate::error::InferenceError::InvalidInput(format!(
"decode-time M-RoPE row width does not match decoder rotary \
half-width: expected {expected_rope_half}"
)));
}
Some((decode_cos_sin.0.as_slice(), decode_cos_sin.1.as_slice()))
} else {
None
};
forward_step_f16(
weights,
cfg,
&rope,
last_token,
physical_pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
None,
mrope_cos_sin,
)?;
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);
}
Ok(GenerateOutput {
text: String::new(),
token_ids: generated_ids.clone(),
prompt_tokens: prompt_len,
generated_tokens: generated_ids.len(),
stopped,
stop_reason: Some(stop_reason),
token_logprobs: vec![],
})
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum PoolingStrategy {
MeanVisualTokens,
LastToken,
}
pub fn prefill_hidden_states_f16(
weights: &F16ModelWeights,
cfg: &Qwen35Config,
request: &Qwen35VisionRequest,
) -> Result<Vec<f32>, crate::error::InferenceError> {
request.validate().map_err(|e| {
crate::error::InferenceError::InvalidInput(format!(
"multimodal request failed validation: {e}"
))
})?;
if let Some(&bad_id) = request
.input_ids
.iter()
.find(|&&id| id as usize >= cfg.vocab_size)
{
return Err(crate::error::InferenceError::InvalidInput(format!(
"input_ids contains out-of-vocabulary token id {bad_id} (vocab_size={})",
cfg.vocab_size
)));
}
let has_image = !request.image_grids.is_empty();
if has_image {
let cfg_image_token_id = cfg.image_token_id.ok_or_else(|| {
crate::error::InferenceError::InvalidInput(
"multimodal request supplied but checkpoint has no image_token_id".to_string(),
)
})?;
if cfg_image_token_id != request.image_token_id {
return Err(crate::error::InferenceError::InvalidInput(format!(
"request image_token_id {} does not match checkpoint image_token_id {cfg_image_token_id}",
request.image_token_id
)));
}
let vision_cfg = cfg.vision_config.as_ref().ok_or_else(|| {
crate::error::InferenceError::InvalidInput(
"multimodal request supplied but checkpoint has no vision_config".to_string(),
)
})?;
if vision_cfg.spatial_merge_size != request.spatial_merge_size {
return Err(crate::error::InferenceError::InvalidInput(format!(
"request spatial_merge_size {} does not match checkpoint \
vision_config.spatial_merge_size {}",
request.spatial_merge_size, vision_cfg.spatial_merge_size
)));
}
if request.decoder_hidden_size != cfg.hidden_size {
return Err(crate::error::InferenceError::InvalidInput(format!(
"request decoder_hidden_size {} does not match checkpoint hidden_size {}",
request.decoder_hidden_size, cfg.hidden_size
)));
}
if vision_cfg.out_hidden_size != cfg.hidden_size {
return Err(crate::error::InferenceError::InvalidInput(format!(
"checkpoint vision_config.out_hidden_size {} does not match decoder \
hidden_size {}",
vision_cfg.out_hidden_size, cfg.hidden_size
)));
}
}
let prompt_len = request.input_ids.len();
crate::model::qwen35::check_prompt_not_empty(prompt_len)?;
let max_context = cfg.max_position_embeddings;
if prompt_len > max_context {
return Err(crate::error::InferenceError::Inference(format!(
"prompt ({prompt_len} tokens) exceeds model context window ({max_context})"
)));
}
let (_positions, tables) = request.build_mrope_tables(cfg)?;
let expected_rope_half = cfg.rope_dim() / 2;
if tables.cos.iter().any(|row| row.len() != expected_rope_half)
|| tables.sin.iter().any(|row| row.len() != expected_rope_half)
{
return Err(crate::error::InferenceError::InvalidInput(format!(
"M-RoPE table row width does not match decoder rotary half-width: expected \
{expected_rope_half}"
)));
}
let prompt_ids = &request.input_ids;
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 rope = RopeTable::new(cfg.rope_dim(), max_context, cfg.rope_theta);
let hidden = cfg.hidden_size;
let mut hidden_states: Vec<f32> = Vec::with_capacity(prompt_len * hidden);
let mut visual_row = 0usize;
for (pos, &token_id) in prompt_ids.iter().enumerate() {
let injected = if token_id == request.image_token_id {
let start = visual_row * request.decoder_hidden_size;
let end = start + request.decoder_hidden_size;
visual_row += 1;
Some(&request.post_merger_rows[start..end])
} else {
None
};
let cos_sin = if has_image {
Some((tables.cos[pos].as_slice(), tables.sin[pos].as_slice()))
} else {
None
};
forward_step_f16(
weights,
cfg,
&rope,
token_id,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
injected,
cos_sin,
)?;
hidden_states.extend_from_slice(&scratch.hidden[..hidden]);
if pos < prompt_len - 1 {
kv_cache.seq_len += 1;
}
}
kv_cache.seq_len = prompt_len;
Ok(hidden_states)
}
fn mean_pool_rows(hidden_states: &[f32], hidden_size: usize, positions: &[usize]) -> Vec<f32> {
debug_assert!(!positions.is_empty());
let mut out = vec![0.0f32; hidden_size];
for &p in positions {
let row = &hidden_states[p * hidden_size..(p + 1) * hidden_size];
for (o, &v) in out.iter_mut().zip(row) {
*o += v;
}
}
let n = positions.len() as f32;
for o in &mut out {
*o /= n;
}
out
}
fn l2_normalize_owned(mut v: Vec<f32>) -> Vec<f32> {
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 && norm.is_finite() {
for x in &mut v {
*x /= norm;
}
}
v
}
fn pool_hidden_states(
hidden_states: &[f32],
hidden_size: usize,
seq_len: usize,
image_pad_positions: &[usize],
pooling: PoolingStrategy,
) -> Vec<f32> {
match pooling {
PoolingStrategy::LastToken => {
hidden_states[(seq_len - 1) * hidden_size..seq_len * hidden_size].to_vec()
}
PoolingStrategy::MeanVisualTokens => {
if image_pad_positions.is_empty() {
let all: Vec<usize> = (0..seq_len).collect();
mean_pool_rows(hidden_states, hidden_size, &all)
} else {
mean_pool_rows(hidden_states, hidden_size, image_pad_positions)
}
}
}
}
pub fn embed_image_f16(
weights: &F16ModelWeights,
cfg: &Qwen35Config,
request: &Qwen35VisionRequest,
pooling: PoolingStrategy,
) -> Result<Vec<f32>, crate::error::InferenceError> {
let hidden_states = prefill_hidden_states_f16(weights, cfg, request)?;
let seq_len = request.input_ids.len();
let image_pad_positions: Vec<usize> = request
.input_ids
.iter()
.enumerate()
.filter(|&(_, &id)| id == request.image_token_id)
.map(|(i, _)| i)
.collect();
let pooled = pool_hidden_states(
&hidden_states,
cfg.hidden_size,
seq_len,
&image_pad_positions,
pooling,
);
Ok(l2_normalize_owned(pooled))
}
pub fn embed_text_vlm_f16(
weights: &F16ModelWeights,
cfg: &Qwen35Config,
tokenizer: &BpeTokenizer,
prompt: &str,
pooling: PoolingStrategy,
) -> Result<Vec<f32>, crate::error::InferenceError> {
let input = tokenizer.tokenize(prompt);
let prompt_ids: Vec<u32> = input.input_ids[..input.real_length].to_vec();
crate::model::qwen35::check_prompt_not_empty(prompt_ids.len())?;
if let Some(&bad_id) = prompt_ids.iter().find(|&&id| id as usize >= cfg.vocab_size) {
return Err(crate::error::InferenceError::InvalidInput(format!(
"prompt contains out-of-vocabulary token id {bad_id} (vocab_size={})",
cfg.vocab_size
)));
}
let image_token_id = cfg.image_token_id.unwrap_or(u32::MAX);
if prompt_ids.contains(&image_token_id) {
return Err(crate::error::InferenceError::InvalidInput(
"tokenized prompt unexpectedly contains the checkpoint's image_token_id".to_string(),
));
}
let request = Qwen35VisionRequest {
input_ids: prompt_ids,
image_grids: vec![],
post_merger_rows: vec![],
image_token_id,
spatial_merge_size: cfg
.vision_config
.as_ref()
.map(|v| v.spatial_merge_size)
.unwrap_or(2),
decoder_hidden_size: cfg.hidden_size,
};
embed_image_f16(weights, cfg, &request, pooling)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
#[allow(clippy::type_complexity)]
fn test_f16_forward_compiles() {
let cfg = Qwen35Config::qwen35_2b();
let _fn_ptr: fn(
&F16ModelWeights,
&Qwen35Config,
&RopeTable,
u32,
usize,
&mut [GatedDeltaNetState],
&mut KvCache,
&mut ForwardScratch,
Option<&[f32]>,
Option<(&[f32], &[f32])>,
) -> Result<(), crate::error::InferenceError> = forward_step_f16;
let _gdn_fn_ptr: fn(
&[f32],
&mut GatedDeltaNetState,
&F16GatedDeltaNetWeights,
&Qwen35Config,
&mut GatedDeltaNetFusedScratch,
&mut [f32],
) = gated_delta_net_step_fused_f16;
let _gen_fn_ptr: fn(
&F16ModelWeights,
&Qwen35Config,
&BpeTokenizer,
&RopeTable,
&str,
&GenerateConfig,
) -> Result<GenerateOutput, crate::error::InferenceError> = generate_f16;
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_full_attn_step_f16_rope_stride_half_parity() {
use crate::model::qwen35_config::LayerType;
use crate::weights::f16_weights::f32_to_f16_slice;
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 to_f16 = |src: &[f32]| -> Vec<u16> {
let mut dst = vec![0u16; src.len()];
f32_to_f16_slice(src, &mut dst);
dst
};
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 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 = F16FullAttentionLayerWeights {
q_proj: to_f16(&q_proj_f32),
k_proj: to_f16(&k_proj_f32),
v_proj: to_f16(&vec![0.0f32; kv_dim * hidden]),
o_proj: to_f16(&vec![0.0f32; 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_f16(
&weights,
0,
position,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
None,
);
let mut k_ref = vec![0.0f32; kv_dim];
matmul_bt_f16(&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 F16 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_f16(
&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 F16 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]
fn test_full_attn_step_f16_nan_cached_score_fails_closed() {
use crate::model::qwen35_config::LayerType;
use crate::weights::f16_weights::f32_to_f16_slice;
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 = 1;
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 to_f16 = |src: &[f32]| -> Vec<u16> {
let mut dst = vec![0u16; src.len()];
f32_to_f16_slice(src, &mut dst);
dst
};
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 mut v_proj_f32 = vec![0.0f32; kv_dim * hidden];
for j in 0..kv_dim {
v_proj_f32[j * hidden + j] = 1.0;
}
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 = F16FullAttentionLayerWeights {
q_proj: to_f16(&q_proj_f32),
k_proj: to_f16(&k_proj_f32),
v_proj: to_f16(&v_proj_f32),
o_proj: to_f16(&vec![0.0f32; 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);
let mut poisoned_k = vec![0.0f32; kv_dim];
poisoned_k[0] = f32::NAN;
let finite_v = vec![5.0f32; kv_dim];
kv_cache.append_kv(0, &poisoned_k, &finite_v);
kv_cache.seq_len = 1;
full_attention_step_f16(
&weights,
0,
position,
&mut kv_cache,
&mut scratch,
&cfg,
&rope,
hidden,
None,
);
assert!(
scratch.context[..head_dim].iter().all(|&v| v == 0.0),
"expected exact-zero context for a NaN-poisoned cached score, \
got {:?}",
&scratch.context[..head_dim]
);
}
#[test]
fn test_gdn_f16_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 weights = F16GatedDeltaNetWeights {
in_proj_qkv: vec![0u16; qkv_dim * hidden],
in_proj_qkv_rows: qkv_dim,
in_proj_qkv_cols: hidden,
in_proj_z: vec![0u16; output_dim * hidden],
in_proj_z_rows: output_dim,
in_proj_z_cols: hidden,
in_proj_b: vec![0u16; num_heads * hidden],
in_proj_b_rows: num_heads,
in_proj_b_cols: hidden,
in_proj_a: vec![0u16; num_heads * hidden],
in_proj_a_rows: num_heads,
in_proj_a_cols: 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; output_dim],
out_proj: vec![0u16; hidden * output_dim],
out_proj_rows: hidden,
out_proj_cols: 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_f16(
&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"
);
}
}
#[test]
fn test_moe_ffn_step_f16_nan_router_fails_closed_no_panic() {
use crate::weights::f16_weights::{
F16, F16MoeLayerWeights, F16MoeRouter, F16RoutedExperts, F16SharedExpert,
};
let num_experts = 4usize;
let hidden = 4usize;
let inter = 2usize;
let shared_inter = 2usize;
let top_k = 2usize;
let nan16 = F16::from_f32(f32::NAN).0;
let zeros = |n: usize| vec![F16::from_f32(0.0).0; n];
let router = F16MoeRouter::new(
vec![nan16; num_experts * hidden],
num_experts,
top_k,
hidden,
)
.unwrap();
let experts = F16RoutedExperts::new(
zeros(num_experts * 2 * inter * hidden),
zeros(num_experts * hidden * inter),
num_experts,
hidden,
inter,
)
.unwrap();
let shared = F16SharedExpert::new(
zeros(shared_inter * hidden),
zeros(shared_inter * hidden),
zeros(hidden * shared_inter),
zeros(hidden),
hidden,
shared_inter,
)
.unwrap();
let moe = F16MoeLayerWeights {
router,
experts,
shared_expert: shared,
};
let mut scratch = ForwardScratch::new();
let buf = inter.max(shared_inter);
scratch.ffn_out.resize(hidden, 1.0);
scratch.input_tmp.resize(hidden, 0.0);
scratch.expert_out.resize(hidden, 0.0);
scratch.gate_buf.resize(buf, 0.0);
scratch.up_buf.resize(buf, 0.0);
scratch.down_input.resize(buf, 0.0);
scratch.router_logits.resize(num_experts, 0.0);
scratch.router_selected.resize(top_k, (usize::MAX, 0.0));
moe_ffn_step_f16(&moe, &mut scratch, hidden);
assert!(
scratch.router_selected[..top_k]
.iter()
.all(|(id, _)| *id < num_experts),
"degenerate f16 router must not leave a usize::MAX sentinel selected"
);
}
#[test]
fn test_moe_ffn_step_f16_finite_max_nan_tail_fails_closed() {
use crate::weights::f16_weights::{
F16, F16MoeLayerWeights, F16MoeRouter, F16RoutedExperts, F16SharedExpert,
};
let num_experts = 4usize;
let hidden = 4usize;
let inter = 2usize;
let shared_inter = 2usize;
let top_k = 2usize;
let nan16 = F16::from_f32(f32::NAN).0;
let zeros = |n: usize| vec![F16::from_f32(0.0).0; n];
let mut gate = zeros(num_experts * hidden);
for w in &mut gate[..hidden] {
*w = nan16;
}
let router = F16MoeRouter::new(gate, num_experts, top_k, hidden).unwrap();
let experts = F16RoutedExperts::new(
zeros(num_experts * 2 * inter * hidden),
zeros(num_experts * hidden * inter),
num_experts,
hidden,
inter,
)
.unwrap();
let shared = F16SharedExpert::new(
zeros(shared_inter * hidden),
zeros(shared_inter * hidden),
zeros(hidden * shared_inter),
zeros(hidden),
hidden,
shared_inter,
)
.unwrap();
let moe = F16MoeLayerWeights {
router,
experts,
shared_expert: shared,
};
let mut scratch = ForwardScratch::new();
let buf = inter.max(shared_inter);
scratch.ffn_out.resize(hidden, 1.0);
scratch.input_tmp.resize(hidden, 0.0);
scratch.expert_out.resize(hidden, 0.0);
scratch.gate_buf.resize(buf, 0.0);
scratch.up_buf.resize(buf, 0.0);
scratch.down_input.resize(buf, 0.0);
scratch.router_logits.resize(num_experts, 0.0);
scratch.router_selected.resize(top_k, (usize::MAX, 0.0));
moe_ffn_step_f16(&moe, &mut scratch, hidden);
assert!(
scratch.router_logits[..num_experts]
.iter()
.all(|p| *p == 0.0),
"finite-max + NaN-tail router row must fail closed to all-zero probs \
(a max-only guard would miss this)"
);
}
fn zero_layer_f16_fixture() -> (Qwen35Config, F16ModelWeights, 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 = F16ModelWeights {
embed_tokens: vec![0u16; 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_f16_rejects_context_overflow() {
let (cfg, weights, rope, tokenizer) = zero_layer_f16_fixture();
let max_context = rope.max_positions(); let gen_cfg = GenerateConfig {
max_new_tokens: max_context,
..Default::default()
};
let err = generate_f16(&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 test_generate_f16_honors_stop_token_ids() {
let (cfg, weights, rope, tokenizer) = zero_layer_f16_fixture();
let gen_cfg = GenerateConfig {
max_new_tokens: 4,
stop_token_ids: vec![0], temperature: 0.0, ..Default::default()
};
let out = generate_f16(&weights, &cfg, &tokenizer, &rope, "h", &gen_cfg)
.expect("generate_f16 must succeed with valid stop_token_ids");
assert_eq!(
out.generated_tokens, 0,
"generate_f16 must stop immediately when the first greedy token (0) \
is in stop_token_ids — got {} generated tokens instead",
out.generated_tokens
);
}
#[test]
fn test_generate_f16_honors_stop_token_ids_decode_loop() {
use crate::weights::f16_weights::f32_to_f16_slice;
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 embed_f32: Vec<f32> = {
let mut v = vec![0.0f32; vocab * hidden];
v[0] = -1.0; v[1] = 1.0; v[hidden] = 1.0; v[hidden + 1] = 1.0; v
};
let mut embed_f16 = vec![0u16; vocab * hidden];
f32_to_f16_slice(&embed_f32, &mut embed_f16);
let weights = F16ModelWeights {
embed_tokens: embed_f16,
final_norm: vec![-2.0f32, 0.0, 0.0, 0.0],
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();
let gen_cfg = GenerateConfig {
max_new_tokens: 10,
stop_token_ids: vec![1], temperature: 0.0, ..Default::default()
};
let out = generate_f16(&weights, &cfg, &tokenizer, &rope, "e", &gen_cfg)
.expect("generate_f16 must succeed");
assert_eq!(
out.generated_tokens, 1,
"generate_f16 must stop at decode-loop step 1 when token 1 is in \
stop_token_ids — got {} tokens; reverting only the decode-loop check \
lets token 1 through and produces ≥ 2 tokens",
out.generated_tokens
);
assert!(
out.stopped,
"generate_f16 must set stopped=true when the decode-loop stop fires"
);
}
#[test]
fn generate_f16_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 = F16ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig::default();
let result = generate_f16(&weights, &cfg, &tokenizer, &rope, "", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::Inference(ref msg)) if msg.contains("empty prompt")),
"generate_f16 must reject an empty prompt with Err(Inference(\"empty \
prompt\")) (#856); got {result:?}"
);
}
#[test]
fn generate_f16_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 = F16ModelWeights {
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_f16(&weights, &cfg, &tokenizer, &rope, "hello", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"generate_f16 must fail closed with InvalidInput when grammar is set (#397/#398); \
got {result:?}"
);
}
#[test]
fn generate_f16_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 = F16ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig {
stop_strings: vec!["</s>".to_string()],
..Default::default()
};
let result = generate_f16(&weights, &cfg, &tokenizer, &rope, "hello", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"generate_f16 must fail closed with InvalidInput when stop_strings is set \
(ADR-080 C3, #783); got {result:?}"
);
}
#[test]
fn generate_f16_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 = F16ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig {
reasoning_budget: Some(16),
..Default::default()
};
let result = generate_f16(&weights, &cfg, &tokenizer, &rope, "hello", &gen_cfg);
assert!(
matches!(result, Err(InferenceError::InvalidInput(_))),
"generate_f16 must fail closed with InvalidInput when reasoning_budget is set \
(ADR-080 C3, #783); got {result:?}"
);
}
#[test]
fn generate_f16_max_new_tokens_zero_returns_empty() {
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 = F16ModelWeights {
embed_tokens: vec![],
final_norm: vec![],
layers: vec![],
};
let gen_cfg = GenerateConfig {
max_new_tokens: 0,
..Default::default()
};
let out = generate_f16(&weights, &cfg, &tokenizer, &rope, "hello", &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.stop_reason,
Some(StopReason::Length),
"max_new_tokens=0 must report stop_reason=Length"
);
}
use crate::model::qwen35_config::{LayerType, RopeParams};
use crate::weights::f16_weights::{
F16CommonLayerWeights, F16FullAttentionLayerWeights, f32_to_f16_slice,
};
fn tiny_vision_splice_model() -> (Qwen35Config, F16ModelWeights) {
let hidden = 8usize;
let vocab = 4usize;
let cfg = Qwen35Config {
hidden_size: hidden,
num_hidden_layers: 1,
vocab_size: vocab,
intermediate_size: 4,
rms_norm_eps: 1e-6,
num_attention_heads: 1,
num_key_value_heads: 1,
head_dim: hidden,
rope_theta: 1.0e7,
partial_rotary_factor: 1.0, rope_parameters: Some(RopeParams {
rope_theta: 1.0e7,
partial_rotary_factor: Some(1.0),
mrope_section: Some(vec![2, 1, 1]),
mrope_interleaved: Some(true),
}),
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: 1,
layer_types: vec![LayerType::FullAttention],
layer_mask: vec![true],
eos_token_id: 999,
max_position_embeddings: 512,
mtp_num_hidden_layers: 0,
mtp_use_dedicated_embeddings: false,
quarot_rotation_seed: None,
vision_config: None,
image_token_id: Some(3),
video_token_id: None,
vision_start_token_id: None,
vision_end_token_id: None,
};
let to_f16 = |src: &[f32]| -> Vec<u16> {
let mut dst = vec![0u16; src.len()];
f32_to_f16_slice(src, &mut dst);
dst
};
let identity = |rows: usize, cols: usize| -> Vec<f32> {
let mut m = vec![0.0f32; rows * cols];
for i in 0..rows.min(cols) {
m[i * cols + i] = 1.0;
}
m
};
let embed_tokens_f32: Vec<f32> = (0..vocab * hidden)
.map(|k| ((k % 11) as f32) * 0.05 - 0.2)
.collect();
let q_dim = hidden;
let mut q_proj_f32 = vec![0.0f32; 2 * q_dim * hidden];
q_proj_f32[..q_dim * hidden].copy_from_slice(&identity(q_dim, hidden));
let full_weights = F16FullAttentionLayerWeights {
q_proj: to_f16(&q_proj_f32),
k_proj: to_f16(&identity(hidden, hidden)),
v_proj: to_f16(&identity(hidden, hidden)),
o_proj: to_f16(&identity(hidden, q_dim)),
q_norm: vec![0.0f32; hidden],
k_norm: vec![0.0f32; hidden],
};
let common = F16CommonLayerWeights {
input_layernorm: vec![0.0f32; hidden],
post_attention_layernorm: vec![0.0f32; hidden],
ffn: F16FeedForwardWeights::Dense {
gate_proj: to_f16(&vec![0.0f32; 4 * hidden]),
up_proj: to_f16(&vec![0.0f32; 4 * hidden]),
down_proj: to_f16(&vec![0.0f32; hidden * 4]),
},
};
let weights = F16ModelWeights {
embed_tokens: to_f16(&embed_tokens_f32),
final_norm: vec![0.0f32; hidden],
layers: vec![(F16AttentionWeights::Full(full_weights), common)],
};
(cfg, weights)
}
#[test]
fn injection_replaces_lookup_and_is_mutation_sensitive() {
let (cfg, mut weights) = tiny_vision_splice_model();
let hidden = cfg.hidden_size;
let run = |weights: &F16ModelWeights, injected: Option<&[f32]>| -> Vec<f32> {
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let mut gdn_states: Vec<GatedDeltaNetState> = vec![];
let mut kv_cache = KvCache::new(cfg.num_full_attention_layers());
let mut scratch = ForwardScratch::new();
forward_step_f16(
weights,
&cfg,
&rope,
0,
0,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
injected,
None,
)
.expect("forward step succeeds");
scratch.hidden[..hidden].to_vec()
};
let baseline = run(&weights, None);
let mut visual_row = vec![0.9f32, -0.4, 0.2, 0.6, -0.1, 0.3, 0.7, -0.8];
let injected_hidden = run(&weights, Some(&visual_row));
assert_ne!(
injected_hidden, baseline,
"an injected embedding must produce a different hidden state than the token-id lookup"
);
visual_row[0] += 1.0;
let mutated_visual_hidden = run(&weights, Some(&visual_row));
assert_ne!(
mutated_visual_hidden, injected_hidden,
"mutating the supplied visual row must change the injected-slot hidden state"
);
visual_row[0] -= 1.0;
let embed_start = hidden; let mutated_row_f32 = vec![1.0f32; hidden];
f32_to_f16_slice(
&mutated_row_f32,
&mut weights.embed_tokens[embed_start..embed_start + hidden],
);
let after_table_mutation = run(&weights, Some(&visual_row));
assert_eq!(
after_table_mutation, injected_hidden,
"mutating an unrelated embedding-table row must not affect the injected slot"
);
let mut gdn_states: Vec<GatedDeltaNetState> = vec![];
let mut kv_cache = KvCache::new(cfg.num_full_attention_layers());
let mut scratch = ForwardScratch::new();
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let short_row = vec![0.0f32; hidden - 1];
assert!(
forward_step_f16(
&weights,
&cfg,
&rope,
0,
0,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
Some(&short_row),
None,
)
.is_err(),
"wrong-length injected_embedding must be rejected"
);
let mut nan_row = vec![0.0f32; hidden];
nan_row[2] = f32::NAN;
assert!(
forward_step_f16(
&weights,
&cfg,
&rope,
0,
0,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
Some(&nan_row),
None,
)
.is_err(),
"non-finite injected_embedding must be rejected"
);
}
#[test]
fn mrope_cos_sin_applied_in_wired_forward_matches_builder() {
use crate::vision::qwen35_mrope::{MRopePositions, build_cos_sin};
let (cfg, weights) = tiny_vision_splice_model();
let hidden = cfg.hidden_size;
let rope_params = cfg.rope_parameters.as_ref().unwrap();
let mrope_section = rope_params.mrope_section.as_ref().unwrap();
let positions = MRopePositions {
positions: vec![(2, 3, 5)],
rope_delta: 0,
};
let tables = build_cos_sin(
&positions,
cfg.head_dim,
rope_params.partial_rotary_factor.unwrap(),
rope_params.rope_theta as f32,
mrope_section,
)
.expect("builds tables");
let (cos_row, sin_row) = (tables.cos[0].as_slice(), tables.sin[0].as_slice());
let rope = RopeTable::new(cfg.rope_dim(), 8, cfg.rope_theta);
let mut gdn_states: Vec<GatedDeltaNetState> = vec![];
let mut kv_cache = KvCache::new(cfg.num_full_attention_layers());
let mut scratch = ForwardScratch::new();
forward_step_f16(
&weights,
&cfg,
&rope,
1,
0,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
None,
Some((cos_row, sin_row)),
)
.expect("forward step succeeds");
let mut k_ref = vec![0.0f32; hidden];
f16_to_f32_slice(&weights.embed_tokens[hidden..2 * hidden], &mut k_ref);
let zero_norm = vec![0.0f32; hidden];
qwen35_rms_norm(&mut k_ref, &zero_norm, hidden, cfg.rms_norm_eps);
let half = cfg.rope_dim() / 2;
for i in 0..half {
let x0 = k_ref[i];
let x1 = k_ref[half + i];
k_ref[i] = x0 * cos_row[i] - x1 * sin_row[i];
k_ref[half + i] = x0 * sin_row[i] + x1 * cos_row[i];
}
let k_cached = &kv_cache.k[0][..hidden];
let max_diff = k_cached
.iter()
.zip(k_ref.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(
max_diff < 1e-3,
"wired M-RoPE rotation diverges from build_cos_sin's own row: max_diff={max_diff}"
);
}
#[test]
fn generate_multimodal_text_only_matches_plain_forward_bit_identical() {
use crate::vision::multimodal::Qwen35VisionRequest;
let (cfg, weights) = tiny_vision_splice_model();
let input_ids: Vec<u32> = vec![0, 1, 2, 0];
let request = Qwen35VisionRequest {
input_ids: input_ids.clone(),
image_grids: vec![],
post_merger_rows: vec![],
image_token_id: 3,
spatial_merge_size: 2,
decoder_hidden_size: cfg.hidden_size,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 2,
temperature: 0.0,
seed: Some(1),
stop_token_ids: vec![],
..Default::default()
};
let multimodal_out = generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg)
.expect("text-only multimodal generate succeeds");
let rope = RopeTable::new(cfg.rope_dim(), 512, cfg.rope_theta);
let mut gdn_states: Vec<GatedDeltaNetState> = vec![];
let mut kv_cache = KvCache::new(cfg.num_full_attention_layers());
let mut scratch = ForwardScratch::new();
for (pos, &token_id) in input_ids.iter().enumerate() {
forward_step_f16(
&weights,
&cfg,
&rope,
token_id,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
None,
None,
)
.expect("reference forward step succeeds");
if pos < input_ids.len() - 1 {
kv_cache.seq_len += 1;
}
}
kv_cache.seq_len = input_ids.len();
let mut rng_state = 1u64;
let mut all_ids = input_ids.clone();
let mut ref_ids = Vec::new();
let next_id = sample_token(
&scratch.logits[..cfg.vocab_size],
&gen_cfg,
&all_ids,
&mut rng_state,
);
ref_ids.push(next_id);
all_ids.push(next_id);
for _ in 1..gen_cfg.max_new_tokens {
let pos = kv_cache.seq_len;
let last_token = *all_ids.last().unwrap();
forward_step_f16(
&weights,
&cfg,
&rope,
last_token,
pos,
&mut gdn_states,
&mut kv_cache,
&mut scratch,
None,
None,
)
.expect("reference decode step succeeds");
kv_cache.seq_len += 1;
let next_id = sample_token(
&scratch.logits[..cfg.vocab_size],
&gen_cfg,
&all_ids,
&mut rng_state,
);
ref_ids.push(next_id);
all_ids.push(next_id);
}
assert_eq!(
multimodal_out.token_ids, ref_ids,
"text-only generate_multimodal_f16 token ids must match the plain forward_step_f16 \
reference bit-for-bit"
);
}
#[test]
fn generate_multimodal_f16_rejects_invalid_request() {
use crate::vision::multimodal::Qwen35VisionRequest;
let (cfg, weights) = tiny_vision_splice_model();
let request = Qwen35VisionRequest {
input_ids: vec![0, 3, 3, 1], image_grids: vec![crate::vision::qwen35_vit::GridThw { t: 1, h: 4, w: 4 }], post_merger_rows: vec![0.0f32; 2 * cfg.hidden_size],
image_token_id: 3,
spatial_merge_size: 2,
decoder_hidden_size: cfg.hidden_size,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
..Default::default()
};
assert!(
generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg).is_err(),
"a request whose image-pad count does not match its grid must be rejected"
);
}
#[test]
fn generate_multimodal_f16_rejects_context_overflow() {
use crate::vision::multimodal::Qwen35VisionRequest;
let (mut cfg, weights) = tiny_vision_splice_model();
cfg.max_position_embeddings = 3;
let request = Qwen35VisionRequest {
input_ids: vec![0, 1, 2],
image_grids: vec![],
post_merger_rows: vec![],
image_token_id: 3,
spatial_merge_size: 2,
decoder_hidden_size: cfg.hidden_size,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 5,
..Default::default()
};
assert!(
generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg).is_err(),
"prompt_len + max_new_tokens exceeding max_position_embeddings must be rejected"
);
}
#[test]
fn generate_multimodal_f16_rejects_out_of_vocab_input_id() {
use crate::vision::multimodal::Qwen35VisionRequest;
let (cfg, weights) = tiny_vision_splice_model();
let request = Qwen35VisionRequest {
input_ids: vec![0, 1, cfg.vocab_size as u32], image_grids: vec![],
post_merger_rows: vec![],
image_token_id: 3,
spatial_merge_size: 2,
decoder_hidden_size: cfg.hidden_size,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
..Default::default()
};
let err = generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg)
.expect_err("an out-of-vocabulary input_id must be rejected, not panic");
assert!(
matches!(err, crate::error::InferenceError::InvalidInput(_)),
"expected InvalidInput, got {err:?}"
);
}
#[test]
fn test_generate_f16_rejects_out_of_vocab_prompt_id() {
use std::collections::HashMap;
let (cfg, weights, rope, _tokenizer) = zero_layer_f16_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_f16(&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:?}"
);
}
fn tiny_vision_splice_model_with_vision_cfg() -> (Qwen35Config, F16ModelWeights) {
use crate::model::qwen35_config::VisionModelConfig;
let (mut cfg, weights) = tiny_vision_splice_model();
cfg.vision_config = Some(VisionModelConfig {
depth: 1,
hidden_size: 8,
num_heads: 1,
patch_size: 1,
spatial_merge_size: 2,
out_hidden_size: cfg.hidden_size,
temporal_patch_size: 1,
num_position_embeddings: 1,
in_channels: 3,
deepstack_visual_indexes: vec![],
intermediate_size: None,
});
(cfg, weights)
}
#[test]
fn generate_multimodal_f16_rejects_mismatched_image_token_id() {
use crate::vision::multimodal::Qwen35VisionRequest;
use crate::vision::qwen35_vit::GridThw;
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
assert_eq!(cfg.image_token_id, Some(3));
let request = Qwen35VisionRequest {
input_ids: vec![0, 2, 2, 2, 2, 1],
image_grids: vec![GridThw { t: 1, h: 4, w: 4 }],
post_merger_rows: vec![0.1f32; 4 * cfg.hidden_size],
image_token_id: 2,
spatial_merge_size: 2,
decoder_hidden_size: cfg.hidden_size,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
..Default::default()
};
let err = generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg)
.expect_err("mismatched image_token_id must be rejected, not silently run");
let msg = format!("{err}");
assert!(
msg.contains("image_token_id"),
"error must name image_token_id; got: {msg}"
);
}
#[test]
fn generate_multimodal_f16_rejects_mismatched_spatial_merge_size() {
use crate::vision::multimodal::Qwen35VisionRequest;
use crate::vision::qwen35_vit::GridThw;
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
assert_eq!(cfg.vision_config.as_ref().unwrap().spatial_merge_size, 2);
let mut input_ids = vec![0u32];
input_ids.extend(std::iter::repeat_n(3u32, 16));
input_ids.push(1);
let request = Qwen35VisionRequest {
input_ids,
image_grids: vec![GridThw { t: 1, h: 4, w: 4 }],
post_merger_rows: vec![0.1f32; 16 * cfg.hidden_size],
image_token_id: 3,
spatial_merge_size: 1,
decoder_hidden_size: cfg.hidden_size,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
..Default::default()
};
let err = generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg)
.expect_err("mismatched spatial_merge_size must be rejected, not silently run");
let msg = format!("{err}");
assert!(
msg.contains("spatial_merge_size"),
"error must name spatial_merge_size; got: {msg}"
);
}
#[test]
fn generate_multimodal_f16_rejects_mismatched_decoder_hidden_size() {
use crate::vision::multimodal::Qwen35VisionRequest;
use crate::vision::qwen35_vit::GridThw;
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
assert_eq!(cfg.hidden_size, 8);
let request = Qwen35VisionRequest {
input_ids: vec![0, 3, 3, 3, 3, 1],
image_grids: vec![GridThw { t: 1, h: 4, w: 4 }],
post_merger_rows: vec![0.1f32; 4 * 4],
image_token_id: 3,
spatial_merge_size: 2,
decoder_hidden_size: 4,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
..Default::default()
};
let err = generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg)
.expect_err("mismatched decoder_hidden_size must be rejected, not silently run");
let msg = format!("{err}");
assert!(
msg.contains("decoder_hidden_size"),
"error must name decoder_hidden_size (not just the unrelated downstream \
injected-row-length message); got: {msg}"
);
}
#[test]
fn generate_multimodal_f16_rejects_mismatched_mrope_row_width() {
use crate::model::qwen35_config::RopeParams;
use crate::vision::multimodal::Qwen35VisionRequest;
let (mut cfg, weights) = tiny_vision_splice_model();
cfg.head_dim = 256;
cfg.partial_rotary_factor = 0.5; cfg.rope_parameters = Some(RopeParams {
rope_theta: 1.0e7,
partial_rotary_factor: Some(0.25), mrope_section: Some(vec![11, 11, 10]),
mrope_interleaved: Some(true),
});
let request = Qwen35VisionRequest {
input_ids: vec![0, 1, 2],
image_grids: vec![],
post_merger_rows: vec![],
image_token_id: 3,
spatial_merge_size: 2,
decoder_hidden_size: cfg.hidden_size,
};
let gen_cfg = GenerateConfig {
max_new_tokens: 1,
..Default::default()
};
let err = generate_multimodal_f16(&weights, &cfg, &request, &gen_cfg).expect_err(
"a config whose rope_parameters and rope_dim() disagree on rotary width must be \
rejected, not panic at attention-lane indexing",
);
assert!(
matches!(err, crate::error::InferenceError::InvalidInput(_)),
"expected InvalidInput, got {err:?}"
);
}
fn cosine(a: &[f32], b: &[f32]) -> f32 {
assert_eq!(a.len(), b.len());
let dot: f32 = a.iter().zip(b).map(|(x, y)| x * y).sum();
let na: f32 = a.iter().map(|x| x * x).sum::<f32>().sqrt();
let nb: f32 = b.iter().map(|x| x * x).sum::<f32>().sqrt();
if na == 0.0 || nb == 0.0 {
return 0.0;
}
dot / (na * nb)
}
fn one_image_request(visual_rows: Vec<f32>) -> Qwen35VisionRequest {
let mut input_ids = vec![0u32];
input_ids.extend(std::iter::repeat_n(3u32, 4));
input_ids.push(1);
Qwen35VisionRequest {
input_ids,
image_grids: vec![crate::vision::qwen35_vit::GridThw { t: 1, h: 4, w: 4 }],
post_merger_rows: visual_rows,
image_token_id: 3,
spatial_merge_size: 2,
decoder_hidden_size: 8,
}
}
#[test]
fn embed_image_f16_is_deterministic() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let request = one_image_request(vec![0.3f32; 4 * cfg.hidden_size]);
let v1 = embed_image_f16(&weights, &cfg, &request, PoolingStrategy::MeanVisualTokens)
.expect("embed_image_f16 succeeds");
let v2 = embed_image_f16(&weights, &cfg, &request, PoolingStrategy::MeanVisualTokens)
.expect("embed_image_f16 succeeds");
assert_eq!(v1, v2, "same input must produce an identical vector");
let v3 = embed_image_f16(&weights, &cfg, &request, PoolingStrategy::LastToken)
.expect("embed_image_f16 succeeds");
let v4 = embed_image_f16(&weights, &cfg, &request, PoolingStrategy::LastToken)
.expect("embed_image_f16 succeeds");
assert_eq!(
v3, v4,
"same input must produce an identical vector (LastToken)"
);
}
#[test]
fn embed_image_f16_is_finite_and_unit_norm() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
for pooling in [
PoolingStrategy::MeanVisualTokens,
PoolingStrategy::LastToken,
] {
let request = one_image_request(vec![0.4f32; 4 * cfg.hidden_size]);
let v = embed_image_f16(&weights, &cfg, &request, pooling)
.expect("embed_image_f16 succeeds");
assert_eq!(v.len(), cfg.hidden_size);
assert!(
v.iter().all(|x| x.is_finite()),
"{pooling:?}: non-finite output"
);
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-4,
"{pooling:?}: expected unit norm, got {norm}"
);
}
}
#[test]
fn embed_image_f16_discriminates_different_images_but_matches_itself() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let request_a = one_image_request(
(0..4 * cfg.hidden_size)
.map(|i| (i as f32) * 0.05 - 0.8)
.collect(),
);
let request_b = one_image_request(
(0..4 * cfg.hidden_size)
.map(|i| -(i as f32) * 0.03 + 0.5)
.collect(),
);
let emb_a = embed_image_f16(
&weights,
&cfg,
&request_a,
PoolingStrategy::MeanVisualTokens,
)
.expect("embed_image_f16 succeeds");
let emb_a_again = embed_image_f16(
&weights,
&cfg,
&request_a,
PoolingStrategy::MeanVisualTokens,
)
.expect("embed_image_f16 succeeds");
let emb_b = embed_image_f16(
&weights,
&cfg,
&request_b,
PoolingStrategy::MeanVisualTokens,
)
.expect("embed_image_f16 succeeds");
let self_cos = cosine(&emb_a, &emb_a_again);
assert!(
(self_cos - 1.0).abs() < 1e-5,
"an image embedded against itself must have cosine ~1.0, got {self_cos}"
);
let cross_cos = cosine(&emb_a, &emb_b);
assert!(
cross_cos < 0.999,
"two different images must not collapse to near-identical embeddings, got cosine {cross_cos}"
);
}
#[test]
fn pool_hidden_states_wrong_positions_change_the_output() {
let hidden_size = 4;
let seq_len = 6;
let hidden_states: Vec<f32> = (0..seq_len)
.flat_map(|i| std::iter::repeat_n(i as f32, hidden_size))
.collect();
let correct_positions = [1usize, 2, 3, 4];
let off_by_one_positions = [2usize, 3, 4, 5];
let correct = pool_hidden_states(
&hidden_states,
hidden_size,
seq_len,
&correct_positions,
PoolingStrategy::MeanVisualTokens,
);
let wrong = pool_hidden_states(
&hidden_states,
hidden_size,
seq_len,
&off_by_one_positions,
PoolingStrategy::MeanVisualTokens,
);
assert_ne!(
correct, wrong,
"pooling over an off-by-one-shifted position window must change the output"
);
}
#[test]
fn embed_image_f16_wrong_pad_run_placement_changes_embedding() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let visual_rows = vec![0.25f32; 4 * cfg.hidden_size];
let correct = one_image_request(visual_rows.clone());
let mut shifted_ids = vec![0u32, 2]; shifted_ids.extend(std::iter::repeat_n(3u32, 4));
shifted_ids.push(1);
let shifted = Qwen35VisionRequest {
input_ids: shifted_ids,
..one_image_request(visual_rows)
};
let emb_correct =
embed_image_f16(&weights, &cfg, &correct, PoolingStrategy::MeanVisualTokens)
.expect("embed_image_f16 succeeds");
let emb_shifted =
embed_image_f16(&weights, &cfg, &shifted, PoolingStrategy::MeanVisualTokens)
.expect("embed_image_f16 succeeds");
assert_ne!(
emb_correct, emb_shifted,
"shifting the image-pad run's position in input_ids must change the pooled embedding"
);
}
#[test]
fn embed_text_vlm_f16_is_deterministic_and_unit_norm() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let mut vocab_map = std::collections::HashMap::new();
for (i, c) in ["a", "b", "c"].iter().enumerate() {
vocab_map.insert((*c).to_string(), i as u32);
}
let tokenizer =
BpeTokenizer::from_vocab_and_merges(vocab_map, vec![]).expect("tokenizer constructs");
for pooling in [
PoolingStrategy::MeanVisualTokens,
PoolingStrategy::LastToken,
] {
let v1 = embed_text_vlm_f16(&weights, &cfg, &tokenizer, "abc", pooling)
.expect("embed_text_vlm_f16 succeeds");
let v2 = embed_text_vlm_f16(&weights, &cfg, &tokenizer, "abc", pooling)
.expect("embed_text_vlm_f16 succeeds");
assert_eq!(
v1, v2,
"{pooling:?}: same prompt must produce an identical vector"
);
assert_eq!(v1.len(), cfg.hidden_size);
let norm: f32 = v1.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!(
(norm - 1.0).abs() < 1e-4,
"{pooling:?}: expected unit norm, got {norm}"
);
}
}
#[test]
fn embed_text_vlm_f16_and_embed_image_f16_share_the_same_space() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let mut vocab_map = std::collections::HashMap::new();
vocab_map.insert("a".to_string(), 0u32);
let tokenizer =
BpeTokenizer::from_vocab_and_merges(vocab_map, vec![]).expect("tokenizer constructs");
let text_emb =
embed_text_vlm_f16(&weights, &cfg, &tokenizer, "a", PoolingStrategy::LastToken)
.expect("embed_text_vlm_f16 succeeds");
let image_emb = embed_image_f16(
&weights,
&cfg,
&one_image_request(vec![0.1f32; 4 * cfg.hidden_size]),
PoolingStrategy::LastToken,
)
.expect("embed_image_f16 succeeds");
assert_eq!(text_emb.len(), image_emb.len());
let cos = cosine(&text_emb, &image_emb);
assert!(cos.is_finite());
}
#[test]
fn embed_text_vlm_f16_rejects_prompt_colliding_with_image_token_id() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
assert_eq!(cfg.image_token_id, Some(3));
let mut vocab_map = std::collections::HashMap::new();
vocab_map.insert("z".to_string(), 3u32);
let tokenizer =
BpeTokenizer::from_vocab_and_merges(vocab_map, vec![]).expect("tokenizer constructs");
let err = embed_text_vlm_f16(&weights, &cfg, &tokenizer, "z", PoolingStrategy::LastToken)
.expect_err("a prompt colliding with image_token_id must be rejected");
assert!(matches!(err, crate::error::InferenceError::InvalidInput(_)));
}
#[test]
fn embed_image_f16_rejects_invalid_request() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let mut request = one_image_request(vec![0.1f32; 4 * cfg.hidden_size]);
request.post_merger_rows.pop(); assert!(
embed_image_f16(&weights, &cfg, &request, PoolingStrategy::MeanVisualTokens).is_err()
);
}
#[test]
fn prefill_rejects_over_context_before_mrope_table_construction() {
let (mut cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let base = one_image_request(vec![0.1f32; 4 * cfg.hidden_size]);
let request = Qwen35VisionRequest {
input_ids: vec![0u32, 3, 1, 3, 3, 3, 1],
image_grids: vec![
crate::vision::qwen35_vit::GridThw { t: 1, h: 2, w: 4 },
crate::vision::qwen35_vit::GridThw { t: 1, h: 2, w: 4 },
],
..base
};
request
.validate()
.expect("request must pass validation so only the builder would catch it");
cfg.max_position_embeddings = request.input_ids.len() - 1;
let err = prefill_hidden_states_f16(&weights, &cfg, &request)
.expect_err("over-context request must be rejected");
let msg = err.to_string();
assert!(
msg.contains("context window"),
"must fail on the context-window check, before M-RoPE table \
construction; got: {msg}"
);
}
#[test]
fn embed_image_f16_matches_independently_computed_golden_pool() {
let (cfg, weights) = tiny_vision_splice_model_with_vision_cfg();
let request = one_image_request(vec![0.37f32; 4 * cfg.hidden_size]);
let hidden_states = prefill_hidden_states_f16(&weights, &cfg, &request)
.expect("prefill_hidden_states_f16 succeeds");
let known_correct_pad_positions = [1usize, 2, 3, 4]; let golden = l2_normalize_owned(mean_pool_rows(
&hidden_states,
cfg.hidden_size,
&known_correct_pad_positions,
));
let got = embed_image_f16(&weights, &cfg, &request, PoolingStrategy::MeanVisualTokens)
.expect("embed_image_f16 succeeds");
assert_eq!(
got, golden,
"embed_image_f16 must pool over exactly the known-correct image-pad positions"
);
}
}