#![forbid(unsafe_code)]
use crate::attention::gdn::GatedDeltaNetWeights;
use crate::error::InferenceError;
use crate::model::qwen35::{
AttentionWeights, CommonLayerWeights, FeedForwardWeights, FullAttentionLayerWeights,
ModelWeights,
};
use crate::model::qwen35_config::{LayerType, Qwen35Config};
use crate::weights::ingress::{IngestedTensor, validate_ingested_tensor};
use std::mem::size_of;
const Q8_MATRIX_SOURCE: &str = "in-memory Q8 quantization";
const Q8_GDN_SOURCE: &str = "GatedDeltaNet Q8 quantization";
const Q8_ATTENTION_SOURCE: &str = "full-attention Q8 quantization";
const Q8_FFN_SOURCE: &str = "dense FFN Q8 quantization";
#[derive(Debug, Clone)]
pub struct Q8Matrix {
pub data: Vec<i8>,
pub scales: Vec<f32>,
pub rows: usize,
pub cols: usize,
}
pub fn quantize_matrix(w: &[f32], rows: usize, cols: usize) -> Result<Q8Matrix, InferenceError> {
quantize_named_matrix(w, rows, cols, Q8_MATRIX_SOURCE, "matrix")
}
fn validate_source_geometry(
len: usize,
rows: usize,
cols: usize,
source: &str,
tensor_name: &str,
) -> Result<(), InferenceError> {
let expected = rows.checked_mul(cols).ok_or_else(|| {
InferenceError::InvalidInput(format!(
"{source}: tensor {tensor_name} shape [{rows}, {cols}] overflows usize element count"
))
})?;
if len != expected {
return Err(InferenceError::InvalidInput(format!(
"{source}: tensor {tensor_name} (Q8) source element count {len} does not match shape \
[{rows}, {cols}] (expected {expected})"
)));
}
Ok(())
}
fn validate_cfg_shape(
rows: usize,
expected_rows: usize,
cols: usize,
expected_cols: usize,
source: &str,
tensor_name: &str,
) -> Result<(), InferenceError> {
if rows != expected_rows || cols != expected_cols {
return Err(InferenceError::InvalidInput(format!(
"{source}: tensor {tensor_name} shape [{rows}, {cols}] does not match config-derived \
shape [{expected_rows}, {expected_cols}]"
)));
}
Ok(())
}
pub(crate) fn validate_cfg_len(
len: usize,
expected: usize,
source: &str,
tensor_name: &str,
) -> Result<(), InferenceError> {
if len != expected {
return Err(InferenceError::InvalidInput(format!(
"{source}: tensor {tensor_name} length {len} does not match config-derived length \
{expected}"
)));
}
Ok(())
}
fn quantize_named_matrix(
w: &[f32],
rows: usize,
cols: usize,
source: &str,
tensor_name: &str,
) -> Result<Q8Matrix, InferenceError> {
validate_source_geometry(w.len(), rows, cols, source, tensor_name)?;
let shape = [rows, cols];
let mut data = Vec::with_capacity(w.len());
let mut scales = Vec::with_capacity(rows);
let mut has_nonfinite = false;
for row_idx in 0..rows {
let start = row_idx * cols;
let end = start + cols;
let row = &w[start..end];
let mut max_abs = 0.0f32;
for &v in row {
has_nonfinite |= !v.is_finite();
let abs_v = v.abs();
if abs_v > max_abs {
max_abs = abs_v;
}
}
let scale = if max_abs == 0.0 { 1.0 } else { max_abs / 127.0 };
scales.push(scale);
for &v in row {
let q = (v / scale).round().clamp(-128.0, 127.0) as i8;
data.push(q);
}
}
if has_nonfinite {
validate_ingested_tensor(IngestedTensor::q8_source(source, tensor_name, &shape, w))?;
unreachable!("has_nonfinite implies validate_ingested_tensor rejects tensor {tensor_name}");
}
validate_ingested_tensor(IngestedTensor::q8(
source,
tensor_name,
&shape,
&data,
&scales,
))?;
Ok(Q8Matrix {
data,
scales,
rows,
cols,
})
}
pub fn matmul_bt_q8(a: &[f32], b_q: &Q8Matrix, c: &mut [f32], m: usize, k: usize, n: usize) {
assert!(m.checked_mul(k).is_some(), "matmul shape overflow: m*k");
assert!(n.checked_mul(k).is_some(), "matmul shape overflow: n*k");
assert!(m.checked_mul(n).is_some(), "matmul shape overflow: m*n");
assert!(a.len() >= m * k, "A length does not match m * k");
assert_eq!(b_q.rows, n, "B_q rows do not match n");
assert_eq!(b_q.cols, k, "B_q cols do not match k");
assert_eq!(
b_q.data.len(),
n * k,
"B_q data length does not match n * k"
);
assert_eq!(b_q.scales.len(), n, "B_q scales length does not match n");
assert!(c.len() >= m * n, "C buffer is too small");
#[cfg(target_os = "macos")]
{
const TILE_N: usize = 64;
assert!(
TILE_N.checked_mul(k).is_some(),
"matmul shape overflow: TILE_N*k"
);
assert!(
m.checked_mul(TILE_N).is_some(),
"matmul shape overflow: m*TILE_N"
);
let mut b_f32 = vec![0.0f32; TILE_N * k];
let mut c_tile = vec![0.0f32; m * TILE_N];
for tile_start in (0..n).step_by(TILE_N) {
let tile_n = (n - tile_start).min(TILE_N);
for j in 0..tile_n {
let global_j = tile_start + j;
let q_row = &b_q.data[global_j * k..(global_j + 1) * k];
let scale = b_q.scales[global_j];
let dst = &mut b_f32[j * k..(j + 1) * k];
for t in 0..k {
dst[t] = q_row[t] as f32 * scale;
}
}
crate::forward::cpu::matmul_bt(
a,
&b_f32[..tile_n * k],
&mut c_tile[..m * tile_n],
m,
k,
tile_n,
);
if m == 1 {
c[tile_start..tile_start + tile_n].copy_from_slice(&c_tile[..tile_n]);
} else {
for i in 0..m {
let c_row = &mut c[i * n + tile_start..i * n + tile_start + tile_n];
let tile_row = &c_tile[i * tile_n..(i + 1) * tile_n];
c_row.copy_from_slice(tile_row);
}
}
}
}
#[cfg(not(target_os = "macos"))]
{
for i in 0..m {
let a_row = &a[i * k..(i + 1) * k];
let c_row = &mut c[i * n..(i + 1) * n];
for j in 0..n {
let q_row = &b_q.data[j * k..(j + 1) * k];
let mut acc = 0.0f32;
for t in 0..k {
acc += a_row[t] * q_row[t] as f32;
}
c_row[j] = acc * b_q.scales[j];
}
}
}
}
#[derive(Debug, Clone)]
pub struct Q8GatedDeltaNetWeights {
pub in_proj_qkv: Q8Matrix,
pub in_proj_z: Q8Matrix,
pub in_proj_b: Q8Matrix,
pub in_proj_a: Q8Matrix,
pub a_log: Vec<f32>,
pub dt_bias: Vec<f32>,
pub conv1d_weight: Vec<f32>,
pub conv_dim: usize,
pub kernel_size: usize,
pub norm_weight: Vec<f32>,
pub out_proj: Q8Matrix,
}
#[derive(Debug, Clone)]
pub struct Q8FullAttentionLayerWeights {
pub q_proj: Q8Matrix,
pub k_proj: Q8Matrix,
pub v_proj: Q8Matrix,
pub o_proj: Q8Matrix,
pub q_norm: Vec<f32>,
pub k_norm: Vec<f32>,
}
#[derive(Debug, Clone)]
pub struct Q8CommonLayerWeights {
pub input_layernorm: Vec<f32>,
pub post_attention_layernorm: Vec<f32>,
pub gate_proj: Q8Matrix,
pub up_proj: Q8Matrix,
pub down_proj: Q8Matrix,
}
#[derive(Debug, Clone)]
pub enum Q8AttentionWeights {
Linear(Q8GatedDeltaNetWeights),
Full(Q8FullAttentionLayerWeights),
}
#[derive(Debug, Clone)]
pub struct Q8ModelWeights {
pub embed_tokens: Vec<f32>,
pub final_norm: Vec<f32>,
pub layers: Vec<(Q8AttentionWeights, Q8CommonLayerWeights)>,
}
pub(crate) fn quantize_model_weights(
weights: &ModelWeights,
cfg: &Qwen35Config,
) -> Result<Q8ModelWeights, InferenceError> {
let layers = weights
.layers
.iter()
.map(|(attn, common)| {
let q8_attn = match attn {
AttentionWeights::Linear(gdn) => {
Q8AttentionWeights::Linear(quantize_gdn_weights(gdn, cfg)?)
}
AttentionWeights::Full(full) => {
Q8AttentionWeights::Full(quantize_full_attn_weights(full, cfg)?)
}
};
let q8_common = quantize_common_weights(common, cfg)?;
Ok((q8_attn, q8_common))
})
.collect::<Result<Vec<_>, InferenceError>>()?;
Ok(Q8ModelWeights {
embed_tokens: weights.embed_tokens.clone(),
final_norm: weights.final_norm.clone(),
layers,
})
}
pub(crate) fn validate_gdn_shapes(
w: &GatedDeltaNetWeights,
cfg: &Qwen35Config,
) -> Result<(), InferenceError> {
let hidden = cfg.hidden_size;
let value_heads = cfg.linear_num_value_heads();
let qkv_dim = cfg.checked_linear_qkv_dim()?;
let output_dim = cfg.checked_linear_output_dim()?;
validate_cfg_shape(
w.in_proj_qkv_rows,
qkv_dim,
w.in_proj_qkv_cols,
hidden,
Q8_GDN_SOURCE,
"in_proj_qkv",
)?;
validate_cfg_shape(
w.in_proj_z_rows,
output_dim,
w.in_proj_z_cols,
hidden,
Q8_GDN_SOURCE,
"in_proj_z",
)?;
validate_cfg_shape(
w.in_proj_b_rows,
value_heads,
w.in_proj_b_cols,
hidden,
Q8_GDN_SOURCE,
"in_proj_b",
)?;
validate_cfg_shape(
w.in_proj_a_rows,
value_heads,
w.in_proj_a_cols,
hidden,
Q8_GDN_SOURCE,
"in_proj_a",
)?;
validate_cfg_shape(
w.out_proj_rows,
hidden,
w.out_proj_cols,
output_dim,
Q8_GDN_SOURCE,
"out_proj",
)?;
validate_source_geometry(
w.in_proj_qkv.len(),
w.in_proj_qkv_rows,
w.in_proj_qkv_cols,
Q8_GDN_SOURCE,
"in_proj_qkv",
)?;
validate_source_geometry(
w.in_proj_z.len(),
w.in_proj_z_rows,
w.in_proj_z_cols,
Q8_GDN_SOURCE,
"in_proj_z",
)?;
validate_source_geometry(
w.in_proj_b.len(),
w.in_proj_b_rows,
w.in_proj_b_cols,
Q8_GDN_SOURCE,
"in_proj_b",
)?;
validate_source_geometry(
w.in_proj_a.len(),
w.in_proj_a_rows,
w.in_proj_a_cols,
Q8_GDN_SOURCE,
"in_proj_a",
)?;
validate_source_geometry(
w.out_proj.len(),
w.out_proj_rows,
w.out_proj_cols,
Q8_GDN_SOURCE,
"out_proj",
)?;
validate_cfg_len(w.a_log.len(), value_heads, Q8_GDN_SOURCE, "a_log")?;
validate_cfg_len(w.dt_bias.len(), value_heads, Q8_GDN_SOURCE, "dt_bias")?;
if w.conv_dim != qkv_dim {
return Err(InferenceError::InvalidInput(format!(
"{Q8_GDN_SOURCE}: tensor conv_dim {} does not match config-derived linear_qkv_dim {qkv_dim}",
w.conv_dim
)));
}
let expected_conv_len = cfg.checked_linear_conv_len()?;
validate_cfg_len(
w.conv1d_weight.len(),
expected_conv_len,
Q8_GDN_SOURCE,
"conv1d_weight",
)?;
validate_cfg_len(
w.norm_weight.len(),
cfg.linear_value_head_dim,
Q8_GDN_SOURCE,
"norm_weight",
)?;
validate_ingested_tensor(IngestedTensor::q8_source(
Q8_GDN_SOURCE,
"a_log",
&[w.a_log.len()],
&w.a_log,
))?;
validate_ingested_tensor(IngestedTensor::q8_source(
Q8_GDN_SOURCE,
"dt_bias",
&[w.dt_bias.len()],
&w.dt_bias,
))?;
validate_ingested_tensor(IngestedTensor::q8_source(
Q8_GDN_SOURCE,
"conv1d_weight",
&[w.conv1d_weight.len()],
&w.conv1d_weight,
))?;
validate_ingested_tensor(IngestedTensor::q8_source(
Q8_GDN_SOURCE,
"norm_weight",
&[w.norm_weight.len()],
&w.norm_weight,
))?;
Ok(())
}
pub fn quantize_gdn_weights(
w: &GatedDeltaNetWeights,
cfg: &Qwen35Config,
) -> Result<Q8GatedDeltaNetWeights, InferenceError> {
validate_gdn_shapes(w, cfg)?;
let in_proj_qkv = quantize_named_matrix(
&w.in_proj_qkv,
w.in_proj_qkv_rows,
w.in_proj_qkv_cols,
Q8_GDN_SOURCE,
"in_proj_qkv",
)?;
let in_proj_z = quantize_named_matrix(
&w.in_proj_z,
w.in_proj_z_rows,
w.in_proj_z_cols,
Q8_GDN_SOURCE,
"in_proj_z",
)?;
let in_proj_b = quantize_named_matrix(
&w.in_proj_b,
w.in_proj_b_rows,
w.in_proj_b_cols,
Q8_GDN_SOURCE,
"in_proj_b",
)?;
let in_proj_a = quantize_named_matrix(
&w.in_proj_a,
w.in_proj_a_rows,
w.in_proj_a_cols,
Q8_GDN_SOURCE,
"in_proj_a",
)?;
let out_proj = quantize_named_matrix(
&w.out_proj,
w.out_proj_rows,
w.out_proj_cols,
Q8_GDN_SOURCE,
"out_proj",
)?;
Ok(Q8GatedDeltaNetWeights {
in_proj_qkv,
in_proj_z,
in_proj_b,
in_proj_a,
a_log: w.a_log.clone(),
dt_bias: w.dt_bias.clone(),
conv1d_weight: w.conv1d_weight.clone(),
conv_dim: w.conv_dim,
kernel_size: w.kernel_size,
norm_weight: w.norm_weight.clone(),
out_proj,
})
}
pub(crate) fn quantize_full_attn_weights(
w: &FullAttentionLayerWeights,
cfg: &Qwen35Config,
) -> Result<Q8FullAttentionLayerWeights, InferenceError> {
let ((q_rows, hidden), (kv_rows, _), (_, _), (o_rows, o_cols)) =
infer_full_attention_shapes(w, cfg)?;
let q_proj = quantize_named_matrix(&w.q_proj, q_rows, hidden, Q8_ATTENTION_SOURCE, "q_proj")?;
let k_proj = quantize_named_matrix(&w.k_proj, kv_rows, hidden, Q8_ATTENTION_SOURCE, "k_proj")?;
let v_proj = quantize_named_matrix(&w.v_proj, kv_rows, hidden, Q8_ATTENTION_SOURCE, "v_proj")?;
let o_proj = quantize_named_matrix(&w.o_proj, o_rows, o_cols, Q8_ATTENTION_SOURCE, "o_proj")?;
Ok(Q8FullAttentionLayerWeights {
q_proj,
k_proj,
v_proj,
o_proj,
q_norm: w.q_norm.clone(),
k_norm: w.k_norm.clone(),
})
}
pub(crate) fn quantize_common_weights(
w: &CommonLayerWeights,
cfg: &Qwen35Config,
) -> Result<Q8CommonLayerWeights, InferenceError> {
let dense = match &w.ffn {
FeedForwardWeights::Dense(d) => d,
FeedForwardWeights::Moe(_) => {
return Err(InferenceError::UnsupportedModel(
"Q8 quantization is dense-only; MoE layers are not supported".to_string(),
));
}
};
let hidden = cfg.hidden_size;
validate_cfg_len(
w.input_layernorm.len(),
hidden,
Q8_FFN_SOURCE,
"input_layernorm",
)?;
validate_cfg_len(
w.post_attention_layernorm.len(),
hidden,
Q8_FFN_SOURCE,
"post_attention_layernorm",
)?;
let ((gate_rows, _), (up_rows, _), (down_rows, down_cols)) = infer_dense_shapes(dense, cfg)?;
let gate_proj = quantize_named_matrix(
&dense.gate_proj,
gate_rows,
hidden,
Q8_FFN_SOURCE,
"gate_proj",
)?;
let up_proj = quantize_named_matrix(&dense.up_proj, up_rows, hidden, Q8_FFN_SOURCE, "up_proj")?;
let down_proj = quantize_named_matrix(
&dense.down_proj,
down_rows,
down_cols,
Q8_FFN_SOURCE,
"down_proj",
)?;
Ok(Q8CommonLayerWeights {
input_layernorm: w.input_layernorm.clone(),
post_attention_layernorm: w.post_attention_layernorm.clone(),
gate_proj,
up_proj,
down_proj,
})
}
pub fn memory_report(cfg: &Qwen35Config) -> (usize, usize, f32) {
let hidden = cfg.hidden_size;
let inter = cfg.intermediate_size;
let qkv_dim = cfg.linear_qkv_dim();
let linear_output_dim = cfg.linear_output_dim();
let linear_heads = cfg.linear_num_value_heads();
let q_dim = cfg.full_q_dim();
let kv_dim = cfg.full_kv_dim();
let mut f32_bytes = 0usize;
let mut q8_bytes = 0usize;
for layer_type in &cfg.layer_types {
add_matrix_bytes(inter, hidden, &mut f32_bytes, &mut q8_bytes); add_matrix_bytes(inter, hidden, &mut f32_bytes, &mut q8_bytes); add_matrix_bytes(hidden, inter, &mut f32_bytes, &mut q8_bytes);
match layer_type {
LayerType::LinearAttention => {
add_matrix_bytes(qkv_dim, hidden, &mut f32_bytes, &mut q8_bytes);
add_matrix_bytes(linear_output_dim, hidden, &mut f32_bytes, &mut q8_bytes);
add_matrix_bytes(linear_heads, hidden, &mut f32_bytes, &mut q8_bytes);
add_matrix_bytes(linear_heads, hidden, &mut f32_bytes, &mut q8_bytes);
add_matrix_bytes(hidden, linear_output_dim, &mut f32_bytes, &mut q8_bytes);
}
LayerType::FullAttention => {
add_matrix_bytes(2 * q_dim, hidden, &mut f32_bytes, &mut q8_bytes);
add_matrix_bytes(kv_dim, hidden, &mut f32_bytes, &mut q8_bytes);
add_matrix_bytes(kv_dim, hidden, &mut f32_bytes, &mut q8_bytes);
add_matrix_bytes(hidden, q_dim, &mut f32_bytes, &mut q8_bytes);
}
}
}
let savings_ratio = if q8_bytes == 0 {
f32::INFINITY
} else {
f32_bytes as f32 / q8_bytes as f32
};
(f32_bytes, q8_bytes, savings_ratio)
}
#[inline]
fn add_matrix_bytes(rows: usize, cols: usize, f32_bytes: &mut usize, q8_bytes: &mut usize) {
*f32_bytes += rows * cols * size_of::<f32>();
*q8_bytes += rows * cols * size_of::<i8>() + rows * size_of::<f32>();
}
type ShapePair = (usize, usize);
fn infer_dense_shapes(
dense: &crate::model::qwen35::DenseFfnWeights,
cfg: &Qwen35Config,
) -> Result<(ShapePair, ShapePair, ShapePair), InferenceError> {
let hidden = cfg.hidden_size;
let inter = cfg.intermediate_size;
if hidden == 0 {
return Err(InferenceError::InvalidInput(format!(
"{Q8_FFN_SOURCE}: cfg.hidden_size must be > 0"
)));
}
validate_source_geometry(
dense.gate_proj.len(),
inter,
hidden,
Q8_FFN_SOURCE,
"gate_proj",
)?;
validate_source_geometry(dense.up_proj.len(), inter, hidden, Q8_FFN_SOURCE, "up_proj")?;
validate_source_geometry(
dense.down_proj.len(),
hidden,
inter,
Q8_FFN_SOURCE,
"down_proj",
)?;
Ok(((inter, hidden), (inter, hidden), (hidden, inter)))
}
fn infer_full_attention_shapes(
w: &FullAttentionLayerWeights,
cfg: &Qwen35Config,
) -> Result<(ShapePair, ShapePair, ShapePair, ShapePair), InferenceError> {
let head_dim = cfg.head_dim;
let hidden = cfg.hidden_size;
let q_dim = cfg.checked_full_q_dim()?;
let kv_dim = cfg.checked_full_kv_dim()?;
let q_rows = crate::model::qwen35_config::checked_double(q_dim, "full_q_dim")?;
validate_cfg_len(w.q_norm.len(), head_dim, Q8_ATTENTION_SOURCE, "q_norm")?;
validate_cfg_len(w.k_norm.len(), head_dim, Q8_ATTENTION_SOURCE, "k_norm")?;
validate_source_geometry(
w.q_proj.len(),
q_rows,
hidden,
Q8_ATTENTION_SOURCE,
"q_proj",
)?;
validate_source_geometry(
w.k_proj.len(),
kv_dim,
hidden,
Q8_ATTENTION_SOURCE,
"k_proj",
)?;
validate_source_geometry(
w.v_proj.len(),
kv_dim,
hidden,
Q8_ATTENTION_SOURCE,
"v_proj",
)?;
validate_source_geometry(w.o_proj.len(), hidden, q_dim, Q8_ATTENTION_SOURCE, "o_proj")?;
Ok((
(q_rows, hidden),
(kv_dim, hidden),
(kv_dim, hidden),
(hidden, q_dim),
))
}
#[cfg(test)]
fn dequantize_matrix(q: &Q8Matrix) -> Vec<f32> {
let mut out = vec![0.0f32; q.rows * q.cols];
for row in 0..q.rows {
let scale = q.scales[row];
let start = row * q.cols;
let end = start + q.cols;
for idx in start..end {
out[idx] = q.data[idx] as f32 * scale;
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::forward::cpu::matmul_bt;
fn approx_eq(a: f32, b: f32, tol: f32) -> bool {
(a - b).abs() <= tol
}
fn xorshift32(state: &mut u32) -> u32 {
let mut x = *state;
x ^= x << 13;
x ^= x >> 17;
x ^= x << 5;
*state = x;
x
}
fn uniform_signed(state: &mut u32) -> f32 {
let x = xorshift32(state);
let u = x as f32 / u32::MAX as f32;
u * 2.0 - 1.0
}
#[test]
fn test_quantize_identity() {
let rows = 4;
let cols = 4;
let w = vec![
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, ];
let q = quantize_matrix(&w, rows, cols).unwrap();
let dq = dequantize_matrix(&q);
assert_eq!(q.rows, rows);
assert_eq!(q.cols, cols);
assert_eq!(q.scales.len(), rows);
assert_eq!(q.data.len(), rows * cols);
for &scale in &q.scales {
assert!(approx_eq(scale, 1.0 / 127.0, 1e-8));
}
for (orig, recon) in w.iter().zip(dq.iter()) {
assert!(approx_eq(*orig, *recon, 1e-6));
}
}
const GDN_VALUE_HEAD_DIM: usize = 2;
fn valid_gdn_weights(value_heads: usize, hidden: usize) -> GatedDeltaNetWeights {
let qkv_rows = value_heads * 4;
let z_rows = value_heads * GDN_VALUE_HEAD_DIM;
GatedDeltaNetWeights {
in_proj_qkv: vec![0.0; qkv_rows * hidden],
in_proj_qkv_rows: qkv_rows,
in_proj_qkv_cols: hidden,
in_proj_z: vec![0.0; z_rows * hidden],
in_proj_z_rows: z_rows,
in_proj_z_cols: hidden,
in_proj_b: vec![0.0; value_heads * hidden],
in_proj_b_rows: value_heads,
in_proj_b_cols: hidden,
in_proj_a: vec![0.0; value_heads * hidden],
in_proj_a_rows: value_heads,
in_proj_a_cols: hidden,
a_log: vec![0.0; value_heads],
dt_bias: vec![0.0; value_heads],
conv1d_weight: vec![0.0; qkv_rows * 4],
conv_dim: qkv_rows,
kernel_size: 4,
norm_weight: vec![1.0; GDN_VALUE_HEAD_DIM],
out_proj: vec![0.0; hidden * z_rows],
out_proj_rows: hidden,
out_proj_cols: z_rows,
}
}
fn matching_gdn_cfg() -> Qwen35Config {
Qwen35Config {
hidden_size: 8,
linear_num_key_heads: 1,
linear_key_head_dim: 3,
linear_num_value_heads: Some(3),
linear_value_head_dim: GDN_VALUE_HEAD_DIM,
linear_conv_kernel_dim: 4,
..Qwen35Config::qwen35_2b()
}
}
#[test]
fn quantize_gdn_weights_accepts_consistent_value_head_shapes() {
let w = valid_gdn_weights(3, 8);
let q = quantize_gdn_weights(&w, &matching_gdn_cfg())
.expect("consistent value-head shapes must quantize");
assert_eq!(q.a_log.len(), 3);
assert_eq!(q.dt_bias.len(), 3);
}
#[test]
fn quantize_gdn_weights_rejects_decay_shape_mismatch() {
let mut w = valid_gdn_weights(3, 8);
w.a_log = vec![0.0; 4];
match quantize_gdn_weights(&w, &matching_gdn_cfg()) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(msg.contains("a_log"), "error must name a_log, got: {msg}");
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!("expected Err for decay shape mismatch, got Ok"),
}
}
#[test]
fn quantize_gdn_weights_rejects_malformed_conv1d_weight() {
let mut w = valid_gdn_weights(3, 8);
w.conv1d_weight.pop(); match quantize_gdn_weights(&w, &matching_gdn_cfg()) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("conv1d_weight"),
"error must name conv1d_weight, got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!("expected Err for malformed conv1d_weight, got Ok"),
}
}
#[test]
fn quantize_gdn_weights_rejects_conv_dim_disagreeing_with_cfg() {
let mut w = valid_gdn_weights(3, 8);
w.conv_dim = 16;
match quantize_gdn_weights(&w, &matching_gdn_cfg()) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("conv_dim"),
"error must name conv_dim, got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!("expected Err for conv_dim disagreeing with cfg, got Ok"),
}
}
#[test]
fn quantize_gdn_weights_accepts_conv_dim_matching_cfg() {
let w = valid_gdn_weights(3, 8);
assert_eq!(w.conv_dim, 12);
assert!(quantize_gdn_weights(&w, &matching_gdn_cfg()).is_ok());
}
#[test]
fn quantize_gdn_weights_rejects_malformed_norm_weight() {
let mut w = valid_gdn_weights(3, 8);
w.norm_weight = vec![1.0; 1]; match quantize_gdn_weights(&w, &matching_gdn_cfg()) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("norm_weight"),
"error must name norm_weight, got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!("expected Err for malformed norm_weight, got Ok"),
}
}
#[test]
fn quantize_gdn_weights_rejects_value_heads_disagreeing_with_cfg() {
let mut w = valid_gdn_weights(3, 8);
w.in_proj_b = vec![0.0; 5 * 8];
w.in_proj_b_rows = 5;
w.in_proj_b_cols = 8;
w.in_proj_a = vec![0.0; 5 * 8];
w.in_proj_a_rows = 5;
w.in_proj_a_cols = 8;
w.a_log = vec![0.0; 5];
w.dt_bias = vec![0.0; 5];
match quantize_gdn_weights(&w, &matching_gdn_cfg()) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("in_proj_b"),
"error must name in_proj_b, got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!(
"expected Err for value-head count disagreeing with cfg, got Ok \
(Q8 matrices incompatible with cfg would have been produced)"
),
}
}
#[test]
fn quantize_gdn_weights_rejects_non_finite_retained_vectors() {
let cfg = matching_gdn_cfg();
let mut w = valid_gdn_weights(3, 8);
w.a_log[0] = f32::NAN;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
let mut w = valid_gdn_weights(3, 8);
w.a_log[0] = f32::INFINITY;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
let mut w = valid_gdn_weights(3, 8);
w.dt_bias[0] = f32::NAN;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
let mut w = valid_gdn_weights(3, 8);
w.dt_bias[0] = f32::INFINITY;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
let mut w = valid_gdn_weights(3, 8);
w.conv1d_weight[0] = f32::NAN;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
let mut w = valid_gdn_weights(3, 8);
w.conv1d_weight[0] = f32::INFINITY;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
let mut w = valid_gdn_weights(3, 8);
w.norm_weight[0] = f32::NAN;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
let mut w = valid_gdn_weights(3, 8);
w.norm_weight[0] = f32::INFINITY;
assert!(matches!(
quantize_gdn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
fn quantize_gdn_weights_accepts_all_finite_retained_vectors() {
let w = valid_gdn_weights(3, 8);
assert!(quantize_gdn_weights(&w, &matching_gdn_cfg()).is_ok());
}
#[test]
fn checked_full_q_dim_rejects_overflow() {
let cfg = Qwen35Config {
num_attention_heads: 1 << 63,
head_dim: 2,
..Qwen35Config::qwen35_2b()
};
assert!(matches!(
cfg.checked_full_q_dim(),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
fn checked_full_q_dim_accepts_valid_config() {
let cfg = Qwen35Config {
num_attention_heads: 16,
head_dim: 64,
..Qwen35Config::qwen35_2b()
};
assert_eq!(cfg.checked_full_q_dim().unwrap(), 16 * 64);
}
#[test]
fn checked_full_kv_dim_rejects_overflow() {
let cfg = Qwen35Config {
num_key_value_heads: 1 << 63,
head_dim: 2,
..Qwen35Config::qwen35_2b()
};
assert!(matches!(
cfg.checked_full_kv_dim(),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
fn checked_full_kv_dim_accepts_valid_config() {
let cfg = Qwen35Config {
num_key_value_heads: 2,
head_dim: 64,
..Qwen35Config::qwen35_2b()
};
assert_eq!(cfg.checked_full_kv_dim().unwrap(), 2 * 64);
}
#[test]
fn checked_linear_qkv_dim_rejects_overflow() {
let cfg = Qwen35Config {
linear_num_key_heads: 1 << 63,
linear_key_head_dim: 2,
..Qwen35Config::qwen35_2b()
};
assert!(matches!(
cfg.checked_linear_qkv_dim(),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
fn checked_linear_qkv_dim_accepts_valid_config() {
let cfg = matching_gdn_cfg();
assert_eq!(cfg.checked_linear_qkv_dim().unwrap(), cfg.linear_qkv_dim());
assert_eq!(cfg.checked_linear_qkv_dim().unwrap(), 12);
}
#[test]
fn checked_linear_output_dim_rejects_overflow() {
let cfg = Qwen35Config {
linear_num_value_heads: Some(1 << 63),
linear_value_head_dim: 2,
..Qwen35Config::qwen35_2b()
};
assert!(matches!(
cfg.checked_linear_output_dim(),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
fn checked_linear_output_dim_accepts_valid_config() {
let cfg = matching_gdn_cfg();
assert_eq!(
cfg.checked_linear_output_dim().unwrap(),
cfg.linear_output_dim()
);
assert_eq!(cfg.checked_linear_output_dim().unwrap(), 6);
}
#[test]
fn checked_linear_conv_len_rejects_overflow() {
let cfg = Qwen35Config {
linear_num_key_heads: 1 << 63,
linear_key_head_dim: 2,
linear_conv_kernel_dim: 4,
..Qwen35Config::qwen35_2b()
};
assert!(matches!(
cfg.checked_linear_conv_len(),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
fn checked_linear_conv_len_accepts_valid_config() {
let cfg = matching_gdn_cfg();
assert_eq!(
cfg.checked_linear_conv_len().unwrap(),
cfg.linear_qkv_dim() * cfg.linear_conv_kernel_dim
);
}
#[test]
fn quantize_full_attn_weights_rejects_overflowing_config() {
let cfg = Qwen35Config {
num_attention_heads: 1 << 63,
num_key_value_heads: 1,
head_dim: 2,
..Qwen35Config::qwen35_2b()
};
let w = FullAttentionLayerWeights {
q_proj: vec![],
k_proj: vec![],
v_proj: vec![],
o_proj: vec![],
q_norm: vec![0.0; cfg.head_dim],
k_norm: vec![0.0; cfg.head_dim],
};
assert!(matches!(
quantize_full_attn_weights(&w, &cfg),
Err(InferenceError::InvalidInput(_))
));
}
#[test]
fn test_quantize_roundtrip_accuracy() {
let rows = 32;
let cols = 64;
let mut seed = 0x1234_5678u32;
let mut w = Vec::with_capacity(rows * cols);
for _ in 0..rows * cols {
w.push(uniform_signed(&mut seed) * 0.03);
}
let q = quantize_matrix(&w, rows, cols).unwrap();
let dq = dequantize_matrix(&q);
let mut max_abs_err = 0.0f32;
let mut mean_abs_err = 0.0f32;
for (orig, recon) in w.iter().zip(dq.iter()) {
let err = (orig - recon).abs();
if err > max_abs_err {
max_abs_err = err;
}
mean_abs_err += err;
}
mean_abs_err /= (rows * cols) as f32;
assert!(max_abs_err < 2.0e-4, "max_abs_err={max_abs_err}");
assert!(mean_abs_err < 7.0e-5, "mean_abs_err={mean_abs_err}");
}
#[test]
fn test_matmul_bt_q8_matches_f32() {
let m = 2;
let k = 256;
let n = 16;
let mut seed = 0xCAFE_BABEu32;
let mut a = Vec::with_capacity(m * k);
let mut b = Vec::with_capacity(n * k);
for _ in 0..m * k {
a.push(uniform_signed(&mut seed) * 0.1);
}
for row in 0..n {
let scale = (row as f32 + 1.0) / 8_192.0;
b.push(127.0 * scale);
for col in 1..k {
let qv = ((row * 17 + col * 13) % 255) as i32 - 127;
b.push(qv as f32 * scale);
}
}
let b_q = quantize_matrix(&b, n, k).unwrap();
let mut c_f32 = vec![0.0f32; m * n];
let mut c_q8 = vec![0.0f32; m * n];
matmul_bt(&a, &b, &mut c_f32, m, k, n);
matmul_bt_q8(&a, &b_q, &mut c_q8, m, k, n);
let mut max_abs_err = 0.0f32;
let mut max_rel_err = 0.0f32;
let mut large_outputs = 0usize;
for (&ref_v, &q_v) in c_f32.iter().zip(c_q8.iter()) {
let abs_err = (ref_v - q_v).abs();
if abs_err > max_abs_err {
max_abs_err = abs_err;
}
if ref_v.abs() > 0.01 {
large_outputs += 1;
let rel_err = abs_err / ref_v.abs();
if rel_err > max_rel_err {
max_rel_err = rel_err;
}
}
}
assert!(large_outputs > 0, "expected some non-trivial outputs");
assert!(max_abs_err < 1.0e-4, "max_abs_err={max_abs_err}");
assert!(max_rel_err < 0.01, "max_rel_err={max_rel_err}");
}
#[test]
fn test_matmul_bt_q8_known_values() {
let a = vec![1.0, 2.0, -1.0]; let b = vec![
1.27, -0.64, 0.0, -2.54, 0.0, 1.28, ];
let q = quantize_matrix(&b, 2, 3).unwrap();
assert!(approx_eq(q.scales[0], 0.01, 1e-8));
assert!(approx_eq(q.scales[1], 0.02, 1e-8));
assert_eq!(&q.data[0..3], &[127i8, -64i8, 0i8]);
assert_eq!(&q.data[3..6], &[-127i8, 0i8, 64i8]);
let mut c = vec![0.0f32; 2];
matmul_bt_q8(&a, &q, &mut c, 1, 3, 2);
let expected0 = 0.01 * (1.0 * 127.0 + 2.0 * -64.0 + -0.0);
let expected1 = 0.02 * (1.0 * -127.0 + 2.0 * 0.0 + -64.0);
assert!(approx_eq(c[0], expected0, 1e-6));
assert!(approx_eq(c[1], expected1, 1e-6));
assert!(approx_eq(c[0], -0.01, 1e-6));
assert!(approx_eq(c[1], -3.82, 1e-6));
}
fn dummy_q8_1x1() -> Q8Matrix {
Q8Matrix {
data: vec![0i8],
scales: vec![0.0f32],
rows: 1,
cols: 1,
}
}
#[test]
#[should_panic(expected = "matmul shape overflow: m*k")]
fn test_matmul_bt_q8_rejects_mk_overflow() {
let q = dummy_q8_1x1();
let mut c = vec![0.0f32; 1];
matmul_bt_q8(&[], &q, &mut c, usize::MAX, 2, 1);
}
#[test]
#[should_panic(expected = "matmul shape overflow: n*k")]
fn test_matmul_bt_q8_rejects_nk_overflow() {
let q = dummy_q8_1x1();
let mut c = vec![0.0f32; 1];
matmul_bt_q8(&[0.0, 0.0], &q, &mut c, 1, 2, usize::MAX);
}
#[test]
#[should_panic(expected = "matmul shape overflow: m*n")]
fn test_matmul_bt_q8_rejects_mn_overflow() {
let q = dummy_q8_1x1();
let mut c = vec![0.0f32; 1];
matmul_bt_q8(&[], &q, &mut c, usize::MAX, 1, 2);
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "matmul shape overflow: TILE_N*k")]
fn test_matmul_bt_q8_rejects_tile_nk_overflow() {
let q = Q8Matrix {
data: vec![],
scales: vec![],
rows: 0,
cols: usize::MAX,
};
matmul_bt_q8(&[], &q, &mut [], 0, usize::MAX, 0);
}
#[cfg(target_os = "macos")]
#[test]
#[should_panic(expected = "matmul shape overflow: m*TILE_N")]
fn test_matmul_bt_q8_rejects_m_tile_n_overflow() {
let q = Q8Matrix {
data: vec![],
scales: vec![],
rows: 0,
cols: 0,
};
matmul_bt_q8(&[], &q, &mut [], usize::MAX, 0, 0);
}
#[test]
fn test_zero_row_handling() {
let b = vec![
0.0, 0.0, 0.0, 0.0, 0.5, -0.5, 0.25, -0.25,
];
let q = quantize_matrix(&b, 2, 4).unwrap();
assert!(approx_eq(q.scales[0], 1.0, 1e-8));
assert_eq!(&q.data[0..4], &[0i8, 0i8, 0i8, 0i8]);
let a = vec![1.0, 2.0, 3.0, 4.0];
let mut c = vec![0.0f32; 2];
matmul_bt_q8(&a, &q, &mut c, 1, 4, 2);
assert!(approx_eq(c[0], 0.0, 1e-8));
assert!(c[1].abs() > 0.0);
}
#[test]
fn test_scale_computation() {
let w = vec![
1.27, -0.63, 0.0, -2.54, 2.0, 0.0, ];
let q = quantize_matrix(&w, 2, 3).unwrap();
assert!(approx_eq(q.scales[0], 0.01, 1e-8));
assert!(approx_eq(q.scales[1], 0.02, 1e-8));
assert_eq!(&q.data[0..3], &[127i8, -63i8, 0i8]);
assert_eq!(&q.data[3..6], &[-127i8, 100i8, 0i8]);
}
#[test]
fn test_memory_report() {
let cfg = Qwen35Config::qwen35_2b();
let (f32_bytes, q8_bytes, ratio) = memory_report(&cfg);
assert_eq!(f32_bytes, 5_490_868_224);
assert_eq!(q8_bytes, 1_375_004_928);
assert!(q8_bytes < f32_bytes);
assert!(ratio > 3.9 && ratio < 4.05, "ratio={ratio}");
}
#[test]
fn test_quantize_large_matrix() {
let rows = 6_144;
let cols = 2_048;
let mut w = Vec::with_capacity(rows * cols);
for r in 0..rows {
let row_scale = 0.005 + (r % 17) as f32 * 0.0005;
for c in 0..cols {
let bucket = ((r * 131 + c * 17) % 257) as i32 - 128;
w.push(bucket as f32 * row_scale / 128.0);
}
}
let q = quantize_matrix(&w, rows, cols).unwrap();
assert_eq!(q.rows, rows);
assert_eq!(q.cols, cols);
assert_eq!(q.data.len(), rows * cols);
assert_eq!(q.scales.len(), rows);
assert!(q.scales.iter().all(|&s| s.is_finite() && s > 0.0));
assert!(q.data.iter().any(|&v| v != 0));
let first_row_scale = q.scales[0];
let last_row_scale = q.scales[rows - 1];
assert!(first_row_scale > 0.0);
assert!(last_row_scale > 0.0);
}
#[test]
fn test_quantize_common_weights_moe_returns_err_not_panic() {
use crate::error::InferenceError;
use crate::model::qwen35::{
CommonLayerWeights, FeedForwardWeights, MoeLayerWeights, MoeRouter, RoutedExperts,
SharedExpert,
};
let hidden = 4usize;
let num_experts = 2usize;
let num_experts_per_tok = 1usize;
let inter = 2usize;
let router = MoeRouter::new(
vec![0.0f32; num_experts * hidden],
num_experts,
num_experts_per_tok,
hidden,
)
.expect("valid router shape");
let experts = RoutedExperts::new(
vec![0.0f32; num_experts * 2 * inter * hidden],
vec![0.0f32; num_experts * hidden * inter],
num_experts,
hidden,
inter,
)
.expect("valid experts shape");
let shared_expert = SharedExpert::new(
vec![0.0f32; inter * hidden],
vec![0.0f32; inter * hidden],
vec![0.0f32; hidden * inter],
vec![0.0f32; hidden],
hidden,
inter,
)
.expect("valid shared expert shape");
let moe_common = CommonLayerWeights {
input_layernorm: vec![0.0f32; hidden],
post_attention_layernorm: vec![0.0f32; hidden],
ffn: FeedForwardWeights::Moe(MoeLayerWeights {
router,
experts,
shared_expert,
}),
};
let cfg = Qwen35Config {
hidden_size: hidden,
intermediate_size: inter,
..Qwen35Config::qwen35_2b()
};
let result = quantize_common_weights(&moe_common, &cfg);
assert!(
matches!(result, Err(InferenceError::UnsupportedModel(_))),
"expected Err(UnsupportedModel), got: {result:?}"
);
}
#[test]
fn test_quantize_matrix_rejects_nonfinite_source_row() {
let w = vec![
1.0,
2.0,
3.0, 1.0,
f32::INFINITY,
0.0, ];
match quantize_matrix(&w, 2, 3) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("non-finite"),
"error must describe the non-finite source; got: {msg}"
);
assert!(msg.contains("element index 4"), "wrong attribution: {msg}");
}
Err(e) => panic!("expected InvalidInput for non-finite source, got: {e}"),
Ok(_) => panic!("expected Err for non-finite source row, got Ok"),
}
}
#[test]
fn test_quantize_matrix_rejects_nan_input() {
let w = vec![
1.0,
2.0,
3.0, 1.0,
f32::NAN,
0.0, ];
match quantize_matrix(&w, 2, 3) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("non-finite"),
"error must describe the non-finite value; got: {msg}"
);
assert!(msg.contains("element index 4"), "wrong attribution: {msg}");
}
Err(e) => panic!("expected InvalidInput for NaN input, got: {e}"),
Ok(_) => panic!("expected Err for NaN input, got Ok"),
}
}
#[test]
fn test_quantize_matrix_rejects_shape_overflow() {
match quantize_matrix(&[], usize::MAX, 2) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("overflows usize"),
"error must describe the geometry overflow; got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput for shape overflow, got: {e}"),
Ok(_) => panic!("expected Err for shape overflow, got Ok"),
}
}
#[test]
fn test_quantize_matrix_rejects_scale_underflow() {
let w = [f32::from_bits(1)];
match quantize_matrix(&w, 1, 1) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("non-positive scale 0"),
"error must describe the underflowed scale; got: {msg}"
);
assert!(msg.contains("row 0"), "wrong attribution: {msg}");
}
Err(e) => panic!("expected InvalidInput for scale underflow, got: {e}"),
Ok(_) => panic!("expected Err for scale underflow, got Ok"),
}
}
#[test]
fn test_quantize_gdn_weights_rejects_nonfinite_source() {
let mut w = valid_gdn_weights(3, 8);
w.in_proj_qkv[0] = f32::INFINITY;
match quantize_gdn_weights(&w, &matching_gdn_cfg()) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("non-finite"),
"error must describe the non-finite source; got: {msg}"
);
assert!(
msg.contains("in_proj_qkv"),
"error must name the offending tensor; got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput for non-finite source, got: {e}"),
Ok(_) => panic!("expected Err for non-finite source weight, got Ok"),
}
}
#[test]
fn test_quantize_common_weights_dense_returns_ok() {
use crate::model::qwen35::{CommonLayerWeights, DenseFfnWeights, FeedForwardWeights};
let hidden = 4usize;
let inter = 2usize;
let dense_common = CommonLayerWeights {
input_layernorm: vec![0.0f32; hidden],
post_attention_layernorm: vec![0.0f32; hidden],
ffn: FeedForwardWeights::Dense(DenseFfnWeights {
gate_proj: vec![0.0f32; inter * hidden],
up_proj: vec![0.0f32; inter * hidden],
down_proj: vec![0.0f32; hidden * inter],
}),
};
let cfg = Qwen35Config {
hidden_size: hidden,
intermediate_size: inter,
..Qwen35Config::qwen35_2b()
};
let result = quantize_common_weights(&dense_common, &cfg);
assert!(
result.is_ok(),
"expected Ok for dense layer, got: {result:?}"
);
}
fn valid_full_attention_weights(
cfg: &Qwen35Config,
head_dim: usize,
hidden: usize,
) -> FullAttentionLayerWeights {
let q_dim = cfg.num_attention_heads * head_dim;
let kv_dim = cfg.num_key_value_heads * head_dim;
let q_rows = 2 * q_dim;
FullAttentionLayerWeights {
q_proj: vec![0.0; q_rows * hidden],
k_proj: vec![0.0; kv_dim * hidden],
v_proj: vec![0.0; kv_dim * hidden],
o_proj: vec![0.0; hidden * q_dim],
q_norm: vec![1.0; head_dim],
k_norm: vec![1.0; head_dim],
}
}
#[test]
fn quantize_full_attn_weights_rejects_empty_q_norm() {
let cfg = Qwen35Config {
head_dim: 4,
hidden_size: 6,
..Qwen35Config::qwen35_2b()
};
let mut w = valid_full_attention_weights(&cfg, 4, 6);
w.q_norm = vec![];
match quantize_full_attn_weights(&w, &cfg) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(msg.contains("q_norm"), "error must name q_norm, got: {msg}");
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!("expected Err for empty q_norm, got Ok"),
}
}
#[test]
fn quantize_full_attn_weights_rejects_mismatched_k_norm() {
let cfg = Qwen35Config {
head_dim: 4,
hidden_size: 6,
..Qwen35Config::qwen35_2b()
};
let mut w = valid_full_attention_weights(&cfg, 4, 6);
w.k_norm = vec![1.0; 2]; match quantize_full_attn_weights(&w, &cfg) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(msg.contains("k_norm"), "error must name k_norm, got: {msg}");
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!("expected Err for mismatched k_norm, got Ok"),
}
}
#[test]
fn quantize_full_attn_weights_rejects_q_norm_disagreeing_with_cfg_head_dim() {
let cfg = Qwen35Config {
head_dim: 4,
hidden_size: 6,
..Qwen35Config::qwen35_2b()
};
let mut w = valid_full_attention_weights(&cfg, 4, 6);
w.q_norm = vec![1.0; 5];
w.k_norm = vec![1.0; 5];
match quantize_full_attn_weights(&w, &cfg) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(msg.contains("q_norm"), "error must name q_norm, got: {msg}");
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!(
"expected Err for q_norm/k_norm disagreeing with cfg.head_dim, got Ok \
(would panic in matmul_bt_q8 on the first full-attention token)"
),
}
}
#[test]
fn quantize_common_weights_rejects_empty_input_layernorm() {
use crate::model::qwen35::{CommonLayerWeights, DenseFfnWeights, FeedForwardWeights};
let cfg = Qwen35Config {
hidden_size: 4,
intermediate_size: 2,
..Qwen35Config::qwen35_2b()
};
let common = CommonLayerWeights {
input_layernorm: vec![],
post_attention_layernorm: vec![0.0f32; 4],
ffn: FeedForwardWeights::Dense(DenseFfnWeights {
gate_proj: vec![0.0f32; 8],
up_proj: vec![0.0f32; 8],
down_proj: vec![0.0f32; 8],
}),
};
match quantize_common_weights(&common, &cfg) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("input_layernorm") || msg.contains("hidden"),
"error must describe the empty hidden dimension, got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!("expected Err for empty input_layernorm, got Ok"),
}
}
#[test]
fn quantize_common_weights_rejects_zero_hidden_size() {
use crate::model::qwen35::{CommonLayerWeights, DenseFfnWeights, FeedForwardWeights};
let cfg = Qwen35Config {
hidden_size: 0,
intermediate_size: 0,
..Qwen35Config::qwen35_2b()
};
let common = CommonLayerWeights {
input_layernorm: vec![],
post_attention_layernorm: vec![],
ffn: FeedForwardWeights::Dense(DenseFfnWeights {
gate_proj: vec![],
up_proj: vec![],
down_proj: vec![],
}),
};
match quantize_common_weights(&common, &cfg) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("hidden_size") || msg.contains("hidden"),
"error must describe the zero hidden dimension, got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!(
"expected Err for hidden_size == 0, got Ok (would panic on x.len() / hidden \
in qwen35_rms_norm on the first forward pass)"
),
}
}
#[test]
fn quantize_common_weights_rejects_short_post_attention_layernorm() {
use crate::model::qwen35::{CommonLayerWeights, DenseFfnWeights, FeedForwardWeights};
let cfg = Qwen35Config {
hidden_size: 4,
intermediate_size: 2,
..Qwen35Config::qwen35_2b()
};
let common = CommonLayerWeights {
input_layernorm: vec![1.0f32; 4],
post_attention_layernorm: vec![1.0f32; 2], ffn: FeedForwardWeights::Dense(DenseFfnWeights {
gate_proj: vec![0.0f32; 8],
up_proj: vec![0.0f32; 8],
down_proj: vec![0.0f32; 8],
}),
};
match quantize_common_weights(&common, &cfg) {
Err(InferenceError::InvalidInput(msg)) => {
assert!(
msg.contains("post_attention_layernorm"),
"error must name post_attention_layernorm, got: {msg}"
);
}
Err(e) => panic!("expected InvalidInput, got: {e}"),
Ok(_) => panic!(
"expected Err for short post_attention_layernorm, got Ok (would silently \
leave trailing hidden values unnormalized)"
),
}
}
fn quantize_matrix_two_pass_reference(w: &[f32], rows: usize, cols: usize) -> Q8Matrix {
assert!(
w.iter().all(|v| v.is_finite()),
"reference requires finite input"
);
let mut data = Vec::with_capacity(w.len());
let mut scales = Vec::with_capacity(rows);
for row_idx in 0..rows {
let start = row_idx * cols;
let end = start + cols;
let row = &w[start..end];
let mut max_abs = 0.0f32;
for &v in row {
let abs_v = v.abs();
if abs_v > max_abs {
max_abs = abs_v;
}
}
let scale = if max_abs == 0.0 { 1.0 } else { max_abs / 127.0 };
scales.push(scale);
for &v in row {
let q = (v / scale).round().clamp(-128.0, 127.0) as i8;
data.push(q);
}
}
Q8Matrix {
data,
scales,
rows,
cols,
}
}
#[test]
fn test_fused_pass_matches_two_pass_reference_bit_exact() {
let cases: [(usize, usize, u32); 4] = [
(4, 4, 0x1111_1111),
(32, 64, 0x1234_5678),
(17, 33, 0xDEAD_BEEF),
(6, 4, 0xCAFE_BABE),
];
for (rows, cols, seed0) in cases {
let mut seed = seed0;
let mut w = Vec::with_capacity(rows * cols);
for _ in 0..rows * cols {
w.push(uniform_signed(&mut seed) * 0.37);
}
for v in &mut w[cols..2 * cols] {
*v = 0.0;
}
let fused = quantize_matrix(&w, rows, cols).expect("valid finite matrix must quantize");
let reference = quantize_matrix_two_pass_reference(&w, rows, cols);
assert_eq!(
fused.data, reference.data,
"quantized bytes diverged for {rows}x{cols}"
);
assert_eq!(
fused.scales, reference.scales,
"scales diverged for {rows}x{cols}"
);
}
}
}