use std::collections::HashMap;
use std::mem::size_of;
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::io::QuarotTensorReader;
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};
#[cfg(test)]
pub(crate) mod pre_admission_allocation_tracking {
use std::alloc::{GlobalAlloc, Layout, System};
use std::cell::Cell;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Phase {
Inactive,
WaitingForConverterBoundary,
BeforeRejection,
AfterRejection,
WaitingForMaterializedWorkingSetBoundary,
MaterializedWorkingSetBoundary,
MaterializedWorkingSetBoundaryCompleted,
}
#[derive(Clone, Copy)]
struct State {
phase: Phase,
before_rejection_allocation_calls: usize,
after_rejection_allocation_calls: usize,
materialized_working_set_allocation_calls: usize,
rejection_seen: bool,
reader_boundary_seen: bool,
materialized_working_set_boundary_seen: bool,
materialized_working_set_boundary_completed: bool,
}
const INACTIVE: State = State {
phase: Phase::Inactive,
before_rejection_allocation_calls: 0,
after_rejection_allocation_calls: 0,
materialized_working_set_allocation_calls: 0,
rejection_seen: false,
reader_boundary_seen: false,
materialized_working_set_boundary_seen: false,
materialized_working_set_boundary_completed: false,
};
std::thread_local! {
static STATE: Cell<State> = const { Cell::new(INACTIVE) };
}
struct TrackingAllocator;
#[global_allocator]
static GLOBAL: TrackingAllocator = TrackingAllocator;
fn record_allocation() {
let _ = STATE.try_with(|cell| {
let mut state = cell.get();
match state.phase {
Phase::BeforeRejection => {
state.before_rejection_allocation_calls =
state.before_rejection_allocation_calls.saturating_add(1);
}
Phase::AfterRejection => {
state.after_rejection_allocation_calls =
state.after_rejection_allocation_calls.saturating_add(1);
}
Phase::MaterializedWorkingSetBoundary => {
state.materialized_working_set_allocation_calls = state
.materialized_working_set_allocation_calls
.saturating_add(1);
}
Phase::Inactive
| Phase::WaitingForConverterBoundary
| Phase::WaitingForMaterializedWorkingSetBoundary
| Phase::MaterializedWorkingSetBoundaryCompleted => return,
}
cell.set(state);
});
}
unsafe impl GlobalAlloc for TrackingAllocator {
unsafe fn alloc(&self, layout: Layout) -> *mut u8 {
record_allocation();
System.alloc(layout)
}
unsafe fn alloc_zeroed(&self, layout: Layout) -> *mut u8 {
record_allocation();
System.alloc_zeroed(layout)
}
unsafe fn dealloc(&self, ptr: *mut u8, layout: Layout) {
System.dealloc(ptr, layout);
}
unsafe fn realloc(&self, ptr: *mut u8, layout: Layout, new_size: usize) -> *mut u8 {
record_allocation();
System.realloc(ptr, layout, new_size)
}
}
pub(crate) struct Guard {
active: bool,
}
pub(crate) struct Observation {
pub(crate) before_rejection_allocation_calls: usize,
pub(crate) after_rejection_allocation_calls: usize,
pub(crate) materialized_working_set_allocation_calls: usize,
pub(crate) rejection_seen: bool,
pub(crate) reader_boundary_seen: bool,
pub(crate) materialized_working_set_boundary_seen: bool,
pub(crate) materialized_working_set_boundary_completed: bool,
}
fn start_in_phase(phase: Phase) -> Guard {
STATE.with(|cell| {
assert_eq!(
cell.get().phase,
Phase::Inactive,
"allocation tracking already active"
);
cell.set(State { phase, ..INACTIVE });
});
Guard { active: true }
}
pub(super) fn start() -> Guard {
start_in_phase(Phase::BeforeRejection)
}
pub(crate) fn start_at_converter_boundary() -> Guard {
start_in_phase(Phase::WaitingForConverterBoundary)
}
pub(crate) fn start_at_materialized_working_set_boundary() -> Guard {
start_in_phase(Phase::WaitingForMaterializedWorkingSetBoundary)
}
pub(crate) fn mark_converter_boundary() {
let _ = STATE.try_with(|cell| {
let mut state = cell.get();
if state.phase == Phase::WaitingForConverterBoundary {
state.phase = Phase::BeforeRejection;
cell.set(state);
}
});
}
pub(crate) fn mark_reader_boundary() {
let _ = STATE.try_with(|cell| {
let mut state = cell.get();
if state.phase != Phase::Inactive {
state.reader_boundary_seen = true;
cell.set(state);
}
});
}
pub(crate) fn mark_materialized_working_set_boundary() {
let _ = STATE.try_with(|cell| {
let mut state = cell.get();
if state.phase != Phase::Inactive {
state.materialized_working_set_boundary_seen = true;
}
if state.phase == Phase::WaitingForMaterializedWorkingSetBoundary {
state.phase = Phase::MaterializedWorkingSetBoundary;
}
cell.set(state);
});
}
pub(crate) fn mark_materialized_working_set_boundary_completed() {
let _ = STATE.try_with(|cell| {
let mut state = cell.get();
if state.phase != Phase::Inactive {
state.materialized_working_set_boundary_completed = true;
}
if state.phase == Phase::MaterializedWorkingSetBoundary {
state.phase = Phase::MaterializedWorkingSetBoundaryCompleted;
}
cell.set(state);
});
}
pub(super) fn mark_rejection() {
let _ = STATE.try_with(|cell| {
let mut state = cell.get();
if state.phase == Phase::BeforeRejection {
state.phase = Phase::AfterRejection;
state.rejection_seen = true;
cell.set(state);
}
});
}
impl Guard {
pub(crate) fn finish(mut self) -> Observation {
self.active = false;
STATE.with(|cell| {
let state = cell.replace(INACTIVE);
Observation {
before_rejection_allocation_calls: state.before_rejection_allocation_calls,
after_rejection_allocation_calls: state.after_rejection_allocation_calls,
materialized_working_set_allocation_calls: state
.materialized_working_set_allocation_calls,
rejection_seen: state.rejection_seen,
reader_boundary_seen: state.reader_boundary_seen,
materialized_working_set_boundary_seen: state
.materialized_working_set_boundary_seen,
materialized_working_set_boundary_completed: state
.materialized_working_set_boundary_completed,
}
})
}
}
impl Drop for Guard {
fn drop(&mut self) {
if self.active {
let _ = STATE.try_with(|cell| cell.set(INACTIVE));
}
}
}
}
const MAX_FORWARD_EQUIVALENCE_PROBE_SNAPSHOT_BYTES: usize = 64 * 1024 * 1024;
#[derive(Debug, Clone)]
pub struct ForwardEquivalenceConfig {
pub num_probe_tokens: usize,
pub tolerance: f64,
pub seed: u64,
}
pub(crate) struct ForwardEquivalenceAdmission<'cfg, 'forward> {
cfg: &'cfg Qwen35Config,
forward_cfg: &'forward ForwardEquivalenceConfig,
}
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(crate) struct ForwardEquivalenceSnapshot<'a> {
cfg: &'a Qwen35Config,
rotation: &'a RandomizedHadamard,
probe_tokens: Vec<u32>,
original_logits: Vec<Vec<f64>>,
fusion_gammas: HashMap<String, TensorEntry>,
tolerance: f64,
}
trait OriginalTensorSource {
fn has_tensor(&self, name: &str) -> bool;
fn load_tensor(&self, name: &str) -> Result<TensorEntry, InferenceError>;
}
impl OriginalTensorSource for HashMap<String, TensorEntry> {
fn has_tensor(&self, name: &str) -> bool {
self.contains_key(name)
}
fn load_tensor(&self, name: &str) -> Result<TensorEntry, InferenceError> {
self.get(name)
.cloned()
.ok_or_else(|| InferenceError::MissingTensor(name.to_string()))
}
}
struct ReaderOriginalTensorSource<'a> {
reader: &'a QuarotTensorReader,
tied_lm_head_from_embed: bool,
}
impl OriginalTensorSource for ReaderOriginalTensorSource<'_> {
fn has_tensor(&self, name: &str) -> bool {
if self.tied_lm_head_from_embed && name == QWEN35_LM_HEAD_NAME {
return false;
}
self.reader.has_tensor(name)
}
fn load_tensor(&self, name: &str) -> Result<TensorEntry, InferenceError> {
let (data, shape) = self.reader.read_tensor_f64(name)?;
Ok(TensorEntry {
name: name.to_string(),
shape,
data,
})
}
}
pub(crate) fn preflight_forward_equivalence_probe_budget(
cfg: &Qwen35Config,
forward_cfg: &ForwardEquivalenceConfig,
) -> Result<(), InferenceError> {
let Some(logit_elements) = forward_cfg.num_probe_tokens.checked_mul(cfg.vocab_size) else {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: retained chain-logit size overflow"
));
};
let Some(logit_bytes) = logit_elements.checked_mul(size_of::<f64>()) else {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: retained chain-logit size overflow"
));
};
let Some(row_descriptor_bytes) = forward_cfg
.num_probe_tokens
.checked_mul(size_of::<Vec<f64>>())
else {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: retained chain-logit size overflow"
));
};
let Some(probe_token_bytes) = forward_cfg.num_probe_tokens.checked_mul(size_of::<u32>()) else {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: retained chain-logit size overflow"
));
};
let Some(retained_bytes) = logit_bytes
.checked_add(row_descriptor_bytes)
.and_then(|bytes| bytes.checked_add(probe_token_bytes))
else {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: retained chain-logit size overflow"
));
};
if retained_bytes > MAX_FORWARD_EQUIVALENCE_PROBE_SNAPSHOT_BYTES {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: retained chain-logit budget exceeded \
(num_probe_tokens={}, vocab_size={}, required_bytes={retained_bytes}, \
max_bytes={MAX_FORWARD_EQUIVALENCE_PROBE_SNAPSHOT_BYTES})",
forward_cfg.num_probe_tokens, cfg.vocab_size
));
}
Ok(())
}
pub(crate) fn validate_forward_equivalence_admission<'cfg, 'forward>(
cfg: &'cfg Qwen35Config,
forward_cfg: &'forward ForwardEquivalenceConfig,
) -> Result<ForwardEquivalenceAdmission<'cfg, 'forward>, InferenceError> {
if cfg.is_moe() {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: MoE configs are deferred to v1 \
(the rotation/fusion pipeline rejects MoE upstream; this probe \
has no expert-mixing path)"
));
}
if forward_cfg.num_probe_tokens == 0 {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: num_probe_tokens must be > 0"
));
}
if !forward_cfg.tolerance.is_finite() || forward_cfg.tolerance <= 0.0 {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: tolerance must be a positive finite value, got {}",
forward_cfg.tolerance
));
}
if cfg.vocab_size == 0 {
return reject_forward_equivalence_admission(format_args!(
"assert_forward_equivalence_qwen35: cfg.vocab_size must be > 0"
));
}
preflight_forward_equivalence_probe_budget(cfg, forward_cfg)?;
Ok(ForwardEquivalenceAdmission { cfg, forward_cfg })
}
fn validate_forward_equivalence_rotation(
cfg: &Qwen35Config,
rotation: &RandomizedHadamard,
) -> Result<(), InferenceError> {
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
)));
}
Ok(())
}
fn reject_forward_equivalence_admission<T>(
message: std::fmt::Arguments<'_>,
) -> Result<T, InferenceError> {
#[cfg(test)]
pre_admission_allocation_tracking::mark_rejection();
Err(InferenceError::Inference(message.to_string()))
}
fn full_fusion_plan(cfg: &Qwen35Config) -> Result<Vec<RmsNormFusionTarget>, InferenceError> {
let mut fusion_plan = qwen35_per_layer_fusion_plan(cfg)?;
fusion_plan.push(qwen35_final_norm_fusion_target());
Ok(fusion_plan)
}
pub(crate) fn prepare_forward_equivalence_qwen35<'a>(
original: &HashMap<String, TensorEntry>,
cfg: &'a Qwen35Config,
rotation: &'a RandomizedHadamard,
forward_cfg: &ForwardEquivalenceConfig,
) -> Result<ForwardEquivalenceSnapshot<'a>, InferenceError> {
let admission = validate_forward_equivalence_admission(cfg, forward_cfg)?;
prepare_forward_equivalence_qwen35_after_admission(original, rotation, admission)
}
pub(crate) fn prepare_forward_equivalence_qwen35_after_admission<'cfg, 'forward>(
original: &HashMap<String, TensorEntry>,
rotation: &'cfg RandomizedHadamard,
admission: ForwardEquivalenceAdmission<'cfg, 'forward>,
) -> Result<ForwardEquivalenceSnapshot<'cfg>, InferenceError> {
let ForwardEquivalenceAdmission { cfg, forward_cfg } = admission;
validate_forward_equivalence_rotation(cfg, rotation)?;
let probe_tokens = deterministic_probe_tokens(
forward_cfg.seed,
forward_cfg.num_probe_tokens,
cfg.vocab_size,
);
let original_logits = probe_tokens
.iter()
.map(|&token| rotation_chain_probe_qwen35(original, cfg, token))
.collect::<Result<Vec<_>, _>>()?;
let mut fusion_gammas = HashMap::new();
for target in full_fusion_plan(cfg)? {
if fusion_gammas.contains_key(&target.norm_tensor) {
continue;
}
let gamma = original.get(&target.norm_tensor).ok_or_else(|| {
InferenceError::Inference(format!(
"prepare_forward_equivalence_qwen35: fusion gamma `{}` not in original \
working set",
target.norm_tensor
))
})?;
fusion_gammas.insert(target.norm_tensor, gamma.clone());
}
Ok(ForwardEquivalenceSnapshot {
cfg,
rotation,
probe_tokens,
original_logits,
fusion_gammas,
tolerance: forward_cfg.tolerance,
})
}
pub fn assert_forward_equivalence_qwen35(
original: &HashMap<String, TensorEntry>,
rotated: &HashMap<String, TensorEntry>,
cfg: &Qwen35Config,
rotation: &RandomizedHadamard,
forward_cfg: &ForwardEquivalenceConfig,
) -> Result<ForwardEquivalenceReport, InferenceError> {
let snapshot = prepare_forward_equivalence_qwen35(original, cfg, rotation, forward_cfg)?;
assert_forward_equivalence_snapshot(snapshot, original, rotated)
}
pub(crate) fn assert_prepared_forward_equivalence_qwen35(
snapshot: ForwardEquivalenceSnapshot<'_>,
reader: &QuarotTensorReader,
rotated: &HashMap<String, TensorEntry>,
) -> Result<ForwardEquivalenceReport, InferenceError> {
let original = ReaderOriginalTensorSource {
reader,
tied_lm_head_from_embed: snapshot.cfg.tie_word_embeddings,
};
assert_forward_equivalence_snapshot(snapshot, &original, rotated)
}
fn assert_forward_equivalence_snapshot<S: OriginalTensorSource>(
snapshot: ForwardEquivalenceSnapshot<'_>,
original: &S,
rotated: &HashMap<String, TensorEntry>,
) -> Result<ForwardEquivalenceReport, InferenceError> {
let ForwardEquivalenceSnapshot {
cfg,
rotation,
probe_tokens,
original_logits,
fusion_gammas,
tolerance,
} = snapshot;
if original_logits.len() != probe_tokens.len() {
return Err(InferenceError::Inference(format!(
"assert_forward_equivalence_qwen35: prepared probe count mismatch \
(tokens={}, original_logits={})",
probe_tokens.len(),
original_logits.len()
)));
}
let mut chain_max_abs = 0.0_f64;
let mut chain_total_abs = 0.0_f64;
let mut chain_count: usize = 0;
for (&token, logits_orig) in probe_tokens.iter().zip(original_logits.iter()) {
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 (logit_index, (a, b)) in logits_orig.iter().zip(logits_rot.iter()).enumerate() {
let d = (a - b).abs();
if !d.is_finite() {
return Err(InferenceError::Inference(format!(
"forward-equivalence refused: chain probe non-finite error \
(token={token}, logit_index={logit_index}, original={a}, rotated={b}, \
abs_error={d}). Do NOT write conversion artifacts."
)));
}
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 > tolerance {
return Err(InferenceError::Inference(format!(
"forward-equivalence refused: chain probe max_abs_error={chain_max_abs} \
exceeds tolerance={tolerance} (mean_abs_error={chain_mean_abs}, \
probe_tokens={probe_tokens:?}). \
Do NOT write conversion artifacts — the pipeline produced logits \
that disagree with the original model."
)));
}
let rotation_plan = RotationPlan::qwen35_residual_stream_linear_layers();
let fusion_plan = full_fusion_plan(cfg)?;
let per_tensor_max_abs = check_per_tensor_rotation_equivalence(
original,
&fusion_gammas,
rotated,
cfg,
rotation,
&rotation_plan,
&fusion_plan,
)?;
if per_tensor_max_abs > tolerance {
return Err(InferenceError::Inference(format!(
"forward-equivalence refused: per-tensor max_abs_error={per_tensor_max_abs} \
exceeds tolerance={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."
)));
}
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,
})
}
fn check_per_tensor_rotation_equivalence<S: OriginalTensorSource>(
original: &S,
fusion_gammas: &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_name = if original.has_tensor(expected_name) {
expected_name.as_str()
} else if expected_name == QWEN35_LM_HEAD_NAME && cfg.tie_word_embeddings {
if !original.has_tensor(QWEN35_EMBED_TOKENS_NAME) {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: tied config requires \
either `{QWEN35_LM_HEAD_NAME}` or `{QWEN35_EMBED_TOKENS_NAME}` in the \
original tensor source as the source for lm_head reconstruction"
)));
}
QWEN35_EMBED_TOKENS_NAME
} else {
return Err(InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: planned tensor `{expected_name}` \
missing from original tensor source"
)));
};
let source = original.load_tensor(source_name).map_err(|err| {
InferenceError::Inference(format!(
"check_per_tensor_rotation_equivalence: failed to load original source \
`{source_name}` for planned tensor `{expected_name}`: {err}"
))
})?;
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;
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 = fusion_gammas.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 prepared \
original snapshot"
))
})?;
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 mut delta = 0.0_f64;
for (element_index, (expected, actual)) in
expected_rot.iter().zip(actual.data.iter()).enumerate()
{
let element_delta = (expected - actual).abs();
if !element_delta.is_finite() {
return Err(InferenceError::Inference(format!(
"forward-equivalence refused: per-tensor non-finite error \
(tensor={expected_name}, element_index={element_index}, \
expected={expected}, actual={actual}, abs_error={element_delta}). \
Do NOT write conversion artifacts."
)));
}
if element_delta > delta {
delta = element_delta;
}
}
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]
#[allow(
clippy::type_complexity,
reason = "the explicit completion signature is the compile-time assertion"
)]
fn prepared_completion_cannot_accept_dense_to_moe_handoff() {
let dense_cfg = tied_tiny_test_cfg();
let original = build_working_set(&dense_cfg, 0xA11C_E5E5);
let rotation = RandomizedHadamard::new(0xD3A5_EA5E, dense_cfg.hidden_size).unwrap();
let forward_cfg = ForwardEquivalenceConfig {
num_probe_tokens: 1,
..Default::default()
};
let snapshot =
prepare_forward_equivalence_qwen35(&original, &dense_cfg, &rotation, &forward_cfg)
.unwrap();
assert!(std::ptr::eq(snapshot.cfg, &dense_cfg));
assert!(std::ptr::eq(snapshot.rotation, &rotation));
let mut mismatched_moe_cfg = dense_cfg.clone();
mismatched_moe_cfg.num_experts = Some(1);
assert!(!dense_cfg.is_moe());
assert!(mismatched_moe_cfg.is_moe());
let completion: for<'snapshot, 'reader, 'rotated> fn(
ForwardEquivalenceSnapshot<'snapshot>,
&'reader QuarotTensorReader,
&'rotated HashMap<String, TensorEntry>,
) -> Result<
ForwardEquivalenceReport,
InferenceError,
> = assert_prepared_forward_equivalence_qwen35;
let _ = (completion, snapshot, mismatched_moe_cfg);
}
#[test]
fn prepared_snapshot_enforces_configured_logit_budget_before_tensor_access() {
let cfg = Qwen35Config::qwen35_0_8b();
let original = HashMap::new();
let rotation = RandomizedHadamard::new(0x51A5_EEED, cfg.hidden_size).unwrap();
let bytes_per_probe =
cfg.vocab_size * size_of::<f64>() + size_of::<Vec<f64>>() + size_of::<u32>();
let max_probe_tokens = MAX_FORWARD_EQUIVALENCE_PROBE_SNAPSHOT_BYTES / bytes_per_probe;
assert_eq!(max_probe_tokens, 33);
assert!(max_probe_tokens * bytes_per_probe <= MAX_FORWARD_EQUIVALENCE_PROBE_SNAPSHOT_BYTES);
assert!(
(max_probe_tokens + 1) * bytes_per_probe > MAX_FORWARD_EQUIVALENCE_PROBE_SNAPSHOT_BYTES
);
let at_limit_cfg = ForwardEquivalenceConfig {
num_probe_tokens: max_probe_tokens,
..Default::default()
};
let at_limit =
prepare_forward_equivalence_qwen35(&original, &cfg, &rotation, &at_limit_cfg)
.err()
.expect("the empty fixture must fail at tensor access")
.to_string();
assert!(
at_limit.contains(QWEN35_EMBED_TOKENS_NAME),
"the largest in-budget probe count must reach tensor access: {at_limit}"
);
let over_limit_cfg = ForwardEquivalenceConfig {
num_probe_tokens: max_probe_tokens + 1,
..Default::default()
};
let over_limit =
prepare_forward_equivalence_qwen35(&original, &cfg, &rotation, &over_limit_cfg)
.err()
.expect("the over-budget request must fail admission")
.to_string();
assert!(
over_limit.contains("retained chain-logit budget"),
"one probe past the configured limit must be rejected before tensor access: \
{over_limit}"
);
assert!(
!over_limit.contains(QWEN35_EMBED_TOKENS_NAME),
"budget rejection must precede tensor access: {over_limit}"
);
}
#[test]
fn prepared_snapshot_rejects_over_budget_before_any_owned_allocation() {
let cfg = Qwen35Config::qwen35_0_8b();
let original = HashMap::new();
let rotation = RandomizedHadamard::new(0x51A5_EEED, cfg.hidden_size).unwrap();
let bytes_per_probe =
cfg.vocab_size * size_of::<f64>() + size_of::<Vec<f64>>() + size_of::<u32>();
let max_probe_tokens = MAX_FORWARD_EQUIVALENCE_PROBE_SNAPSHOT_BYTES / bytes_per_probe;
let forward_cfg = ForwardEquivalenceConfig {
num_probe_tokens: max_probe_tokens + 1,
..Default::default()
};
let tracking = pre_admission_allocation_tracking::start();
let result = prepare_forward_equivalence_qwen35(&original, &cfg, &rotation, &forward_cfg);
let observation = tracking.finish();
assert!(
observation.rejection_seen,
"the measured call must reach budget rejection"
);
assert_eq!(
observation.before_rejection_allocation_calls, 0,
"over-budget preparation allocated before rejecting"
);
assert!(
observation.after_rejection_allocation_calls > 0,
"the diagnostic allocation after budget rejection must be observed"
);
let error = result
.err()
.expect("the over-budget request must fail admission")
.to_string();
assert!(
error.contains("retained chain-logit budget"),
"unexpected error: {error}"
);
}
#[test]
fn prepared_snapshot_rejects_logit_budget_arithmetic_overflow() {
let original = HashMap::new();
for (num_probe_tokens, vocab_size) in [(2, usize::MAX / 2 + 1), (1, usize::MAX / 8 + 1)] {
let mut cfg = tied_tiny_test_cfg();
cfg.vocab_size = vocab_size;
let rotation = RandomizedHadamard::new(0x51A5_EEED, cfg.hidden_size).unwrap();
let forward_cfg = ForwardEquivalenceConfig {
num_probe_tokens,
..Default::default()
};
let err = prepare_forward_equivalence_qwen35(&original, &cfg, &rotation, &forward_cfg)
.err()
.expect("overflowing budget arithmetic must fail admission")
.to_string();
assert!(
err.contains("retained chain-logit size overflow"),
"overflowing budget arithmetic must return an error: {err}"
);
}
}
#[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_non_finite_chain_probe_error() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 32);
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(0xCA11_0BAD, cfg.hidden_size).unwrap();
absorb_rotations(&mut rotated, &rot_plan, &rotation).unwrap();
rotated.get_mut(QWEN35_FINAL_NORM_NAME).unwrap().data[0] = f64::NAN;
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("chain probe non-finite error"),
"unexpected error: {msg}"
);
assert!(msg.contains("abs_error=NaN"), "unexpected error: {msg}");
}
#[test]
fn forward_equivalence_refuses_non_finite_per_tensor_error() {
let cfg = tied_tiny_test_cfg();
let original = build_working_set(&cfg, 33);
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(0xFACE_FEED, 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";
rotated.get_mut(k_name).unwrap().data[0] = f64::NAN;
let err = assert_forward_equivalence_qwen35(
&original,
&rotated,
&cfg,
&rotation,
&ForwardEquivalenceConfig::default(),
)
.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("per-tensor non-finite error"),
"unexpected error: {msg}"
);
assert!(msg.contains(k_name), "unexpected error: {msg}");
assert!(msg.contains("abs_error=NaN"), "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}"
);
}
}