use std::collections::HashMap;
use crate::error::InferenceError;
use crate::model::qwen35::qwen_required_tensor_names;
use crate::model::qwen35_config::Qwen35Config;
use crate::quant::quarot::hadamard::RandomizedHadamard;
use crate::quant::quarot::lm_head::{
QWEN35_EMBED_TOKENS_NAME, QWEN35_FINAL_NORM_NAME, QWEN35_LM_HEAD_NAME,
qwen35_final_norm_fusion_target,
};
use crate::quant::quarot::pipeline::TensorEntry;
use crate::quant::quarot::plan::{AbsorptionSide, RotationPlan};
use crate::quant::quarot::rmsnorm_fusion::{
RmsNormFusionTarget, fuse_shifted_rmsnorm_into_next_layer_f64, qwen35_per_layer_fusion_plan,
};
use crate::quant::quarot::rotation::{absorb_input_rotation_f64, absorb_output_rotation_f64};
#[derive(Debug, Clone)]
pub struct ForwardEquivalenceConfig {
pub num_probe_tokens: usize,
pub tolerance: f64,
pub seed: u64,
}
impl Default for ForwardEquivalenceConfig {
fn default() -> Self {
Self {
num_probe_tokens: 4,
tolerance: 1e-5,
seed: 0xCAFE_BABE_DEAD_BEEF,
}
}
}
#[derive(Debug, Clone)]
pub struct ForwardEquivalenceReport {
pub max_abs_error: f64,
pub mean_abs_error: f64,
pub probe_tokens: Vec<u32>,
pub tolerance: f64,
}
pub fn assert_forward_equivalence_qwen35(
original: &HashMap<String, TensorEntry>,
rotated: &HashMap<String, TensorEntry>,
cfg: &Qwen35Config,
rotation: &RandomizedHadamard,
forward_cfg: &ForwardEquivalenceConfig,
) -> Result<ForwardEquivalenceReport, InferenceError> {
if cfg.is_moe() {
return Err(InferenceError::Inference(
"assert_forward_equivalence_qwen35: MoE configs are deferred to v1 \
(the rotation/fusion pipeline rejects MoE upstream; this probe \
has no expert-mixing path)"
.to_string(),
));
}
if forward_cfg.num_probe_tokens == 0 {
return Err(InferenceError::Inference(
"assert_forward_equivalence_qwen35: num_probe_tokens must be > 0".to_string(),
));
}
if !forward_cfg.tolerance.is_finite() || forward_cfg.tolerance <= 0.0 {
return Err(InferenceError::Inference(format!(
"assert_forward_equivalence_qwen35: tolerance must be a positive finite value, got {}",
forward_cfg.tolerance
)));
}
if cfg.vocab_size == 0 {
return Err(InferenceError::Inference(
"assert_forward_equivalence_qwen35: cfg.vocab_size must be > 0".to_string(),
));
}
if rotation.dim() != cfg.hidden_size {
return Err(InferenceError::Inference(format!(
"assert_forward_equivalence_qwen35: rotation.dim()={} != cfg.hidden_size={}",
rotation.dim(),
cfg.hidden_size
)));
}
let probe_tokens = deterministic_probe_tokens(
forward_cfg.seed,
forward_cfg.num_probe_tokens,
cfg.vocab_size,
);
let mut chain_max_abs = 0.0_f64;
let mut chain_total_abs = 0.0_f64;
let mut chain_count: usize = 0;
for &token in &probe_tokens {
let logits_orig = rotation_chain_probe_qwen35(original, cfg, token)?;
let logits_rot = rotation_chain_probe_qwen35(rotated, cfg, token)?;
if logits_orig.len() != logits_rot.len() {
return Err(InferenceError::Inference(format!(
"assert_forward_equivalence_qwen35: probe logits length mismatch \
(original={}, rotated={}) on token {token}",
logits_orig.len(),
logits_rot.len()
)));
}
for (a, b) in logits_orig.iter().zip(logits_rot.iter()) {
let d = (a - b).abs();
if d > chain_max_abs {
chain_max_abs = d;
}
chain_total_abs += d;
chain_count += 1;
}
}
let chain_mean_abs = if chain_count > 0 {
chain_total_abs / chain_count as f64
} else {
0.0
};
if chain_max_abs > forward_cfg.tolerance {
return Err(InferenceError::Inference(format!(
"forward-equivalence refused: chain probe max_abs_error={chain_max_abs} \
exceeds tolerance={} (mean_abs_error={chain_mean_abs}, probe_tokens={:?}). \
Do NOT write conversion artifacts — the pipeline produced logits \
that disagree with the original model.",
forward_cfg.tolerance, probe_tokens
)));
}
let rotation_plan = RotationPlan::qwen35_residual_stream_linear_layers();
let mut fusion_plan = qwen35_per_layer_fusion_plan(cfg)?;
fusion_plan.push(qwen35_final_norm_fusion_target());
let per_tensor_max_abs = check_per_tensor_rotation_equivalence(
original,
rotated,
cfg,
rotation,
&rotation_plan,
&fusion_plan,
)?;
if per_tensor_max_abs > forward_cfg.tolerance {
return Err(InferenceError::Inference(format!(
"forward-equivalence refused: per-tensor max_abs_error={per_tensor_max_abs} \
exceeds tolerance={} (chain probe max={chain_max_abs}, mean={chain_mean_abs}). \
At least one planned tensor disagrees with the rotation/fusion algebra. \
Do NOT write conversion artifacts.",
forward_cfg.tolerance
)));
}
let max_abs = chain_max_abs.max(per_tensor_max_abs);
Ok(ForwardEquivalenceReport {
max_abs_error: max_abs,
mean_abs_error: chain_mean_abs,
probe_tokens,
tolerance: forward_cfg.tolerance,
})
}
fn check_per_tensor_rotation_equivalence(
original: &HashMap<String, TensorEntry>,
rotated: &HashMap<String, TensorEntry>,
cfg: &Qwen35Config,
rotation: &RandomizedHadamard,
rotation_plan: &RotationPlan,
fusion_plan: &[RmsNormFusionTarget],
) -> Result<f64, InferenceError> {
let hidden = cfg.hidden_size;
let mut fusion_gamma: HashMap<&str, &str> = HashMap::new();
for target in fusion_plan {
for downstream in &target.downstream_weights {
fusion_gamma.insert(downstream.as_str(), target.norm_tensor.as_str());
}
}
let required = qwen_required_tensor_names(cfg);
let mut expected_planned: Vec<String> = required
.into_iter()
.filter(|n| rotation_plan.for_tensor(n).is_some())
.collect();
let lm_head_name = QWEN35_LM_HEAD_NAME.to_string();
if !expected_planned.iter().any(|n| n == &lm_head_name) {
expected_planned.push(lm_head_name);
}
let mut max_abs = 0.0_f64;
for expected_name in &expected_planned {
let tensor_rotation = rotation_plan.for_tensor(expected_name).ok_or_else(|| {
InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: expected planned tensor `{expected_name}` \
has no rotation rule (qwen_required_tensor_names/rotation_plan inconsistency)"
))
})?;
let source = if let Some(t) = original.get(expected_name) {
t
} else if expected_name == QWEN35_LM_HEAD_NAME && cfg.tie_word_embeddings {
original.get(QWEN35_EMBED_TOKENS_NAME).ok_or_else(|| {
InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: tied config requires \
either `{QWEN35_LM_HEAD_NAME}` or `{QWEN35_EMBED_TOKENS_NAME}` in the \
original working set as the source for lm_head reconstruction"
))
})?
} else {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: planned tensor `{expected_name}` \
missing from original working set"
)));
};
let actual = rotated.get(expected_name).ok_or_else(|| {
InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: planned tensor `{expected_name}` \
missing from rotated working set"
))
})?;
if source.shape.len() != 2 || actual.shape != source.shape {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: tensor `{expected_name}` shape mismatch \
(source={:?}, rotated={:?})",
source.shape, actual.shape
)));
}
let rows = source.shape[0];
let cols = source.shape[1];
let expected_len = rows.checked_mul(cols).ok_or_else(|| {
InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: rows*cols overflow on `{expected_name}` \
({rows}*{cols})"
))
})?;
if source.data.len() != expected_len || actual.data.len() != expected_len {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: tensor `{expected_name}` data.len() mismatch \
(source={}, rotated={}, rows*cols={expected_len})",
source.data.len(),
actual.data.len()
)));
}
let mut expected_rot = source.data.clone();
match tensor_rotation.side {
AbsorptionSide::InputSide => {
if cols != hidden {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: input-side tensor `{expected_name}` \
cols={cols} != hidden={hidden} (rotation plan invariant violated)"
)));
}
if let Some(norm_name) = fusion_gamma.get(expected_name.as_str()) {
let norm = original.get(*norm_name).ok_or_else(|| {
InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: fusion gamma source \
`{norm_name}` (for downstream `{expected_name}`) not in original \
working set"
))
})?;
if norm.shape.len() != 1 || norm.shape[0] != cols || norm.data.len() != cols {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: fusion gamma `{norm_name}` \
shape/data mismatch (shape={:?}, data.len()={}, expected cols={cols})",
norm.shape,
norm.data.len()
)));
}
fuse_shifted_rmsnorm_into_next_layer_f64(
&mut expected_rot,
rows,
cols,
&norm.data,
)?;
}
absorb_input_rotation_f64(&mut expected_rot, rows, cols, rotation)?;
}
AbsorptionSide::OutputSide => {
if rows != hidden {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: output-side tensor `{expected_name}` \
rows={rows} != hidden={hidden} (rotation plan invariant violated)"
)));
}
if fusion_gamma.contains_key(expected_name.as_str()) {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: output-side tensor `{expected_name}` \
unexpectedly has a fusion rule (rotation/fusion plan inconsistency)"
)));
}
absorb_output_rotation_f64(&mut expected_rot, rows, cols, rotation)?;
}
}
let delta = expected_rot
.iter()
.zip(actual.data.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0_f64, f64::max);
if delta > max_abs {
max_abs = delta;
}
}
Ok(max_abs)
}
fn deterministic_probe_tokens(seed: u64, n: usize, vocab_size: usize) -> Vec<u32> {
let mut state = seed;
(0..n)
.map(|_| {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
((state >> 32) % vocab_size as u64) as u32
})
.collect()
}
fn rotation_chain_probe_qwen35(
tensors: &HashMap<String, TensorEntry>,
cfg: &Qwen35Config,
token: u32,
) -> Result<Vec<f64>, InferenceError> {
let hidden = cfg.hidden_size;
let vocab = cfg.vocab_size;
let intermediate = cfg.intermediate_size;
let full_q_dim = cfg.full_q_dim();
let linear_output_dim = cfg.linear_output_dim();
let eps = f64::from(cfg.rms_norm_eps);
if (token as usize) >= vocab {
return Err(InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: token id {token} out of range (vocab_size={vocab})"
)));
}
let embed = get_tensor_2d(tensors, QWEN35_EMBED_TOKENS_NAME, vocab, hidden)?;
let row = (token as usize) * hidden;
let mut h: Vec<f64> = embed.data[row..row + hidden].to_vec();
for layer in 0..cfg.num_hidden_layers {
let prefix = format!("model.language_model.layers.{layer}");
let gamma_in = get_tensor_1d(tensors, &format!("{prefix}.input_layernorm.weight"), hidden)?;
let h_pre = rms_normalize_shifted(&h, &gamma_in.data, eps);
let attn_out = if cfg.is_full_attention(layer) {
let q_proj = get_tensor_2d(
tensors,
&format!("{prefix}.self_attn.q_proj.weight"),
2 * full_q_dim,
hidden,
)?;
let o_proj = get_tensor_2d(
tensors,
&format!("{prefix}.self_attn.o_proj.weight"),
hidden,
full_q_dim,
)?;
let q_full = matvec_f64(&q_proj.data, 2 * full_q_dim, hidden, &h_pre);
matvec_f64(&o_proj.data, hidden, full_q_dim, &q_full[..full_q_dim])
} else {
let in_proj_z = get_tensor_2d(
tensors,
&format!("{prefix}.linear_attn.in_proj_z.weight"),
linear_output_dim,
hidden,
)?;
let out_proj = get_tensor_2d(
tensors,
&format!("{prefix}.linear_attn.out_proj.weight"),
hidden,
linear_output_dim,
)?;
let z = matvec_f64(&in_proj_z.data, linear_output_dim, hidden, &h_pre);
matvec_f64(&out_proj.data, hidden, linear_output_dim, &z)
};
add_in_place(&mut h, &attn_out);
let gamma_post = get_tensor_1d(
tensors,
&format!("{prefix}.post_attention_layernorm.weight"),
hidden,
)?;
let h_pre_mlp = rms_normalize_shifted(&h, &gamma_post.data, eps);
let gate_proj = get_tensor_2d(
tensors,
&format!("{prefix}.mlp.gate_proj.weight"),
intermediate,
hidden,
)?;
let up_proj = get_tensor_2d(
tensors,
&format!("{prefix}.mlp.up_proj.weight"),
intermediate,
hidden,
)?;
let down_proj = get_tensor_2d(
tensors,
&format!("{prefix}.mlp.down_proj.weight"),
hidden,
intermediate,
)?;
let gate = matvec_f64(&gate_proj.data, intermediate, hidden, &h_pre_mlp);
let up = matvec_f64(&up_proj.data, intermediate, hidden, &h_pre_mlp);
let mid: Vec<f64> = gate.iter().zip(up.iter()).map(|(a, b)| a + b).collect();
let mlp_out = matvec_f64(&down_proj.data, hidden, intermediate, &mid);
add_in_place(&mut h, &mlp_out);
}
let gamma_final = get_tensor_1d(tensors, QWEN35_FINAL_NORM_NAME, hidden)?;
let h_final = rms_normalize_shifted(&h, &gamma_final.data, eps);
let lm_tensor = if tensors.contains_key(QWEN35_LM_HEAD_NAME) {
get_tensor_2d(tensors, QWEN35_LM_HEAD_NAME, vocab, hidden)?
} else if cfg.tie_word_embeddings {
embed
} else {
return Err(InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: untied config requires `{QWEN35_LM_HEAD_NAME}` \
in the working set (config says `tie_word_embeddings=false` and no fallback \
is valid in that case)"
)));
};
Ok(matvec_f64(&lm_tensor.data, vocab, hidden, &h_final))
}
fn get_tensor_2d<'a>(
tensors: &'a HashMap<String, TensorEntry>,
name: &str,
rows: usize,
cols: usize,
) -> Result<&'a TensorEntry, InferenceError> {
let t = tensors.get(name).ok_or_else(|| {
InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: tensor `{name}` not in working set"
))
})?;
if t.shape.len() != 2 || t.shape[0] != rows || t.shape[1] != cols {
return Err(InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: tensor `{name}` shape {:?} != expected [{rows}, {cols}]",
t.shape
)));
}
let expected = rows.checked_mul(cols).ok_or_else(|| {
InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: rows*cols overflow on `{name}` ({rows}*{cols})"
))
})?;
if t.data.len() != expected {
return Err(InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: tensor `{name}` data.len()={} != rows*cols {expected}",
t.data.len()
)));
}
Ok(t)
}
fn get_tensor_1d<'a>(
tensors: &'a HashMap<String, TensorEntry>,
name: &str,
len: usize,
) -> Result<&'a TensorEntry, InferenceError> {
let t = tensors.get(name).ok_or_else(|| {
InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: tensor `{name}` not in working set"
))
})?;
if t.shape.len() != 1 || t.shape[0] != len {
return Err(InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: tensor `{name}` shape {:?} != expected [{len}]",
t.shape
)));
}
if t.data.len() != len {
return Err(InferenceError::Inference(format!(
"rotation_chain_probe_qwen35: tensor `{name}` data.len()={} != expected {len}",
t.data.len()
)));
}
Ok(t)
}
fn rms_normalize_shifted(h: &[f64], gamma: &[f64], eps: f64) -> Vec<f64> {
debug_assert_eq!(h.len(), gamma.len());
let n = h.len();
let sum_sq: f64 = h.iter().map(|v| v * v).sum();
let rms = (sum_sq / n as f64 + eps).sqrt();
let inv_rms = 1.0 / rms;
h.iter()
.zip(gamma.iter())
.map(|(v, g)| v * inv_rms * (1.0 + g))
.collect()
}
fn matvec_f64(w: &[f64], rows: usize, cols: usize, x: &[f64]) -> Vec<f64> {
debug_assert_eq!(x.len(), cols);
debug_assert_eq!(w.len(), rows * cols);
let mut y = vec![0.0_f64; rows];
for r in 0..rows {
let row = &w[r * cols..(r + 1) * cols];
y[r] = row.iter().zip(x.iter()).map(|(a, b)| a * b).sum();
}
y
}
fn add_in_place(h: &mut [f64], addend: &[f64]) {
debug_assert_eq!(h.len(), addend.len());
for (a, b) in h.iter_mut().zip(addend.iter()) {
*a += b;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::qwen35_config::{LayerType, Qwen35Config, compute_layer_types};
use crate::quant::quarot::hadamard::RandomizedHadamard;
use crate::quant::quarot::lm_head::{
materialize_lm_head_for_qwen35, qwen35_final_norm_fusion_target,
};
use crate::quant::quarot::pipeline::{absorb_rotations, fuse_rmsnorms};
use crate::quant::quarot::plan::RotationPlan;
use crate::quant::quarot::rmsnorm_fusion::qwen35_per_layer_fusion_plan;
fn tiny_test_cfg() -> Qwen35Config {
let mut cfg = Qwen35Config::qwen35_0_8b();
cfg.hidden_size = 8;
cfg.num_hidden_layers = 2;
cfg.vocab_size = 4;
cfg.intermediate_size = 16;
cfg.num_attention_heads = 2;
cfg.num_key_value_heads = 1;
cfg.head_dim = 4;
cfg.linear_num_key_heads = 1;
cfg.linear_key_head_dim = 2;
cfg.linear_value_head_dim = 2;
cfg.linear_num_value_heads = Some(1);
cfg.full_attention_interval = 2;
cfg.layer_types = compute_layer_types(cfg.num_hidden_layers, cfg.full_attention_interval);
cfg.layer_mask = vec![true; cfg.num_hidden_layers];
cfg.tie_word_embeddings = false;
cfg.rms_norm_eps = 1e-6;
cfg
}
fn tied_tiny_test_cfg() -> Qwen35Config {
let mut cfg = tiny_test_cfg();
cfg.tie_word_embeddings = true;
cfg
}
fn synthetic_f64(n: usize, seed: u64) -> Vec<f64> {
let mut state = seed;
(0..n)
.map(|_| {
state = state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
let bits = (state >> 11) as u32;
(bits as f64 / u32::MAX as f64) - 0.5
})
.collect()
}
fn insert_tensor(
tensors: &mut HashMap<String, TensorEntry>,
name: &str,
shape: Vec<usize>,
data: Vec<f64>,
) {
tensors.insert(
name.to_string(),
TensorEntry {
name: name.to_string(),
shape,
data,
},
);
}
fn build_working_set(cfg: &Qwen35Config, seed: u64) -> HashMap<String, TensorEntry> {
let hidden = cfg.hidden_size;
let vocab = cfg.vocab_size;
let intermediate = cfg.intermediate_size;
let full_q_dim = cfg.full_q_dim();
let full_kv_dim = cfg.full_kv_dim();
let linear_qkv_dim = cfg.linear_qkv_dim();
let linear_output_dim = cfg.linear_output_dim();
let linear_num_heads = cfg.linear_num_key_heads;
let mut tensors = HashMap::new();
insert_tensor(
&mut tensors,
QWEN35_EMBED_TOKENS_NAME,
vec![vocab, hidden],
synthetic_f64(vocab * hidden, seed.wrapping_add(1)),
);
insert_tensor(
&mut tensors,
QWEN35_FINAL_NORM_NAME,
vec![hidden],
synthetic_f64(hidden, seed.wrapping_add(2)),
);
for i in 0..cfg.num_hidden_layers {
let prefix = format!("model.language_model.layers.{i}");
let layer_seed = seed.wrapping_add(100 + i as u64);
insert_tensor(
&mut tensors,
&format!("{prefix}.input_layernorm.weight"),
vec![hidden],
synthetic_f64(hidden, layer_seed.wrapping_add(1)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.post_attention_layernorm.weight"),
vec![hidden],
synthetic_f64(hidden, layer_seed.wrapping_add(2)),
);
if cfg.is_full_attention(i) {
insert_tensor(
&mut tensors,
&format!("{prefix}.self_attn.q_proj.weight"),
vec![2 * full_q_dim, hidden],
synthetic_f64(2 * full_q_dim * hidden, layer_seed.wrapping_add(10)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.self_attn.k_proj.weight"),
vec![full_kv_dim, hidden],
synthetic_f64(full_kv_dim * hidden, layer_seed.wrapping_add(11)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.self_attn.v_proj.weight"),
vec![full_kv_dim, hidden],
synthetic_f64(full_kv_dim * hidden, layer_seed.wrapping_add(12)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.self_attn.o_proj.weight"),
vec![hidden, full_q_dim],
synthetic_f64(hidden * full_q_dim, layer_seed.wrapping_add(13)),
);
} else {
insert_tensor(
&mut tensors,
&format!("{prefix}.linear_attn.in_proj_qkv.weight"),
vec![linear_qkv_dim, hidden],
synthetic_f64(linear_qkv_dim * hidden, layer_seed.wrapping_add(20)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.linear_attn.in_proj_z.weight"),
vec![linear_output_dim, hidden],
synthetic_f64(linear_output_dim * hidden, layer_seed.wrapping_add(21)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.linear_attn.in_proj_a.weight"),
vec![linear_num_heads, hidden],
synthetic_f64(linear_num_heads * hidden, layer_seed.wrapping_add(22)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.linear_attn.in_proj_b.weight"),
vec![linear_num_heads, hidden],
synthetic_f64(linear_num_heads * hidden, layer_seed.wrapping_add(23)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.linear_attn.out_proj.weight"),
vec![hidden, linear_output_dim],
synthetic_f64(hidden * linear_output_dim, layer_seed.wrapping_add(24)),
);
}
insert_tensor(
&mut tensors,
&format!("{prefix}.mlp.gate_proj.weight"),
vec![intermediate, hidden],
synthetic_f64(intermediate * hidden, layer_seed.wrapping_add(30)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.mlp.up_proj.weight"),
vec![intermediate, hidden],
synthetic_f64(intermediate * hidden, layer_seed.wrapping_add(31)),
);
insert_tensor(
&mut tensors,
&format!("{prefix}.mlp.down_proj.weight"),
vec![hidden, intermediate],
synthetic_f64(hidden * intermediate, layer_seed.wrapping_add(32)),
);
}
tensors
}
fn full_pipeline_plans(
cfg: &Qwen35Config,
) -> (
Vec<crate::quant::quarot::rmsnorm_fusion::RmsNormFusionTarget>,
RotationPlan,
) {
let mut fusion = qwen35_per_layer_fusion_plan(cfg).unwrap();
fusion.push(qwen35_final_norm_fusion_target());
(fusion, RotationPlan::qwen35_residual_stream_linear_layers())
}
#[test]
fn tiny_cfg_has_one_full_and_one_linear_layer() {
let cfg = tiny_test_cfg();
assert_eq!(cfg.num_hidden_layers, 2);
assert_eq!(cfg.layer_types[0], LayerType::LinearAttention);
assert_eq!(cfg.layer_types[1], LayerType::FullAttention);
}
#[test]
fn forward_equivalence_passes_on_full_untied_pipeline() {
let cfg = tiny_test_cfg();
let mut original = build_working_set(&cfg, 1);
insert_tensor(
&mut original,
QWEN35_LM_HEAD_NAME,
vec![cfg.vocab_size, cfg.hidden_size],
synthetic_f64(cfg.vocab_size * cfg.hidden_size, 999),
);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation = RandomizedHadamard::new(0xC0FFEE, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
let report = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap();
assert!(report.max_abs_error < 1e-5, "unexpected delta: {report:?}");
assert_eq!(report.probe_tokens.len(), 4);
for &t in &report.probe_tokens {
assert!((t as usize) < cfg.vocab_size);
}
}
#[test]
fn forward_equivalence_passes_on_full_tied_pipeline() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 2);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation = RandomizedHadamard::new(0xFEED_FACE, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
let report = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap();
assert!(report.max_abs_error < 1e-5, "unexpected delta: {report:?}");
}
#[test]
fn forward_equivalence_refuses_when_final_norm_fusion_skipped() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 4);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let per_layer = qwen35_per_layer_fusion_plan(&cfg).unwrap();
fuse_rmsnorms(&mut rotated, &per_layer).unwrap();
let rotation = RandomizedHadamard::new(0xBADC0DE, cfg.hidden_size).unwrap();
absorb_rotations(
&mut rotated,
&RotationPlan::qwen35_residual_stream_linear_layers(),
&rotation,
)
.unwrap();
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("forward-equivalence refused"),
"unexpected error: {msg}"
);
assert!(msg.contains("exceeds tolerance"), "unexpected error: {msg}");
}
#[test]
fn forward_equivalence_refuses_on_corrupted_rotated_tensor() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 5);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation = RandomizedHadamard::new(0x55AA_55AA, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
let lm = rotated
.get_mut(QWEN35_LM_HEAD_NAME)
.expect("lm_head should exist after materialize");
for v in lm.data.iter_mut() {
*v *= 1.25;
}
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("forward-equivalence refused"),
"unexpected error: {msg}"
);
}
#[test]
fn forward_equivalence_refuses_when_per_layer_fusion_skipped() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 6);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let final_only = vec![qwen35_final_norm_fusion_target()];
fuse_rmsnorms(&mut rotated, &final_only).unwrap();
let rotation = RandomizedHadamard::new(0xABCDEF, cfg.hidden_size).unwrap();
absorb_rotations(
&mut rotated,
&RotationPlan::qwen35_residual_stream_linear_layers(),
&rotation,
)
.unwrap();
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("forward-equivalence refused"),
"unexpected error: {msg}"
);
}
#[test]
fn forward_equivalence_errors_on_untied_original_missing_lm_head() {
let cfg = tiny_test_cfg(); let original = build_working_set(&cfg, 7); let mut rotated = original.clone();
insert_tensor(
&mut rotated,
QWEN35_LM_HEAD_NAME,
vec![cfg.vocab_size, cfg.hidden_size],
synthetic_f64(cfg.vocab_size * cfg.hidden_size, 7777),
);
let rotation = RandomizedHadamard::new(1, cfg.hidden_size).unwrap();
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains(QWEN35_LM_HEAD_NAME), "unexpected error: {msg}");
assert!(msg.contains("untied"), "unexpected error: {msg}");
}
#[test]
fn forward_equivalence_errors_on_missing_required_tensor() {
let cfg = tiny_test_cfg();
let mut original = build_working_set(&cfg, 8);
insert_tensor(
&mut original,
QWEN35_LM_HEAD_NAME,
vec![cfg.vocab_size, cfg.hidden_size],
synthetic_f64(cfg.vocab_size * cfg.hidden_size, 888),
);
let rotated = original.clone();
original.remove(QWEN35_FINAL_NORM_NAME);
let rotation = RandomizedHadamard::new(2, cfg.hidden_size).unwrap();
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains(QWEN35_FINAL_NORM_NAME),
"unexpected error: {msg}"
);
assert!(
msg.contains("not in working set"),
"unexpected error: {msg}"
);
}
#[test]
fn forward_equivalence_errors_on_shape_mismatch() {
let cfg = tiny_test_cfg();
let mut original = build_working_set(&cfg, 9);
original.insert(
QWEN35_EMBED_TOKENS_NAME.to_string(),
TensorEntry {
name: QWEN35_EMBED_TOKENS_NAME.to_string(),
shape: vec![cfg.vocab_size, cfg.hidden_size + 1],
data: vec![0.0; cfg.vocab_size * (cfg.hidden_size + 1)],
},
);
insert_tensor(
&mut original,
QWEN35_LM_HEAD_NAME,
vec![cfg.vocab_size, cfg.hidden_size],
synthetic_f64(cfg.vocab_size * cfg.hidden_size, 999),
);
let rotated = original.clone();
let rotation = RandomizedHadamard::new(3, cfg.hidden_size).unwrap();
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains(QWEN35_EMBED_TOKENS_NAME),
"unexpected error: {msg}"
);
assert!(msg.contains("shape"), "unexpected error: {msg}");
}
#[test]
fn forward_equivalence_rejects_moe_config() {
let cfg = Qwen35Config::qwen36_35b_a3b();
assert!(cfg.is_moe());
let original = HashMap::new();
let rotated = HashMap::new();
let rotation = RandomizedHadamard::new(0, 8).unwrap();
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("MoE"), "unexpected error: {msg}");
}
#[test]
fn forward_equivalence_rejects_zero_probe_tokens() {
let cfg = tiny_test_cfg();
let original = HashMap::new();
let rotated = HashMap::new();
let rotation = RandomizedHadamard::new(0, cfg.hidden_size).unwrap();
let fc = ForwardEquivalenceConfig {
num_probe_tokens: 0,
..Default::default()
};
let err = assert_forward_equivalence_qwen35(&original, &rotated, &cfg, &rotation, &fc)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("num_probe_tokens must be > 0"),
"unexpected error: {msg}"
);
}
#[test]
fn forward_equivalence_rejects_non_positive_tolerance() {
let cfg = tiny_test_cfg();
let original = HashMap::new();
let rotated = HashMap::new();
let rotation = RandomizedHadamard::new(0, cfg.hidden_size).unwrap();
for bad in [0.0_f64, -1e-5, f64::NAN] {
let fc = ForwardEquivalenceConfig {
tolerance: bad,
..Default::default()
};
let err = assert_forward_equivalence_qwen35(&original, &rotated, &cfg, &rotation, &fc)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("tolerance must be a positive finite value"),
"unexpected error for tolerance={bad}: {msg}"
);
}
}
#[test]
fn forward_equivalence_rejects_rotation_dim_mismatch() {
let cfg = tiny_test_cfg();
let original = build_working_set(&cfg, 11);
let rotated = original.clone();
let rotation = RandomizedHadamard::new(0, cfg.hidden_size * 2).unwrap();
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("rotation.dim()"), "unexpected error: {msg}");
assert!(
msg.contains(&format!("cfg.hidden_size={}", cfg.hidden_size)),
"unexpected error: {msg}"
);
}
#[test]
fn probe_tokens_are_deterministic_in_seed() {
let a = deterministic_probe_tokens(0xDEAD_BEEF, 4, 100);
let b = deterministic_probe_tokens(0xDEAD_BEEF, 4, 100);
assert_eq!(a, b);
let c = deterministic_probe_tokens(0xDEAD_BEEF_u64.wrapping_add(1), 4, 100);
assert_ne!(
a, c,
"different seeds should produce different probe tokens"
);
}
#[test]
fn refuse_error_message_includes_max_abs_and_tolerance() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 10);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let per_layer = qwen35_per_layer_fusion_plan(&cfg).unwrap();
fuse_rmsnorms(&mut rotated, &per_layer).unwrap();
let rotation = RandomizedHadamard::new(0x1234_5678, cfg.hidden_size).unwrap();
absorb_rotations(
&mut rotated,
&RotationPlan::qwen35_residual_stream_linear_layers(),
&rotation,
)
.unwrap();
let fc = ForwardEquivalenceConfig {
tolerance: 1e-12,
..Default::default()
};
let err = assert_forward_equivalence_qwen35(&original, &rotated, &cfg, &rotation, &fc)
.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("max_abs_error="), "unexpected error: {msg}");
assert!(
msg.contains("exceeds tolerance="),
"unexpected error: {msg}"
);
assert!(
msg.contains("chain probe") || msg.contains("per-tensor"),
"expected refuse message to name the failing check: {msg}"
);
}
fn assert_corrupting_planned_tensor_refuses(victim_name: &str, factor: f64, seed: u64) {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, seed);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation =
RandomizedHadamard::new(seed.wrapping_mul(0x9E37_79B9_7F4A_7C15), cfg.hidden_size)
.unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_or_else(|e| {
panic!("pre-corruption pipeline must pass for victim `{victim_name}`: {e}")
});
let victim = rotated
.get_mut(victim_name)
.unwrap_or_else(|| panic!("victim tensor `{victim_name}` missing from working set"));
for v in victim.data.iter_mut() {
*v *= factor;
}
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("forward-equivalence refused"),
"expected refuse for corrupted `{victim_name}`: {msg}"
);
assert!(
msg.contains("per-tensor"),
"corruption of `{victim_name}` should be caught by the per-tensor check: {msg}"
);
}
#[test]
fn per_tensor_check_catches_corrupted_k_proj() {
assert_corrupting_planned_tensor_refuses(
"model.language_model.layers.1.self_attn.k_proj.weight",
1000.0,
21,
);
}
#[test]
fn per_tensor_check_catches_corrupted_v_proj() {
assert_corrupting_planned_tensor_refuses(
"model.language_model.layers.1.self_attn.v_proj.weight",
-2.0,
22,
);
}
#[test]
fn per_tensor_check_catches_corrupted_q_proj_gate_z_half() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 23);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation = RandomizedHadamard::new(0xAB12_34CD, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
let full_layer = (0..cfg.num_hidden_layers)
.find(|&i| cfg.is_full_attention(i))
.expect("tied_tiny_test_cfg must have at least one full-attention layer");
let q_name = format!("model.language_model.layers.{full_layer}.self_attn.q_proj.weight");
let full_q_dim = cfg.full_q_dim();
let hidden = cfg.hidden_size;
let victim = rotated.get_mut(&q_name).unwrap();
for r in full_q_dim..(2 * full_q_dim) {
for c in 0..hidden {
victim.data[r * hidden + c] *= 3.0;
}
}
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("per-tensor"), "unexpected error: {msg}");
}
#[test]
fn per_tensor_check_catches_corrupted_in_proj_qkv() {
assert_corrupting_planned_tensor_refuses(
"model.language_model.layers.0.linear_attn.in_proj_qkv.weight",
10.0,
24,
);
}
#[test]
fn per_tensor_check_catches_corrupted_in_proj_a() {
assert_corrupting_planned_tensor_refuses(
"model.language_model.layers.0.linear_attn.in_proj_a.weight",
0.5,
25,
);
}
#[test]
fn per_tensor_check_catches_corrupted_in_proj_b() {
assert_corrupting_planned_tensor_refuses(
"model.language_model.layers.0.linear_attn.in_proj_b.weight",
-1.0,
26,
);
}
#[test]
fn per_tensor_check_catches_single_element_perturbation_orthogonal_to_probe_vector() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 30);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation = RandomizedHadamard::new(0x9876_5432, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap();
let k_name = "model.language_model.layers.1.self_attn.k_proj.weight";
let victim = rotated.get_mut(k_name).unwrap();
victim.data[0] += 0.5;
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains("per-tensor"), "unexpected error: {msg}");
}
#[test]
fn per_tensor_check_errors_when_planned_tensor_missing_from_both_maps() {
let cfg = tied_tiny_test_cfg();
let mut original = build_working_set(&cfg, 31);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation = RandomizedHadamard::new(0xFEED_0BAD, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
let k_name = "model.language_model.layers.1.self_attn.k_proj.weight".to_string();
assert!(original.remove(&k_name).is_some());
assert!(rotated.remove(&k_name).is_some());
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(msg.contains(&k_name), "unexpected error: {msg}");
assert!(msg.contains("missing"), "unexpected error: {msg}");
}
#[test]
fn either_check_catches_corrupted_embed_tokens() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 27);
let mut rotated = original.clone();
materialize_lm_head_for_qwen35(&mut rotated, &cfg).unwrap();
let (fusion, rot_plan) = full_pipeline_plans(&cfg);
fuse_rmsnorms(&mut rotated, &fusion).unwrap();
let rotation = RandomizedHadamard::new(0xEDEDED, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
let embed = rotated.get_mut(QWEN35_EMBED_TOKENS_NAME).unwrap();
for v in embed.data.iter_mut() {
*v *= 1.5;
}
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("forward-equivalence refused"),
"unexpected error: {msg}"
);
assert!(
msg.contains("chain probe") || msg.contains("per-tensor"),
"expected refuse message to name the failing check: {msg}"
);
}
}