use super::*;
use crate::ops::quantized;
const QUANT_IN: usize = 64;
fn dense_weight(out: usize) -> Array {
let mut data = Vec::with_capacity(out * QUANT_IN);
for o in 0..out {
for i in 0..QUANT_IN {
data.push(((o * 10 + i) as f32) * 0.001);
}
}
Array::from_slice::<f32>(&data, &(out, QUANT_IN)).unwrap()
}
fn input(n: usize) -> Array {
let mut data = Vec::with_capacity(n * QUANT_IN);
for ni in 0..n {
for i in 0..QUANT_IN {
data.push(((ni * 50 + i) as f32) * 0.01);
}
}
Array::from_slice::<f32>(&data, &(n, QUANT_IN)).unwrap()
}
#[allow(clippy::too_many_arguments)]
fn dequant_then_matmul(
x: &Array,
w_q: &Array,
scales: &Array,
q_biases: Option<&Array>,
bias: Option<&Array>,
group_size: i32,
bits: i32,
mode: &str,
) -> Vec<f32> {
let dense =
quantized::dequantize(w_q, scales, q_biases, group_size, bits, mode, None, None).unwrap();
let wt = dense.transpose().unwrap();
let mut y = x.matmul(&wt).unwrap();
if let Some(b) = bias {
y = y.add(b).unwrap();
}
y.to_vec::<f32>().unwrap()
}
fn assert_within_quant_error(reference: &[f32], got: &[f32]) {
assert_eq!(reference.len(), got.len(), "length mismatch");
let max_abs = reference.iter().fold(0.0f32, |m, v| m.max(v.abs()));
for (r, g) in reference.iter().zip(got.iter()) {
assert!(
(r - g).abs() <= 0.1 * max_abs + 1e-3,
"quantized Linear drift too large: reference={r} got={g}"
);
}
}
#[test]
fn linear_forward_no_bias() {
let weight = Array::from_slice::<f32>(
&[
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, ],
&(2usize, 3usize),
)
.unwrap();
let layer = Linear::new(weight, None);
let x = Array::from_slice::<f32>(&[10.0, 20.0, 30.0], &(1usize, 3usize)).unwrap();
let mut y = layer.forward(&x).unwrap();
assert_eq!(y.shape(), vec![1, 2]);
assert_eq!(y.to_vec::<f32>().unwrap(), vec![10.0, 30.0]);
}
#[test]
fn linear_forward_with_bias() {
let weight =
Array::from_slice::<f32>(&[1.0, 0.0, 0.0, 0.0, 0.0, 1.0], &(2usize, 3usize)).unwrap();
let bias = Array::from_slice::<f32>(&[100.0, 200.0], &(2usize,)).unwrap();
let layer = Linear::new(weight, Some(bias));
let x = Array::from_slice::<f32>(&[10.0, 20.0, 30.0], &(1usize, 3usize)).unwrap();
let mut y = layer.forward(&x).unwrap();
assert_eq!(y.to_vec::<f32>().unwrap(), vec![110.0, 230.0]);
}
#[test]
fn quantized_linear_forward_matches_dequant_matmul_no_bias() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
assert!(q_biases.is_some(), "affine produces per-group biases");
let layer = QuantizedLinear::from_parts(
w_q.try_clone().unwrap(),
scales.try_clone().unwrap(),
q_biases.as_ref().map(|b| b.try_clone().unwrap()),
None,
64,
4,
"affine",
)
.unwrap();
let x = input(2);
let mut got = layer.forward(&x).unwrap();
assert_eq!(got.shape(), vec![2, 8]);
let reference = dequant_then_matmul(&x, &w_q, &scales, q_biases.as_ref(), None, 64, 4, "affine");
assert_within_quant_error(&reference, &got.to_vec::<f32>().unwrap());
}
#[test]
fn quantized_linear_forward_matches_dequant_matmul_with_bias() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let bias =
Array::from_slice::<f32>(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0], &(8usize,)).unwrap();
let layer = QuantizedLinear::from_parts(
w_q.try_clone().unwrap(),
scales.try_clone().unwrap(),
q_biases.as_ref().map(|b| b.try_clone().unwrap()),
Some(bias.try_clone().unwrap()),
64,
4,
"affine",
)
.unwrap();
let x = input(2);
let mut got = layer.forward(&x).unwrap();
let reference = dequant_then_matmul(
&x,
&w_q,
&scales,
q_biases.as_ref(),
Some(&bias),
64,
4,
"affine",
);
assert_within_quant_error(&reference, &got.to_vec::<f32>().unwrap());
}
#[test]
fn quantized_linear_from_parts_rejects_non_u32_weight() {
let dense_w = dense_weight(8);
let (_w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let err = QuantizedLinear::from_parts(
dense_w, scales, q_biases, None, 64, 4, "affine",
)
.unwrap_err();
assert!(matches!(err, crate::Error::InvariantViolation(_)));
}
#[test]
fn quantized_linear_from_parts_rejects_rank3_weight() {
let dense3 = Array::from_slice::<f32>(
&vec![0.001f32; 2 * 4 * QUANT_IN],
&(2usize, 4usize, QUANT_IN),
)
.unwrap();
let (w_q3, scales3, qb3) = quantized::quantize(&dense3, 64, 4, "affine", None).unwrap();
let err = QuantizedLinear::from_parts(w_q3, scales3, qb3, None, 64, 4, "affine").unwrap_err();
assert!(matches!(err, crate::Error::RankMismatch(_)));
}
#[test]
fn quantized_linear_from_parts_rejects_affine_without_biases() {
let dense_w = dense_weight(8);
let (w_q, scales, _q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let err = QuantizedLinear::from_parts(w_q, scales, None, None, 64, 4, "affine").unwrap_err();
assert!(matches!(err, crate::Error::InvariantViolation(_)));
}
#[test]
fn quantized_linear_from_parts_rejects_unknown_mode() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let err = QuantizedLinear::from_parts(w_q, scales, q_biases, None, 64, 4, "garbage").unwrap_err();
assert!(matches!(err, crate::Error::UnknownEnumValue(_)));
}
#[test]
fn quantized_linear_from_parts_rejects_zero_group_size() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let err = QuantizedLinear::from_parts(w_q, scales, q_biases, None, 0, 4, "affine").unwrap_err();
assert!(matches!(err, crate::Error::OutOfRange(_)));
}
#[test]
fn quantized_linear_from_parts_rejects_higher_rank_bias() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let bad_bias = Array::from_slice::<f32>(&[1.0; 16], &(8usize, 2usize)).unwrap();
let err = QuantizedLinear::from_parts(w_q, scales, q_biases, Some(bad_bias), 64, 4, "affine")
.unwrap_err();
assert!(matches!(err, crate::Error::RankMismatch(_)));
}
#[test]
fn quantized_linear_from_parts_rejects_length_one_bias() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let bad_bias = Array::from_slice::<f32>(&[1.0], &(1usize,)).unwrap();
let err = QuantizedLinear::from_parts(w_q, scales, q_biases, Some(bad_bias), 64, 4, "affine")
.unwrap_err();
assert!(matches!(err, crate::Error::ShapePairMismatch(_)));
}
#[test]
fn quantized_linear_from_parts_accepts_exact_length_bias() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let bias = Array::from_slice::<f32>(&[1.0; 8], &(8usize,)).unwrap();
assert!(QuantizedLinear::from_parts(w_q, scales, q_biases, Some(bias), 64, 4, "affine").is_ok());
}
#[test]
fn quantized_linear_from_parts_rejects_mismatched_scales_leading_dim() {
let dense_w = dense_weight(8);
let (w_q, _scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let other_dense = dense_weight(4);
let (_w2, bad_scales, _qb2) = quantized::quantize(&other_dense, 64, 4, "affine", None).unwrap();
let err =
QuantizedLinear::from_parts(w_q, bad_scales, q_biases, None, 64, 4, "affine").unwrap_err();
assert!(matches!(err, crate::Error::ShapePairMismatch(_)));
}
#[test]
fn quantized_linear_from_parts_rejects_wrong_scales_trailing_dim() {
let dense_w = dense_weight(8);
let (w_q, _scales64, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let (_w32, wrong_scales, _qb32) = quantized::quantize(&dense_w, 32, 4, "affine", None).unwrap();
assert_eq!(
wrong_scales.shape(),
vec![8, 2],
"fixture: group_size-32 scales"
);
let err =
QuantizedLinear::from_parts(w_q, wrong_scales, q_biases, None, 64, 4, "affine").unwrap_err();
assert!(
matches!(err, crate::Error::ShapePairMismatch(_)),
"expected ShapePairMismatch for a wrong scales trailing dim, got {err:?}"
);
}
#[test]
fn quantized_linear_from_parts_accepts_affine_integer_scales_floating_biases() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let int_scales = scales.astype(Dtype::I32).unwrap();
let got = QuantizedLinear::from_parts(w_q, int_scales, q_biases, None, 64, 4, "affine");
assert!(
got.is_ok(),
"expected integer-scales + floating-biases affine triple to construct, got {got:?}"
);
}
#[test]
fn quantized_linear_from_parts_accepts_affine_floating_scales_integer_biases() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let int_biases = q_biases.unwrap().astype(Dtype::I32).unwrap();
let got = QuantizedLinear::from_parts(w_q, scales, Some(int_biases), None, 64, 4, "affine");
assert!(
got.is_ok(),
"expected floating-scales + integer-biases affine triple to construct, got {got:?}"
);
}
#[test]
fn quantized_linear_from_parts_rejects_fp_mode_non_uint8_scales() {
let dense_w = dense_weight(8);
let (w_q, scales, _q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
assert!(
scales.dtype().unwrap() != Dtype::U8,
"fixture: affine scales are floating, not uint8"
);
let err = QuantizedLinear::from_parts(w_q, scales, None, None, 64, 4, "mxfp4").unwrap_err();
assert!(
matches!(err, crate::Error::UnsupportedDtype(_)),
"expected UnsupportedDtype for non-uint8 fp-mode scales, got {err:?}"
);
}
#[test]
fn maybe_quantized_picks_dense_when_no_scales() {
let mut weights: HashMap<String, Array> = HashMap::new();
let weight = Array::from_slice::<f32>(&[1.0, 0.0, 0.0, 1.0], &(2usize, 2usize)).unwrap();
let bias = Array::from_slice::<f32>(&[7.0, 9.0], &(2usize,)).unwrap();
weights.insert("blk.q.weight".to_string(), weight);
weights.insert("blk.q.bias".to_string(), bias);
let layer = MaybeQuantizedLinear::from_weights(&mut weights, "blk.q", None).unwrap();
assert!(!layer.is_quantized());
assert!(matches!(layer, MaybeQuantizedLinear::Dense(_)));
assert!(weights.is_empty());
let x = Array::from_slice::<f32>(&[3.0, 5.0], &(1usize, 2usize)).unwrap();
let mut y = layer.forward(&x).unwrap();
assert_eq!(y.to_vec::<f32>().unwrap(), vec![10.0, 14.0]);
}
#[test]
fn maybe_quantized_picks_quantized_when_scales_present() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let q_biases = q_biases.expect("affine biases");
let mut weights: HashMap<String, Array> = HashMap::new();
weights.insert("blk.q.weight".to_string(), w_q.try_clone().unwrap());
weights.insert("blk.q.scales".to_string(), scales.try_clone().unwrap());
weights.insert("blk.q.biases".to_string(), q_biases.try_clone().unwrap());
let layer =
MaybeQuantizedLinear::from_weights(&mut weights, "blk.q", Some((64, 4, "affine"))).unwrap();
assert!(layer.is_quantized());
assert!(matches!(layer, MaybeQuantizedLinear::Quantized(_)));
assert!(weights.is_empty());
let x = input(2);
let mut got = layer.forward(&x).unwrap();
let reference = dequant_then_matmul(&x, &w_q, &scales, Some(&q_biases), None, 64, 4, "affine");
assert_within_quant_error(&reference, &got.to_vec::<f32>().unwrap());
}
#[test]
fn maybe_quantized_carries_dense_bias_on_quantized_layer() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let q_biases = q_biases.expect("affine biases");
let dense_bias = Array::from_slice::<f32>(&[1.0; 8], &(8usize,)).unwrap();
let mut weights: HashMap<String, Array> = HashMap::new();
weights.insert("blk.q.weight".to_string(), w_q);
weights.insert("blk.q.scales".to_string(), scales);
weights.insert("blk.q.biases".to_string(), q_biases);
weights.insert("blk.q.bias".to_string(), dense_bias);
let layer =
MaybeQuantizedLinear::from_weights(&mut weights, "blk.q", Some((64, 4, "affine"))).unwrap();
match layer {
MaybeQuantizedLinear::Quantized(q) => {
assert!(
q.bias().is_some(),
"dense `.bias` must land in the bias slot"
);
assert!(q.quant_biases().is_some(), "affine `.biases` present");
}
_ => panic!("expected quantized variant"),
}
assert!(weights.is_empty());
}
#[test]
fn maybe_quantized_scales_present_but_no_config_errors() {
let dense_w = dense_weight(8);
let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let mut weights: HashMap<String, Array> = HashMap::new();
weights.insert("blk.q.weight".to_string(), w_q);
weights.insert("blk.q.scales".to_string(), scales);
weights.insert("blk.q.biases".to_string(), q_biases.unwrap());
let err = MaybeQuantizedLinear::from_weights(&mut weights, "blk.q", None).unwrap_err();
assert!(matches!(err, crate::Error::InvariantViolation(_)));
}
#[test]
fn maybe_quantized_dense_missing_weight_errors() {
let mut weights: HashMap<String, Array> = HashMap::new();
let err = MaybeQuantizedLinear::from_weights(&mut weights, "blk.q", None).unwrap_err();
assert!(matches!(err, crate::Error::MissingKey(_)));
}
fn embedding_triple(num_embeddings: usize) -> (Array, Array, Array) {
let mut data = Vec::with_capacity(num_embeddings * QUANT_IN);
for v in 0..num_embeddings {
for s in 0..QUANT_IN {
data.push(((v * 5 + s) as f32) * 0.001);
}
}
let dense = Array::from_slice::<f32>(&data, &(num_embeddings, QUANT_IN)).unwrap();
let (w_q, scales, biases) = quantized::quantize(&dense, 64, 8, "affine", None).unwrap();
(
w_q,
scales,
biases.expect("affine produces per-group biases"),
)
}
#[test]
fn maybe_quantized_embedding_dense_gather_and_table() {
let table = dense_weight(6); let emb = MaybeQuantizedEmbedding::dense(table.try_clone().unwrap());
assert!(!emb.is_quantized());
let ids = Array::from_slice::<i32>(&[0, 2, 5], &(3usize,)).unwrap();
let rows = emb.gather(&ids).unwrap();
assert_eq!(rows.shape(), vec![3, QUANT_IN]);
let full = emb.dense_table(None).unwrap();
assert_eq!(full.shape(), vec![6, QUANT_IN]);
}
#[test]
fn maybe_quantized_embedding_quantized_gather_and_table_finite() {
let (w_q, scales, biases) = embedding_triple(8);
let emb =
MaybeQuantizedEmbedding::from_parts(w_q, scales, Some(biases), 64, 8, "affine").unwrap();
assert!(emb.is_quantized());
let ids = Array::from_slice::<i32>(&[0, 3, 7], &(3usize,)).unwrap();
let mut rows = emb.gather(&ids).unwrap();
assert_eq!(rows.shape(), vec![3, QUANT_IN]);
rows.eval().unwrap();
for v in rows.to_vec::<f32>().unwrap() {
assert!(v.is_finite(), "dequantized embedding row non-finite: {v}");
}
let mut full = emb.dense_table(Some(crate::Dtype::F32)).unwrap();
assert_eq!(full.shape(), vec![8, QUANT_IN]);
full.eval().unwrap();
for v in full.to_vec::<f32>().unwrap() {
assert!(v.is_finite(), "dequantized embedding table non-finite: {v}");
}
}
#[test]
fn maybe_quantized_embedding_from_parts_rejects_affine_without_biases() {
let (w_q, scales, _biases) = embedding_triple(8);
let err = MaybeQuantizedEmbedding::from_parts(w_q, scales, None, 64, 8, "affine").unwrap_err();
assert!(matches!(err, crate::Error::InvariantViolation(_)));
}
#[test]
fn maybe_quantized_embedding_from_weights_quantized_path() {
let (w_q, scales, biases) = embedding_triple(8);
let mut weights: HashMap<String, Array> = HashMap::new();
weights.insert("emb.weight".to_string(), w_q);
weights.insert("emb.scales".to_string(), scales);
weights.insert("emb.biases".to_string(), biases);
let emb =
MaybeQuantizedEmbedding::from_weights(&mut weights, "emb", Some((64, 8, "affine"))).unwrap();
assert!(emb.is_quantized());
assert!(weights.is_empty());
}
#[test]
fn maybe_quantized_embedding_from_weights_dense_path() {
let table = dense_weight(6);
let mut weights: HashMap<String, Array> = HashMap::new();
weights.insert("emb.weight".to_string(), table);
let emb =
MaybeQuantizedEmbedding::from_weights(&mut weights, "emb", Some((64, 8, "affine"))).unwrap();
assert!(!emb.is_quantized());
assert!(weights.is_empty());
}
#[test]
fn maybe_quantized_embedding_logical_shape_both_arms() {
let dense = MaybeQuantizedEmbedding::dense(dense_weight(6)); assert_eq!(dense.logical_shape().unwrap(), (6, QUANT_IN as i32));
let (w_q, scales, biases) = embedding_triple(8); let quant =
MaybeQuantizedEmbedding::from_parts(w_q, scales, Some(biases), 64, 8, "affine").unwrap();
assert_eq!(quant.logical_shape().unwrap(), (8, QUANT_IN as i32));
}
#[test]
fn maybe_quantized_linear_logical_shape_both_arms() {
let dense = MaybeQuantizedLinear::Dense(Linear::new(dense_weight(5), None)); assert_eq!(dense.logical_shape().unwrap(), (5, QUANT_IN as i32));
let dense_w = dense_weight(8); let (w_q, scales, q_biases) = quantized::quantize(&dense_w, 64, 4, "affine", None).unwrap();
let q = QuantizedLinear::from_parts(w_q, scales, q_biases, None, 64, 4, "affine").unwrap();
let quant = MaybeQuantizedLinear::Quantized(q);
assert_eq!(quant.logical_shape().unwrap(), (8, QUANT_IN as i32));
}
#[test]
fn maybe_quantized_embedding_scales_present_but_no_config_errors() {
let (w_q, scales, biases) = embedding_triple(8);
let mut weights: HashMap<String, Array> = HashMap::new();
weights.insert("emb.weight".to_string(), w_q);
weights.insert("emb.scales".to_string(), scales);
weights.insert("emb.biases".to_string(), biases);
let err = MaybeQuantizedEmbedding::from_weights(&mut weights, "emb", None).unwrap_err();
assert!(matches!(err, crate::Error::InvariantViolation(_)));
}