Skip to main content

ferrotorch_core/
quantize.rs

1//! Post-training quantization (PTQ) for ferrotorch tensors.
2//!
3//! Provides symmetric and asymmetric quantization to INT8, INT4, and UINT8,
4//! with per-tensor or per-channel granularity. Designed for inference-time
5//! model compression — quantize once after training, then run forward passes
6//! with reduced memory and (on supported hardware) faster matmul.
7//!
8//! ## REQ status (per `.design/ferrotorch-core/quantize.md`)
9//!
10//! | REQ | Status | Evidence |
11//! |---|---|---|
12//! | REQ-1 | SHIPPED | impl `QuantScheme`, `QuantDtype`; non-test consumer re-exported at `lib.rs:179-181`. |
13//! | REQ-2 | SHIPPED | impl `QuantizedTensor`; non-test consumer re-exported at `lib.rs:179-181`; threaded through `quantize_named_tensors`. |
14//! | REQ-3 | SHIPPED | impl `quantize`; non-test consumer `quantize_named_tensors`, `FakeQuantize::forward` chain via `grad_fns::quantize_grad`. |
15//! | REQ-4 | SHIPPED | impl `dequantize`; non-test consumer `quantized_matmul`, `FakeQuantize::forward`. |
16//! | REQ-5 | SHIPPED | impl `quantized_matmul`; non-test consumer re-exported at `lib.rs:179-181`. |
17//! | REQ-6 | SHIPPED | impl `QParams`; non-test consumer threaded through every observer + `QatModel::step`. |
18//! | REQ-7 | SHIPPED | impl `trait Observer` + `MinMaxObserver` + `PerChannelMinMaxObserver` + `HistogramObserver`; non-test consumer `QatLayer`. |
19//! | REQ-8 | SHIPPED | impl `FakeQuantize`; non-test consumer `Tensor::fake_quantize_per_tensor_affine_t` at `methods.rs:596` via `grad_fns::quantize_grad`. |
20//! | REQ-9 | SHIPPED | impl `QatLayer`, `QatModel`, `prepare_qat`; non-test consumer pub-API QAT entry point at `lib.rs:179-181`. |
21//! | REQ-10 | SHIPPED | impl `quantize_named_tensors`; non-test consumer quantized-state-dict save flow. |
22
23use std::collections::HashMap;
24
25use crate::dtype::Float;
26use crate::error::{FerrotorchError, FerrotorchResult};
27use crate::storage::TensorStorage;
28use crate::tensor::Tensor;
29
30// ---------------------------------------------------------------------------
31// Enums
32// ---------------------------------------------------------------------------
33
34/// Granularity of quantization parameters (scale / zero_point).
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
36pub enum QuantScheme {
37    /// One scale and zero_point for the entire tensor.
38    PerTensor,
39    /// One scale and zero_point per slice along the given axis.
40    PerChannel(usize),
41}
42
43/// Target integer dtype for quantized storage.
44#[derive(Debug, Clone, Copy, PartialEq, Eq)]
45pub enum QuantDtype {
46    /// Signed 8-bit: [-128, 127].
47    Int8,
48    /// Signed 4-bit: [-8, 7].  Stored packed in `i8` values.
49    Int4,
50    /// Unsigned 8-bit: [0, 255].
51    Uint8,
52}
53
54impl QuantDtype {
55    /// Minimum representable value.
56    #[inline]
57    fn qmin(self) -> i32 {
58        match self {
59            QuantDtype::Int8 => -128,
60            QuantDtype::Int4 => -8,
61            QuantDtype::Uint8 => 0,
62        }
63    }
64
65    /// Maximum representable value.
66    #[inline]
67    fn qmax(self) -> i32 {
68        match self {
69            QuantDtype::Int8 => 127,
70            QuantDtype::Int4 => 7,
71            QuantDtype::Uint8 => 255,
72        }
73    }
74}
75
76// ---------------------------------------------------------------------------
77// QuantizedTensor
78// ---------------------------------------------------------------------------
79
80/// A tensor stored in quantized (integer) representation.
81///
82/// The real value is recovered by `x = (q - zero_point) * scale`.
83///
84/// `scale` and `zero_point` are vectors whose length equals:
85/// * 1 for `PerTensor`
86/// * `shape[axis]` for `PerChannel(axis)`
87#[derive(Debug, Clone)]
88pub struct QuantizedTensor {
89    /// Quantized values stored as `i8` regardless of logical dtype.
90    /// For `Uint8`, the stored `i8` is reinterpreted as `u8` via
91    /// wrapping cast; for `Int4` only the low 4 bits are significant.
92    data: Vec<i8>,
93    /// Per-tensor or per-channel scales.
94    scale: Vec<f32>,
95    /// Per-tensor or per-channel zero points (in quantized domain).
96    zero_point: Vec<i32>,
97    /// Original tensor shape.
98    shape: Vec<usize>,
99    /// Quantization granularity.
100    scheme: QuantScheme,
101    /// Target quantized dtype.
102    dtype: QuantDtype,
103}
104
105impl QuantizedTensor {
106    /// Number of elements.
107    #[inline]
108    pub fn numel(&self) -> usize {
109        self.shape.iter().product()
110    }
111
112    /// Borrow the shape.
113    #[inline]
114    pub fn shape(&self) -> &[usize] {
115        &self.shape
116    }
117
118    /// Borrow the quantized data.
119    #[inline]
120    pub fn data(&self) -> &[i8] {
121        &self.data
122    }
123
124    /// Borrow the scale vector.
125    #[inline]
126    pub fn scale(&self) -> &[f32] {
127        &self.scale
128    }
129
130    /// Borrow the zero-point vector.
131    #[inline]
132    pub fn zero_point(&self) -> &[i32] {
133        &self.zero_point
134    }
135
136    /// The quantization scheme used.
137    #[inline]
138    pub fn scheme(&self) -> QuantScheme {
139        self.scheme
140    }
141
142    /// The quantized dtype.
143    #[inline]
144    pub fn qdtype(&self) -> QuantDtype {
145        self.dtype
146    }
147}
148
149// ---------------------------------------------------------------------------
150// Helpers
151// ---------------------------------------------------------------------------
152
153/// Compute scale and zero_point for a given (min, max) range and target dtype.
154///
155/// Uses the standard asymmetric affine quantization formula:
156///   scale = (max - min) / (qmax - qmin)
157///   zero_point = round(qmin - min / scale)
158///
159/// The range is always expanded to include zero so that `0.0` maps exactly
160/// to an integer quantized value (important for zero-padding and ReLU outputs).
161/// When min == max the range would collapse to zero, so this expansion also
162/// prevents division-by-zero.
163fn compute_scale_zp(min_val: f32, max_val: f32, dtype: QuantDtype) -> (f32, i32) {
164    let qmin = dtype.qmin();
165    let qmax = dtype.qmax();
166
167    // Ensure the range includes zero (standard PyTorch behaviour).
168    let min_val = min_val.min(0.0);
169    let max_val = max_val.max(0.0);
170
171    // After including zero the range is at least max(|min|, |max|) > 0,
172    // but guard against the degenerate all-zeros case.
173    let range = (max_val - min_val).max(f32::EPSILON);
174    let scale = range / (qmax - qmin) as f32;
175
176    // zero_point is intentionally NOT clamped to [qmin, qmax]. It is stored
177    // as i32 and may lie outside the quantized integer range. This is correct
178    // for asymmetric affine quantization — clamping the zero_point distorts
179    // the mapping when the float range doesn't straddle zero.
180    let zp = (qmin as f32 - min_val / scale).round() as i32;
181
182    (scale, zp)
183}
184
185/// Clamp and round a float to the quantized integer range.
186///
187/// Returns the result as `i8`. For `Uint8` the caller passes `qmin=0`,
188/// `qmax=255`; the clamped i32 is cast to `u8` first then transmuted to `i8`
189/// so that values 128..=255 are preserved through the bit pattern.
190#[inline]
191fn quantize_val(x: f32, scale: f32, zp: i32, qmin: i32, qmax: i32, is_unsigned: bool) -> i8 {
192    let q = (x / scale + zp as f32).round() as i32;
193    let clamped = q.clamp(qmin, qmax);
194    if is_unsigned {
195        (clamped as u8) as i8
196    } else {
197        clamped as i8
198    }
199}
200
201/// Recover the i32 quantized value from the stored `i8`, accounting for
202/// unsigned dtypes where the bit pattern represents a `u8`.
203#[inline]
204fn stored_to_i32(val: i8, is_unsigned: bool) -> i32 {
205    if is_unsigned {
206        (val as u8) as i32
207    } else {
208        val as i32
209    }
210}
211
212/// Map a linear flat index to per-channel parameters.
213///
214/// For a tensor of shape `[d0, d1, ..., dn]` with channel axis `axis`,
215/// returns the channel index for the element at `flat_index`.
216#[inline]
217fn channel_index(flat_index: usize, shape: &[usize], axis: usize) -> usize {
218    // stride of the channel axis = product of dims after axis.
219    let stride: usize = shape[axis + 1..].iter().product();
220    (flat_index / stride) % shape[axis]
221}
222
223// ---------------------------------------------------------------------------
224// Quantize
225// ---------------------------------------------------------------------------
226
227/// Quantize a floating-point tensor.
228///
229/// # Per-tensor
230///
231/// Computes a single (scale, zero_point) pair from the global min/max.
232///
233/// # Per-channel
234///
235/// Computes one (scale, zero_point) per slice along the given axis. This is
236/// common for weight tensors where each output channel has its own range.
237pub fn quantize<T: Float>(
238    tensor: &Tensor<T>,
239    scheme: QuantScheme,
240    dtype: QuantDtype,
241) -> FerrotorchResult<QuantizedTensor> {
242    let data = tensor.data()?;
243    let shape = tensor.shape().to_vec();
244    let numel = tensor.numel();
245    let qmin = dtype.qmin();
246    let qmax = dtype.qmax();
247
248    let is_unsigned = dtype == QuantDtype::Uint8;
249
250    match scheme {
251        QuantScheme::PerTensor => {
252            // Global min/max.
253            let mut min_val = f32::INFINITY;
254            let mut max_val = f32::NEG_INFINITY;
255            for &v in data {
256                let f = v.to_f32().unwrap();
257                if f < min_val {
258                    min_val = f;
259                }
260                if f > max_val {
261                    max_val = f;
262                }
263            }
264
265            let (scale, zp) = compute_scale_zp(min_val, max_val, dtype);
266
267            let qdata: Vec<i8> = data
268                .iter()
269                .map(|&v| quantize_val(v.to_f32().unwrap(), scale, zp, qmin, qmax, is_unsigned))
270                .collect();
271
272            Ok(QuantizedTensor {
273                data: qdata,
274                scale: vec![scale],
275                zero_point: vec![zp],
276                shape,
277                scheme,
278                dtype,
279            })
280        }
281
282        QuantScheme::PerChannel(axis) => {
283            if axis >= shape.len() {
284                return Err(FerrotorchError::InvalidArgument {
285                    message: format!(
286                        "PerChannel axis {axis} out of range for {}-d tensor",
287                        shape.len()
288                    ),
289                });
290            }
291
292            let num_channels = shape[axis];
293            let mut mins = vec![f32::INFINITY; num_channels];
294            let mut maxs = vec![f32::NEG_INFINITY; num_channels];
295
296            for (i, &v) in data.iter().enumerate() {
297                let ch = channel_index(i, &shape, axis);
298                let f = v.to_f32().unwrap();
299                if f < mins[ch] {
300                    mins[ch] = f;
301                }
302                if f > maxs[ch] {
303                    maxs[ch] = f;
304                }
305            }
306
307            let params: Vec<(f32, i32)> = mins
308                .iter()
309                .zip(maxs.iter())
310                .map(|(&mn, &mx)| compute_scale_zp(mn, mx, dtype))
311                .collect();
312
313            let scales: Vec<f32> = params.iter().map(|&(s, _)| s).collect();
314            let zps: Vec<i32> = params.iter().map(|&(_, z)| z).collect();
315
316            let mut qdata = Vec::with_capacity(numel);
317            for (i, &v) in data.iter().enumerate() {
318                let ch = channel_index(i, &shape, axis);
319                qdata.push(quantize_val(
320                    v.to_f32().unwrap(),
321                    scales[ch],
322                    zps[ch],
323                    qmin,
324                    qmax,
325                    is_unsigned,
326                ));
327            }
328
329            Ok(QuantizedTensor {
330                data: qdata,
331                scale: scales,
332                zero_point: zps,
333                shape,
334                scheme,
335                dtype,
336            })
337        }
338    }
339}
340
341// ---------------------------------------------------------------------------
342// Dequantize
343// ---------------------------------------------------------------------------
344
345/// Dequantize back to a floating-point tensor.
346///
347/// Applies the inverse mapping: `x = (q - zero_point) * scale`.
348pub fn dequantize<T: Float>(qtensor: &QuantizedTensor) -> FerrotorchResult<Tensor<T>> {
349    let numel = qtensor.numel();
350    let mut result = Vec::with_capacity(numel);
351    let is_unsigned = qtensor.dtype == QuantDtype::Uint8;
352
353    match qtensor.scheme {
354        QuantScheme::PerTensor => {
355            let scale = qtensor.scale[0];
356            let zp = qtensor.zero_point[0];
357            for &q in &qtensor.data {
358                let val = (stored_to_i32(q, is_unsigned) - zp) as f32 * scale;
359                result.push(T::from(val).unwrap());
360            }
361        }
362        QuantScheme::PerChannel(axis) => {
363            for (i, &q) in qtensor.data.iter().enumerate() {
364                let ch = channel_index(i, &qtensor.shape, axis);
365                let val = (stored_to_i32(q, is_unsigned) - qtensor.zero_point[ch]) as f32
366                    * qtensor.scale[ch];
367                result.push(T::from(val).unwrap());
368            }
369        }
370    }
371
372    Tensor::from_storage(TensorStorage::cpu(result), qtensor.shape.clone(), false)
373}
374
375// ---------------------------------------------------------------------------
376// Quantized matmul
377// ---------------------------------------------------------------------------
378
379/// Multiply two quantized 2-D matrices and return a quantized result.
380///
381/// Strategy: accumulate in `i32` to avoid overflow, then rescale to the output
382/// quantized domain. This avoids a full dequantize-matmul-requantize round-trip
383/// while remaining numerically correct for INT8.
384///
385/// Both inputs must be 2-D, with compatible inner dimensions (standard matmul
386/// rules: `[M, K] x [K, N] -> [M, N]`).
387pub fn quantized_matmul(
388    a: &QuantizedTensor,
389    b: &QuantizedTensor,
390) -> FerrotorchResult<QuantizedTensor> {
391    // Validate shapes.
392    if a.shape.len() != 2 || b.shape.len() != 2 {
393        return Err(FerrotorchError::InvalidArgument {
394            message: format!(
395                "quantized_matmul requires 2-D tensors, got shapes {:?} and {:?}",
396                a.shape, b.shape
397            ),
398        });
399    }
400
401    let m = a.shape[0];
402    let k = a.shape[1];
403    let k2 = b.shape[0];
404    let n = b.shape[1];
405
406    if k != k2 {
407        return Err(FerrotorchError::ShapeMismatch {
408            message: format!(
409                "quantized_matmul inner dimensions mismatch: [{m}, {k}] x [{k2}, {n}]"
410            ),
411        });
412    }
413
414    // Both inputs must be PerTensor for the fast path.
415    if a.scale.len() != 1 || b.scale.len() != 1 {
416        return Err(FerrotorchError::InvalidArgument {
417            message: "quantized_matmul currently requires PerTensor-quantized inputs".into(),
418        });
419    }
420
421    let a_scale = a.scale[0];
422    let a_zp = a.zero_point[0];
423    let b_scale = b.scale[0];
424    let b_zp = b.zero_point[0];
425
426    let a_unsigned = a.dtype == QuantDtype::Uint8;
427    let b_unsigned = b.dtype == QuantDtype::Uint8;
428
429    // Accumulate in i32.
430    let mut acc = vec![0i32; m * n];
431    for i in 0..m {
432        for j in 0..n {
433            let mut sum = 0i32;
434            for p in 0..k {
435                let qa = stored_to_i32(a.data[i * k + p], a_unsigned) - a_zp;
436                let qb = stored_to_i32(b.data[p * n + j], b_unsigned) - b_zp;
437                sum += qa * qb;
438            }
439            acc[i * n + j] = sum;
440        }
441    }
442
443    // The real-valued result element is: acc[i,j] * a_scale * b_scale.
444    // Requantize: pick INT8 output with its own scale/zp.
445    let combined_scale = a_scale * b_scale;
446
447    // Find the real-valued min/max of the output.
448    let mut out_min = f32::INFINITY;
449    let mut out_max = f32::NEG_INFINITY;
450    for &a_val in &acc {
451        let real = a_val as f32 * combined_scale;
452        if real < out_min {
453            out_min = real;
454        }
455        if real > out_max {
456            out_max = real;
457        }
458    }
459
460    let out_dtype = QuantDtype::Int8;
461    let (out_scale, out_zp) = compute_scale_zp(out_min, out_max, out_dtype);
462    let qmin = out_dtype.qmin();
463    let qmax = out_dtype.qmax();
464
465    let qdata: Vec<i8> = acc
466        .iter()
467        .map(|&a_val| {
468            let real = a_val as f32 * combined_scale;
469            quantize_val(real, out_scale, out_zp, qmin, qmax, false)
470        })
471        .collect();
472
473    Ok(QuantizedTensor {
474        data: qdata,
475        scale: vec![out_scale],
476        zero_point: vec![out_zp],
477        shape: vec![m, n],
478        scheme: QuantScheme::PerTensor,
479        dtype: out_dtype,
480    })
481}
482
483// ---------------------------------------------------------------------------
484// Module-level quantization utility
485// ---------------------------------------------------------------------------
486
487/// Quantize every weight tensor in a module, returning a name -> QuantizedTensor
488/// map suitable for serialization or quantized inference.
489///
490/// This accepts any type implementing the `Module` trait from `ferrotorch-nn`.
491/// Because `ferrotorch-core` does not depend on `ferrotorch-nn`, we accept a
492/// generic iterator of named tensors instead.
493pub fn quantize_named_tensors<T: Float>(
494    named_tensors: impl IntoIterator<Item = (String, Tensor<T>)>,
495    scheme: QuantScheme,
496    dtype: QuantDtype,
497) -> FerrotorchResult<HashMap<String, QuantizedTensor>> {
498    let mut result = HashMap::new();
499    for (name, tensor) in named_tensors {
500        let qtensor = quantize(&tensor, scheme, dtype)?;
501        result.insert(name, qtensor);
502    }
503    Ok(result)
504}
505
506// ===========================================================================
507// QParams — quantization parameters
508// ===========================================================================
509
510/// Computed quantization parameters (scale and zero_point).
511#[derive(Debug, Clone)]
512pub struct QParams {
513    /// Per-tensor or per-channel scales.
514    pub scale: Vec<f32>,
515    /// Per-tensor or per-channel zero points.
516    pub zero_point: Vec<i32>,
517}
518
519impl QParams {
520    /// Compute symmetric quantization parameters.
521    ///
522    /// For symmetric quantization the range is `[-max_abs, max_abs]` and:
523    /// - INT8: `zero_point = 0`, `scale = max_abs / 127`
524    /// - INT4: `zero_point = 0`, `scale = max_abs / 7`
525    /// - UINT8: `zero_point = 128`, `scale = max_abs / 128`
526    pub fn symmetric(max_abs: f32, dtype: QuantDtype) -> Self {
527        let max_abs = max_abs.max(f32::EPSILON);
528        match dtype {
529            QuantDtype::Int8 => QParams {
530                scale: vec![max_abs / 127.0],
531                zero_point: vec![0],
532            },
533            QuantDtype::Int4 => QParams {
534                scale: vec![max_abs / 7.0],
535                zero_point: vec![0],
536            },
537            QuantDtype::Uint8 => QParams {
538                scale: vec![max_abs / 128.0],
539                zero_point: vec![128],
540            },
541        }
542    }
543
544    /// Compute asymmetric quantization parameters from observed min/max.
545    pub fn asymmetric(min_val: f32, max_val: f32, dtype: QuantDtype) -> Self {
546        let (scale, zp) = compute_scale_zp(min_val, max_val, dtype);
547        QParams {
548            scale: vec![scale],
549            zero_point: vec![zp],
550        }
551    }
552}
553
554// ===========================================================================
555// Observers — collect statistics for quantization calibration
556// ===========================================================================
557
558/// Trait for quantization observers that collect data statistics.
559pub trait Observer {
560    /// Update the observer with a batch of floating-point values.
561    fn observe(&mut self, data: &[f32]);
562    /// Calculate quantization parameters from collected statistics.
563    fn calculate_qparams(&self, dtype: QuantDtype) -> QParams;
564    /// Reset the observer state.
565    fn reset(&mut self);
566}
567
568// ---------------------------------------------------------------------------
569// MinMaxObserver
570// ---------------------------------------------------------------------------
571
572/// Tracks the running min/max of observed values.
573///
574/// Filters out NaN and Inf values before updating min/max.
575#[derive(Debug, Clone)]
576pub struct MinMaxObserver {
577    min_val: f32,
578    max_val: f32,
579}
580
581impl MinMaxObserver {
582    pub fn new() -> Self {
583        Self {
584            min_val: f32::INFINITY,
585            max_val: f32::NEG_INFINITY,
586        }
587    }
588}
589
590impl Default for MinMaxObserver {
591    fn default() -> Self {
592        Self::new()
593    }
594}
595
596impl Observer for MinMaxObserver {
597    fn observe(&mut self, data: &[f32]) {
598        for &x in data {
599            if !x.is_finite() {
600                continue;
601            }
602            if x < self.min_val {
603                self.min_val = x;
604            }
605            if x > self.max_val {
606                self.max_val = x;
607            }
608        }
609    }
610
611    fn calculate_qparams(&self, dtype: QuantDtype) -> QParams {
612        QParams::asymmetric(self.min_val, self.max_val, dtype)
613    }
614
615    fn reset(&mut self) {
616        self.min_val = f32::INFINITY;
617        self.max_val = f32::NEG_INFINITY;
618    }
619}
620
621// ---------------------------------------------------------------------------
622// PerChannelMinMaxObserver
623// ---------------------------------------------------------------------------
624
625/// Tracks per-channel running min/max of observed values.
626///
627/// Filters out NaN and Inf values before updating min/max.
628/// Logs a warning and returns an error when the channel count of incoming
629/// data doesn't match the configured number of channels.
630#[derive(Debug, Clone)]
631pub struct PerChannelMinMaxObserver {
632    num_channels: usize,
633    axis: usize,
634    min_vals: Vec<f32>,
635    max_vals: Vec<f32>,
636}
637
638impl PerChannelMinMaxObserver {
639    /// Create a new per-channel observer.
640    ///
641    /// * `num_channels` — expected number of channels.
642    /// * `axis` — the axis along which channels are sliced.
643    pub fn new(num_channels: usize, axis: usize) -> Self {
644        Self {
645            num_channels,
646            axis,
647            min_vals: vec![f32::INFINITY; num_channels],
648            max_vals: vec![f32::NEG_INFINITY; num_channels],
649        }
650    }
651
652    /// Observe a tensor's data with the given shape.
653    ///
654    /// Returns `Err` if the channel count along `self.axis` doesn't match.
655    pub fn observe_with_shape(&mut self, data: &[f32], shape: &[usize]) -> FerrotorchResult<()> {
656        if self.axis >= shape.len() {
657            return Err(FerrotorchError::InvalidArgument {
658                message: format!(
659                    "PerChannelMinMaxObserver axis {} out of range for {}-d tensor",
660                    self.axis,
661                    shape.len()
662                ),
663            });
664        }
665        let actual_channels = shape[self.axis];
666        if actual_channels != self.num_channels {
667            // The `Err` below carries the same information as the previous
668            // `eprintln!` (channel count, axis, observed value); duplicating
669            // it on stderr is `print_stderr` noise per `rust-quality` §4.
670            return Err(FerrotorchError::InvalidArgument {
671                message: format!(
672                    "PerChannelMinMaxObserver expected {} channels on axis {}, got {}",
673                    self.num_channels, self.axis, actual_channels
674                ),
675            });
676        }
677
678        for (i, &x) in data.iter().enumerate() {
679            if !x.is_finite() {
680                continue;
681            }
682            let ch = channel_index(i, shape, self.axis);
683            if x < self.min_vals[ch] {
684                self.min_vals[ch] = x;
685            }
686            if x > self.max_vals[ch] {
687                self.max_vals[ch] = x;
688            }
689        }
690        Ok(())
691    }
692}
693
694impl Observer for PerChannelMinMaxObserver {
695    fn observe(&mut self, data: &[f32]) {
696        // Without shape info, we treat data as [num_channels, N] where N = len / num_channels.
697        // If `data` isn't divisible by `num_channels`, skip silently — the
698        // caller can use the shape-aware `observe_with_shape` if they need a
699        // reportable error. The previous `eprintln!` was unactionable noise
700        // and is forbidden by `rust-quality` §4 (`print_stderr` lint).
701        if !data.len().is_multiple_of(self.num_channels) {
702            return;
703        }
704        let per_channel = data.len() / self.num_channels;
705        for (i, &x) in data.iter().enumerate() {
706            if !x.is_finite() {
707                continue;
708            }
709            let ch = i / per_channel;
710            if ch >= self.num_channels {
711                continue;
712            }
713            if x < self.min_vals[ch] {
714                self.min_vals[ch] = x;
715            }
716            if x > self.max_vals[ch] {
717                self.max_vals[ch] = x;
718            }
719        }
720    }
721
722    fn calculate_qparams(&self, dtype: QuantDtype) -> QParams {
723        let params: Vec<(f32, i32)> = self
724            .min_vals
725            .iter()
726            .zip(self.max_vals.iter())
727            .map(|(&mn, &mx)| compute_scale_zp(mn, mx, dtype))
728            .collect();
729        QParams {
730            scale: params.iter().map(|&(s, _)| s).collect(),
731            zero_point: params.iter().map(|&(_, z)| z).collect(),
732        }
733    }
734
735    fn reset(&mut self) {
736        self.min_vals.fill(f32::INFINITY);
737        self.max_vals.fill(f32::NEG_INFINITY);
738    }
739}
740
741// ---------------------------------------------------------------------------
742// HistogramObserver
743// ---------------------------------------------------------------------------
744
745/// Histogram-based observer that collects a distribution of values.
746///
747/// When the observed range expands, existing bin counts are redistributed
748/// into the new bin layout via linear interpolation rather than being zeroed.
749#[derive(Debug, Clone)]
750pub struct HistogramObserver {
751    num_bins: usize,
752    bins: Vec<u64>,
753    min_val: f32,
754    max_val: f32,
755    /// Whether we've seen any data yet.
756    initialized: bool,
757}
758
759impl HistogramObserver {
760    pub fn new(num_bins: usize) -> Self {
761        Self {
762            num_bins,
763            bins: vec![0u64; num_bins],
764            min_val: f32::INFINITY,
765            max_val: f32::NEG_INFINITY,
766            initialized: false,
767        }
768    }
769
770    /// Redistribute old bins into a new range via linear interpolation.
771    fn redistribute(&mut self, new_min: f32, new_max: f32) {
772        if !self.initialized || self.bins.iter().all(|&c| c == 0) {
773            self.min_val = new_min;
774            self.max_val = new_max;
775            return;
776        }
777
778        let old_min = self.min_val;
779        let old_max = self.max_val;
780        let old_range = old_max - old_min;
781        let new_range = new_max - new_min;
782
783        if old_range <= 0.0 || new_range <= 0.0 {
784            self.min_val = new_min;
785            self.max_val = new_max;
786            return;
787        }
788
789        let n = self.num_bins;
790        let old_bins = self.bins.clone();
791        self.bins.fill(0);
792
793        let old_bin_width = old_range / n as f32;
794        let new_bin_width = new_range / n as f32;
795
796        for (old_idx, &old_count) in old_bins.iter().enumerate().take(n) {
797            if old_count == 0 {
798                continue;
799            }
800            // Center of the old bin in value space.
801            let old_center = old_min + (old_idx as f32 + 0.5) * old_bin_width;
802            // Map to new bin index.
803            let new_frac = (old_center - new_min) / new_bin_width;
804            let new_idx = (new_frac as usize).min(n - 1);
805            self.bins[new_idx] += old_count;
806        }
807
808        self.min_val = new_min;
809        self.max_val = new_max;
810    }
811}
812
813impl Observer for HistogramObserver {
814    fn observe(&mut self, data: &[f32]) {
815        // First pass: find min/max of new data, filtering NaN/Inf.
816        let mut batch_min = f32::INFINITY;
817        let mut batch_max = f32::NEG_INFINITY;
818        for &x in data {
819            if !x.is_finite() {
820                continue;
821            }
822            if x < batch_min {
823                batch_min = x;
824            }
825            if x > batch_max {
826                batch_max = x;
827            }
828        }
829
830        if batch_min > batch_max {
831            // No finite values in this batch.
832            return;
833        }
834
835        // Check if range needs expanding.
836        let new_min = if self.initialized {
837            self.min_val.min(batch_min)
838        } else {
839            batch_min
840        };
841        let new_max = if self.initialized {
842            self.max_val.max(batch_max)
843        } else {
844            batch_max
845        };
846
847        if self.initialized && (new_min < self.min_val || new_max > self.max_val) {
848            // Range expanded — redistribute existing counts into new layout.
849            self.redistribute(new_min, new_max);
850        } else if !self.initialized {
851            self.min_val = new_min;
852            self.max_val = new_max;
853            self.initialized = true;
854        }
855
856        // Insert new data into bins.
857        let range = (self.max_val - self.min_val).max(f32::EPSILON);
858        let n = self.num_bins;
859        for &x in data {
860            if !x.is_finite() {
861                continue;
862            }
863            let frac = (x - self.min_val) / range;
864            let idx = ((frac * n as f32) as usize).min(n - 1);
865            self.bins[idx] += 1;
866        }
867    }
868
869    fn calculate_qparams(&self, dtype: QuantDtype) -> QParams {
870        QParams::asymmetric(self.min_val, self.max_val, dtype)
871    }
872
873    fn reset(&mut self) {
874        self.bins.fill(0);
875        self.min_val = f32::INFINITY;
876        self.max_val = f32::NEG_INFINITY;
877        self.initialized = false;
878    }
879}
880
881// ===========================================================================
882// FakeQuantize — differentiable quantize/dequantize for QAT
883// ===========================================================================
884
885/// Simulates quantization during training by quantizing and immediately
886/// dequantizing values, while allowing gradients to flow through via the
887/// straight-through estimator (STE).
888///
889/// Implements clipped STE: gradients are passed through unchanged for
890/// values within the quantization range `[dequantize(qmin), dequantize(qmax)]`,
891/// and zeroed for out-of-range values.
892#[derive(Debug, Clone)]
893pub struct FakeQuantize {
894    /// Target quantized dtype.
895    pub dtype: QuantDtype,
896    /// Cached quantization parameters.
897    pub qparams: Option<QParams>,
898    /// Whether the observer is enabled (collects statistics).
899    pub observer_enabled: bool,
900    /// Whether fake quantization is enabled.
901    pub fake_quant_enabled: bool,
902    /// The observer used to compute qparams.
903    observer: MinMaxObserver,
904}
905
906impl FakeQuantize {
907    /// Create a new FakeQuantize module.
908    pub fn new(dtype: QuantDtype) -> Self {
909        Self {
910            dtype,
911            qparams: None,
912            observer_enabled: true,
913            fake_quant_enabled: true,
914            observer: MinMaxObserver::new(),
915        }
916    }
917
918    /// Forward pass: observe data, fake-quantize, and return the result.
919    ///
920    /// Returns the fake-quantized data and a gradient mask for clipped STE.
921    /// The mask is 1.0 for in-range values and 0.0 for out-of-range values.
922    pub fn forward(&mut self, data: &[f32]) -> (Vec<f32>, Vec<f32>) {
923        if !self.fake_quant_enabled {
924            let ones = vec![1.0f32; data.len()];
925            return (data.to_vec(), ones);
926        }
927
928        // Observe if enabled.
929        if self.observer_enabled {
930            self.observer.observe(data);
931        }
932
933        // Calculate or use cached qparams.
934        // When observer is disabled and we have cached params, skip recalculation.
935        let qparams = if let Some(cached) = self.qparams.as_ref().filter(|_| !self.observer_enabled)
936        {
937            cached.clone()
938        } else {
939            let qp = self.observer.calculate_qparams(self.dtype);
940            self.qparams = Some(qp.clone());
941            qp
942        };
943
944        let scale = qparams.scale[0];
945        let zp = qparams.zero_point[0];
946        let qmin = self.dtype.qmin();
947        let qmax = self.dtype.qmax();
948
949        // Compute the dequantized range boundaries for clipped STE.
950        let range_min = (qmin as f32 - zp as f32) * scale;
951        let range_max = (qmax as f32 - zp as f32) * scale;
952
953        let mut output = Vec::with_capacity(data.len());
954        let mut grad_mask = Vec::with_capacity(data.len());
955
956        for &x in data {
957            // Fake quantize: quantize then dequantize.
958            let q = (x / scale + zp as f32)
959                .round()
960                .clamp(qmin as f32, qmax as f32);
961            let dq = (q - zp as f32) * scale;
962            output.push(dq);
963
964            // Clipped STE: zero gradient for out-of-range inputs.
965            if x >= range_min && x <= range_max {
966                grad_mask.push(1.0);
967            } else {
968                grad_mask.push(0.0);
969            }
970        }
971
972        (output, grad_mask)
973    }
974}
975
976// ===========================================================================
977// QatModel — quantization-aware training wrapper
978// ===========================================================================
979
980/// A layer with associated FakeQuantize modules for QAT.
981#[derive(Debug, Clone)]
982pub struct QatLayer {
983    /// FakeQuantize for this layer's weights.
984    pub weight_fq: FakeQuantize,
985    /// FakeQuantize for this layer's activations (applied after forward).
986    pub activation_fq: FakeQuantize,
987}
988
989/// Wraps a collection of named weight tensors for quantization-aware training.
990///
991/// Applies `FakeQuantize` to weights before forward and to activations after
992/// each layer's forward pass. Original weights are saved before fake-quantization
993/// and restored after forward to preserve full-precision values for gradient
994/// updates.
995#[derive(Debug)]
996pub struct QatModel {
997    /// Per-layer FakeQuantize state, keyed by layer name.
998    pub layers: HashMap<String, QatLayer>,
999    /// Target quantized dtype.
1000    pub dtype: QuantDtype,
1001}
1002
1003impl QatModel {
1004    /// Create a new QAT model wrapper.
1005    pub fn new(dtype: QuantDtype) -> Self {
1006        Self {
1007            layers: HashMap::new(),
1008            dtype,
1009        }
1010    }
1011
1012    /// Register a layer for QAT.
1013    pub fn register_layer(&mut self, name: &str) {
1014        self.layers.insert(
1015            name.to_string(),
1016            QatLayer {
1017                weight_fq: FakeQuantize::new(self.dtype),
1018                activation_fq: FakeQuantize::new(self.dtype),
1019            },
1020        );
1021    }
1022
1023    /// Fake-quantize weights for a named layer.
1024    ///
1025    /// Returns `(fake_quantized_weights, original_weights)` so the caller
1026    /// can restore originals after the forward pass.
1027    pub fn fake_quantize_weights(
1028        &mut self,
1029        layer_name: &str,
1030        weights: &[f32],
1031    ) -> FerrotorchResult<(Vec<f32>, Vec<f32>)> {
1032        let layer =
1033            self.layers
1034                .get_mut(layer_name)
1035                .ok_or_else(|| FerrotorchError::InvalidArgument {
1036                    message: format!("layer '{layer_name}' not registered for QAT"),
1037                })?;
1038
1039        // Save original weights.
1040        let originals = weights.to_vec();
1041
1042        // Fake-quantize (gradient mask is used during backward, not here).
1043        let (fq_weights, _mask) = layer.weight_fq.forward(weights);
1044
1045        Ok((fq_weights, originals))
1046    }
1047
1048    /// Fake-quantize activations for a named layer.
1049    ///
1050    /// Applied after each layer's forward output, not just the last layer.
1051    pub fn fake_quantize_activations(
1052        &mut self,
1053        layer_name: &str,
1054        activations: &[f32],
1055    ) -> FerrotorchResult<(Vec<f32>, Vec<f32>)> {
1056        let layer =
1057            self.layers
1058                .get_mut(layer_name)
1059                .ok_or_else(|| FerrotorchError::InvalidArgument {
1060                    message: format!("layer '{layer_name}' not registered for QAT"),
1061                })?;
1062
1063        let (fq_activations, grad_mask) = layer.activation_fq.forward(activations);
1064        Ok((fq_activations, grad_mask))
1065    }
1066}
1067
1068/// Prepare a set of named parameters for quantization-aware training.
1069///
1070/// Creates a `QatModel` and registers layers. Only parameters whose name
1071/// contains "weight" get weight FakeQuantize; bias parameters are skipped.
1072pub fn prepare_qat(param_names: &[&str], dtype: QuantDtype) -> QatModel {
1073    let mut model = QatModel::new(dtype);
1074
1075    for &name in param_names {
1076        // Extract the layer name (everything before the last `.weight` or `.bias`).
1077        let layer_name = if let Some(prefix) = name.strip_suffix(".weight") {
1078            prefix
1079        } else if let Some(prefix) = name.strip_suffix(".bias") {
1080            // Only register the layer if not already registered — don't apply
1081            // weight FakeQuantize to bias parameters.
1082            if !model.layers.contains_key(prefix) {
1083                model.register_layer(prefix);
1084            }
1085            continue;
1086        } else {
1087            name
1088        };
1089
1090        model.register_layer(layer_name);
1091    }
1092
1093    model
1094}
1095
1096// ===========================================================================
1097// CUDA RNG — fork/join for reproducible GPU random state
1098// ===========================================================================
1099
1100/// Thread-safe GPU RNG state for fork/join semantics.
1101///
1102/// Uses `Mutex` with graceful poisoning recovery to avoid panics
1103/// when a thread panics while holding the lock.
1104pub mod cuda_rng {
1105    use std::sync::Mutex;
1106
1107    /// Global RNG state — a simple seed counter.
1108    static RNG_STATE: Mutex<u64> = Mutex::new(0xdeadbeef_cafebabe);
1109
1110    /// Saved RNG states for fork/join.
1111    static RNG_STACK: Mutex<Vec<u64>> = Mutex::new(Vec::new());
1112
1113    /// Get the current RNG state, recovering gracefully from mutex poisoning.
1114    pub fn get_state() -> u64 {
1115        let guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1116        *guard
1117    }
1118
1119    /// Set the RNG state.
1120    pub fn set_state(state: u64) {
1121        let mut guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1122        *guard = state;
1123    }
1124
1125    /// Save the current RNG state to a stack and set a new state.
1126    ///
1127    /// Uses `unwrap_or_else(|e| e.into_inner())` to handle poisoned mutexes
1128    /// gracefully instead of panicking.
1129    pub fn fork_rng(new_seed: u64) {
1130        let current = {
1131            let guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1132            *guard
1133        };
1134
1135        {
1136            let mut stack = RNG_STACK.lock().unwrap_or_else(|e| e.into_inner());
1137            stack.push(current);
1138        }
1139
1140        set_state(new_seed);
1141    }
1142
1143    /// Restore the previously saved RNG state from the stack.
1144    ///
1145    /// Uses `unwrap_or_else(|e| e.into_inner())` to handle poisoned mutexes
1146    /// gracefully instead of panicking.
1147    pub fn join_rng() {
1148        let saved = {
1149            let mut stack = RNG_STACK.lock().unwrap_or_else(|e| e.into_inner());
1150            stack.pop()
1151        };
1152
1153        if let Some(state) = saved {
1154            set_state(state);
1155        }
1156    }
1157
1158    /// Advance the RNG state and return the new value.
1159    pub fn next_seed() -> u64 {
1160        let mut guard = RNG_STATE.lock().unwrap_or_else(|e| e.into_inner());
1161        // Simple splitmix64 step.
1162        *guard = guard.wrapping_add(0x9e3779b97f4a7c15);
1163        let mut z = *guard;
1164        z = (z ^ (z >> 30)).wrapping_mul(0xbf58476d1ce4e5b9);
1165        z = (z ^ (z >> 27)).wrapping_mul(0x94d049bb133111eb);
1166        z ^ (z >> 31)
1167    }
1168}
1169
1170// ---------------------------------------------------------------------------
1171// Tests
1172// ---------------------------------------------------------------------------
1173
1174#[cfg(test)]
1175mod tests {
1176    use super::*;
1177
1178    /// Helper: create a tensor from f32 data.
1179    fn make_tensor(data: &[f32], shape: &[usize]) -> Tensor<f32> {
1180        crate::from_slice(data, shape).unwrap()
1181    }
1182
1183    // ----- Round-trip quantize/dequantize -----
1184
1185    #[test]
1186    fn test_per_tensor_int8_roundtrip() {
1187        let data: Vec<f32> = (-10..=10).map(|x| x as f32 * 0.5).collect();
1188        let t = make_tensor(&data, &[data.len()]);
1189        let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1190        let rt: Tensor<f32> = dequantize(&qt).unwrap();
1191
1192        assert_eq!(rt.shape(), t.shape());
1193        let orig = t.data().unwrap();
1194        let recovered = rt.data().unwrap();
1195        for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1196            let err = (o - r).abs();
1197            // INT8 over [-5, 5]: step ≈ 10/255 ≈ 0.04, max error ≈ half step ≈ 0.02
1198            assert!(
1199                err < 0.05,
1200                "element {i}: original={o}, recovered={r}, error={err}"
1201            );
1202        }
1203    }
1204
1205    #[test]
1206    fn test_per_tensor_uint8_roundtrip() {
1207        let data: Vec<f32> = (0..=20).map(|x| x as f32 * 0.1).collect();
1208        let t = make_tensor(&data, &[data.len()]);
1209        let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Uint8).unwrap();
1210        let rt: Tensor<f32> = dequantize(&qt).unwrap();
1211
1212        let orig = t.data().unwrap();
1213        let recovered = rt.data().unwrap();
1214        for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1215            let err = (o - r).abs();
1216            // UINT8 over [0, 2]: step ≈ 2/255 ≈ 0.008
1217            assert!(
1218                err < 0.02,
1219                "element {i}: original={o}, recovered={r}, error={err}"
1220            );
1221        }
1222    }
1223
1224    #[test]
1225    fn test_per_tensor_int4_roundtrip() {
1226        // INT4 has only 16 levels, so larger quantization error is expected.
1227        let data: Vec<f32> = (-8..=7).map(|x| x as f32).collect();
1228        let t = make_tensor(&data, &[data.len()]);
1229        let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int4).unwrap();
1230        let rt: Tensor<f32> = dequantize(&qt).unwrap();
1231
1232        let orig = t.data().unwrap();
1233        let recovered = rt.data().unwrap();
1234        for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1235            let err = (o - r).abs();
1236            // INT4 over [-8, 7]: step = 15/15 = 1.0, max error ≈ 0.5
1237            assert!(
1238                err < 1.01,
1239                "element {i}: original={o}, recovered={r}, error={err}"
1240            );
1241        }
1242    }
1243
1244    // ----- Per-channel -----
1245
1246    #[test]
1247    fn test_per_channel_int8_roundtrip() {
1248        // Shape [3, 4]: 3 channels along axis 0, each with different ranges.
1249        #[rustfmt::skip]
1250        let data: Vec<f32> = vec![
1251            // channel 0: range [0, 3]
1252            0.0, 1.0, 2.0, 3.0,
1253            // channel 1: range [-10, 10]
1254            -10.0, -5.0, 5.0, 10.0,
1255            // channel 2: range [100, 200]
1256            100.0, 130.0, 170.0, 200.0,
1257        ];
1258        let t = make_tensor(&data, &[3, 4]);
1259        let qt = quantize(&t, QuantScheme::PerChannel(0), QuantDtype::Int8).unwrap();
1260        let rt: Tensor<f32> = dequantize(&qt).unwrap();
1261
1262        assert_eq!(qt.scale.len(), 3);
1263        assert_eq!(qt.zero_point.len(), 3);
1264
1265        let orig = t.data().unwrap();
1266        let recovered = rt.data().unwrap();
1267        for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1268            let err = (o - r).abs();
1269            // Each channel has its own scale, so error is relative to the
1270            // channel's range. Worst case channel 2: 100/255 ≈ 0.39.
1271            assert!(
1272                err < 0.5,
1273                "element {i}: original={o}, recovered={r}, error={err}"
1274            );
1275        }
1276    }
1277
1278    #[test]
1279    fn test_per_channel_axis_out_of_bounds() {
1280        let t = make_tensor(&[1.0, 2.0, 3.0], &[3]);
1281        let result = quantize(&t, QuantScheme::PerChannel(5), QuantDtype::Int8);
1282        assert!(result.is_err());
1283    }
1284
1285    // ----- Quantized matmul -----
1286
1287    #[test]
1288    fn test_quantized_matmul_identity() {
1289        // A * I should ≈ A after quantize -> matmul -> dequantize.
1290        let a_data = vec![1.0f32, 2.0, 3.0, 4.0];
1291        let a = make_tensor(&a_data, &[2, 2]);
1292        let eye = make_tensor(&[1.0, 0.0, 0.0, 1.0], &[2, 2]);
1293
1294        let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1295        let qi = quantize(&eye, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1296        let qc = quantized_matmul(&qa, &qi).unwrap();
1297        let c: Tensor<f32> = dequantize(&qc).unwrap();
1298
1299        assert_eq!(c.shape(), &[2, 2]);
1300        let c_data = c.data().unwrap();
1301        for (i, (&expected, &got)) in a_data.iter().zip(c_data.iter()).enumerate() {
1302            let err = (expected - got).abs();
1303            assert!(
1304                err < 0.5,
1305                "element {i}: expected={expected}, got={got}, error={err}"
1306            );
1307        }
1308    }
1309
1310    #[test]
1311    fn test_quantized_matmul_correctness() {
1312        // [2,3] x [3,2] -> [2,2]
1313        // A = [[1, 2, 3],
1314        //      [4, 5, 6]]
1315        // B = [[7,  8],
1316        //      [9, 10],
1317        //      [11, 12]]
1318        // A @ B = [[ 58,  64],
1319        //          [139, 154]]
1320        let a = make_tensor(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
1321        let b = make_tensor(&[7.0, 8.0, 9.0, 10.0, 11.0, 12.0], &[3, 2]);
1322
1323        let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1324        let qb = quantize(&b, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1325        let qc = quantized_matmul(&qa, &qb).unwrap();
1326        let c: Tensor<f32> = dequantize(&qc).unwrap();
1327
1328        let expected = [58.0f32, 64.0, 139.0, 154.0];
1329        let c_data = c.data().unwrap();
1330        assert_eq!(c.shape(), &[2, 2]);
1331        for (i, (&e, &g)) in expected.iter().zip(c_data.iter()).enumerate() {
1332            let err = (e - g).abs();
1333            // Quantization introduces some error; for small integers in INT8
1334            // the error should be small relative to the values.
1335            assert!(err < 3.0, "element {i}: expected={e}, got={g}, error={err}");
1336        }
1337    }
1338
1339    #[test]
1340    fn test_quantized_matmul_shape_mismatch() {
1341        let a = make_tensor(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
1342        let b = make_tensor(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
1343
1344        let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1345        let qb = quantize(&b, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1346        let result = quantized_matmul(&qa, &qb);
1347        assert!(result.is_err());
1348    }
1349
1350    #[test]
1351    fn test_quantized_matmul_non_2d() {
1352        let a = make_tensor(&[1.0, 2.0, 3.0], &[3]);
1353        let b = make_tensor(&[4.0, 5.0, 6.0], &[3]);
1354
1355        let qa = quantize(&a, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1356        let qb = quantize(&b, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1357        let result = quantized_matmul(&qa, &qb);
1358        assert!(result.is_err());
1359    }
1360
1361    // ----- Module quantization utility -----
1362
1363    #[test]
1364    fn test_quantize_named_tensors() {
1365        let w1 = make_tensor(&[1.0, 2.0, 3.0, 4.0], &[2, 2]);
1366        let w2 = make_tensor(&[-1.0, 0.0, 1.0, 2.0, 3.0, 4.0], &[3, 2]);
1367
1368        let named = vec![
1369            ("layer.weight".to_string(), w1),
1370            ("layer2.weight".to_string(), w2),
1371        ];
1372
1373        let qmap = quantize_named_tensors(named, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1374
1375        assert_eq!(qmap.len(), 2);
1376        assert!(qmap.contains_key("layer.weight"));
1377        assert!(qmap.contains_key("layer2.weight"));
1378        assert_eq!(qmap["layer.weight"].shape(), &[2, 2]);
1379        assert_eq!(qmap["layer2.weight"].shape(), &[3, 2]);
1380    }
1381
1382    // ----- Constant values / edge cases -----
1383
1384    #[test]
1385    fn test_quantize_constant_tensor() {
1386        // All values identical — scale should not be zero.
1387        let t = make_tensor(&[5.0, 5.0, 5.0, 5.0], &[4]);
1388        let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1389        let rt: Tensor<f32> = dequantize(&qt).unwrap();
1390
1391        let recovered = rt.data().unwrap();
1392        for &r in recovered {
1393            assert!(
1394                (r - 5.0).abs() < 0.1,
1395                "constant tensor dequantized to {r}, expected 5.0"
1396            );
1397        }
1398    }
1399
1400    #[test]
1401    fn test_quantize_single_element() {
1402        let t = make_tensor(&[42.0], &[1]);
1403        let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1404        let rt: Tensor<f32> = dequantize(&qt).unwrap();
1405        assert!((rt.data().unwrap()[0] - 42.0).abs() < 0.5);
1406    }
1407
1408    #[test]
1409    fn test_per_channel_int4() {
1410        // 2 channels, 3 elements each.
1411        let data = vec![0.0, 1.0, 2.0, -4.0, 0.0, 4.0];
1412        let t = make_tensor(&data, &[2, 3]);
1413        let qt = quantize(&t, QuantScheme::PerChannel(0), QuantDtype::Int4).unwrap();
1414
1415        assert_eq!(qt.scale.len(), 2);
1416        assert_eq!(qt.zero_point.len(), 2);
1417
1418        let rt: Tensor<f32> = dequantize(&qt).unwrap();
1419        let orig = t.data().unwrap();
1420        let recovered = rt.data().unwrap();
1421        for (i, (&o, &r)) in orig.iter().zip(recovered.iter()).enumerate() {
1422            let err = (o - r).abs();
1423            // INT4 has coarse resolution, but channel-level ranges are small.
1424            assert!(
1425                err < 1.0,
1426                "element {i}: original={o}, recovered={r}, error={err}"
1427            );
1428        }
1429    }
1430
1431    #[test]
1432    fn test_dequantize_f64() {
1433        let data = vec![1.0f32, 2.0, 3.0, 4.0];
1434        let t = crate::from_slice(&data, &[4]).unwrap();
1435        let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1436        let rt: Tensor<f64> = dequantize(&qt).unwrap();
1437
1438        assert_eq!(rt.shape(), &[4]);
1439        let recovered = rt.data().unwrap();
1440        for (i, &r) in recovered.iter().enumerate() {
1441            let expected = data[i] as f64;
1442            let err = (expected - r).abs();
1443            assert!(
1444                err < 0.05,
1445                "element {i}: expected={expected}, recovered={r}, error={err}"
1446            );
1447        }
1448    }
1449
1450    #[test]
1451    fn test_quantized_tensor_accessors() {
1452        let t = make_tensor(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
1453        let qt = quantize(&t, QuantScheme::PerTensor, QuantDtype::Int8).unwrap();
1454
1455        assert_eq!(qt.numel(), 6);
1456        assert_eq!(qt.shape(), &[2, 3]);
1457        assert_eq!(qt.data().len(), 6);
1458        assert_eq!(qt.scale().len(), 1);
1459        assert_eq!(qt.zero_point().len(), 1);
1460        assert_eq!(qt.scheme(), QuantScheme::PerTensor);
1461        assert_eq!(qt.qdtype(), QuantDtype::Int8);
1462    }
1463
1464    // ----- QParams -----
1465
1466    #[test]
1467    fn test_qparams_symmetric_int8() {
1468        let qp = QParams::symmetric(5.0, QuantDtype::Int8);
1469        assert_eq!(qp.zero_point, vec![0]);
1470        assert!((qp.scale[0] - 5.0 / 127.0).abs() < 1e-7);
1471    }
1472
1473    #[test]
1474    fn test_qparams_symmetric_uint8() {
1475        let qp = QParams::symmetric(5.0, QuantDtype::Uint8);
1476        assert_eq!(qp.zero_point, vec![128]);
1477        assert!((qp.scale[0] - 5.0 / 128.0).abs() < 1e-7);
1478    }
1479
1480    #[test]
1481    fn test_qparams_symmetric_int4() {
1482        let qp = QParams::symmetric(7.0, QuantDtype::Int4);
1483        assert_eq!(qp.zero_point, vec![0]);
1484        assert!((qp.scale[0] - 1.0).abs() < 1e-7);
1485    }
1486
1487    // ----- MinMaxObserver -----
1488
1489    #[test]
1490    fn test_minmax_observer() {
1491        let mut obs = MinMaxObserver::new();
1492        obs.observe(&[1.0, 2.0, 3.0]);
1493        obs.observe(&[-1.0, 5.0]);
1494        let qp = obs.calculate_qparams(QuantDtype::Int8);
1495        // Range includes zero: min=-1, max=5.
1496        assert_eq!(qp.scale.len(), 1);
1497        assert_eq!(qp.zero_point.len(), 1);
1498    }
1499
1500    #[test]
1501    fn test_minmax_observer_filters_nan_inf() {
1502        let mut obs = MinMaxObserver::new();
1503        obs.observe(&[1.0, f32::NAN, 2.0, f32::INFINITY, -1.0, f32::NEG_INFINITY]);
1504        let qp = obs.calculate_qparams(QuantDtype::Int8);
1505        // Should only see range [-1, 2], NaN/Inf filtered.
1506        let expected_range = 2.0 - (-1.0); // = 3.0
1507        let expected_scale = expected_range / 255.0;
1508        assert!((qp.scale[0] - expected_scale).abs() < 1e-5);
1509    }
1510
1511    // ----- PerChannelMinMaxObserver -----
1512
1513    #[test]
1514    fn test_per_channel_observer_with_shape() {
1515        let mut obs = PerChannelMinMaxObserver::new(2, 0);
1516        // Shape [2, 3]: channel 0 = [0, 1, 2], channel 1 = [10, 20, 30]
1517        obs.observe_with_shape(&[0.0, 1.0, 2.0, 10.0, 20.0, 30.0], &[2, 3])
1518            .unwrap();
1519        let qp = obs.calculate_qparams(QuantDtype::Int8);
1520        assert_eq!(qp.scale.len(), 2);
1521        assert_eq!(qp.zero_point.len(), 2);
1522    }
1523
1524    #[test]
1525    fn test_per_channel_observer_shape_mismatch() {
1526        let mut obs = PerChannelMinMaxObserver::new(3, 0);
1527        // Shape [2, 3] has 2 channels on axis 0, but observer expects 3.
1528        let result = obs.observe_with_shape(&[1.0; 6], &[2, 3]);
1529        assert!(result.is_err());
1530    }
1531
1532    #[test]
1533    fn test_per_channel_observer_axis() {
1534        let mut obs = PerChannelMinMaxObserver::new(3, 1);
1535        // Shape [2, 3]: axis 1 has 3 channels.
1536        obs.observe_with_shape(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3])
1537            .unwrap();
1538        let qp = obs.calculate_qparams(QuantDtype::Int8);
1539        assert_eq!(qp.scale.len(), 3);
1540    }
1541
1542    #[test]
1543    fn test_per_channel_observer_filters_nan_inf() {
1544        let mut obs = PerChannelMinMaxObserver::new(2, 0);
1545        obs.observe_with_shape(&[f32::NAN, 1.0, 2.0, 10.0, f32::INFINITY, 30.0], &[2, 3])
1546            .unwrap();
1547        // Channel 0 should only see [1, 2], channel 1 should only see [10, 30].
1548        let qp = obs.calculate_qparams(QuantDtype::Int8);
1549        assert_eq!(qp.scale.len(), 2);
1550    }
1551
1552    // ----- HistogramObserver -----
1553
1554    #[test]
1555    fn test_histogram_observer_basic() {
1556        let mut obs = HistogramObserver::new(100);
1557        obs.observe(&[0.0, 0.5, 1.0]);
1558        let qp = obs.calculate_qparams(QuantDtype::Int8);
1559        assert_eq!(qp.scale.len(), 1);
1560    }
1561
1562    #[test]
1563    fn test_histogram_observer_range_expansion() {
1564        let mut obs = HistogramObserver::new(100);
1565        obs.observe(&[0.0, 1.0]);
1566        // Initial range is [0, 1].
1567        let bins_after_first = obs.bins.clone();
1568        let total_first: u64 = bins_after_first.iter().sum();
1569        assert_eq!(total_first, 2);
1570
1571        obs.observe(&[-1.0, 2.0]);
1572        // Range expanded to [-1, 2]. Old counts should be redistributed, not zeroed.
1573        let total_second: u64 = obs.bins.iter().sum();
1574        // Should have 4 total counts (2 original redistributed + 2 new).
1575        assert_eq!(total_second, 4);
1576    }
1577
1578    #[test]
1579    fn test_histogram_observer_filters_nan_inf() {
1580        let mut obs = HistogramObserver::new(50);
1581        obs.observe(&[f32::NAN, 1.0, f32::INFINITY, 2.0]);
1582        let total: u64 = obs.bins.iter().sum();
1583        // Only 2 finite values should be counted.
1584        assert_eq!(total, 2);
1585    }
1586
1587    // ----- FakeQuantize -----
1588
1589    #[test]
1590    fn test_fake_quantize_roundtrip() {
1591        let mut fq = FakeQuantize::new(QuantDtype::Int8);
1592        let data = vec![0.0, 0.5, 1.0, 1.5, 2.0];
1593        let (output, mask) = fq.forward(&data);
1594        assert_eq!(output.len(), 5);
1595        assert_eq!(mask.len(), 5);
1596
1597        // Output should be close to input (quantize then dequantize).
1598        for (i, (&o, &d)) in output.iter().zip(data.iter()).enumerate() {
1599            assert!((o - d).abs() < 0.1, "element {i}: output={o}, data={d}");
1600        }
1601    }
1602
1603    #[test]
1604    // reason: STE mask is binary 0.0/1.0 — written as exact bit patterns,
1605    // never the result of arithmetic, so equality is the right check.
1606    #[allow(clippy::float_cmp)]
1607    fn test_fake_quantize_ste_clipping() {
1608        let mut fq = FakeQuantize::new(QuantDtype::Int8);
1609        // First, observe a range [0, 2].
1610        let (_, _) = fq.forward(&[0.0, 1.0, 2.0]);
1611
1612        // Disable observer so range stays locked at [0, 2].
1613        fq.observer_enabled = false;
1614
1615        // Now forward with values outside the observed range.
1616        let (_, mask) = fq.forward(&[0.5, 1.0, 100.0, -100.0]);
1617        // In-range values should have mask = 1.0.
1618        assert_eq!(mask[0], 1.0);
1619        assert_eq!(mask[1], 1.0);
1620        // Out-of-range values should have mask = 0.0.
1621        assert_eq!(mask[2], 0.0);
1622        assert_eq!(mask[3], 0.0);
1623    }
1624
1625    #[test]
1626    fn test_fake_quantize_observer_disabled_uses_cached() {
1627        let mut fq = FakeQuantize::new(QuantDtype::Int8);
1628        // Observe initial range.
1629        let (_, _) = fq.forward(&[0.0, 10.0]);
1630        let cached_scale = fq.qparams.as_ref().unwrap().scale[0];
1631
1632        // Disable observer.
1633        fq.observer_enabled = false;
1634
1635        // Forward with a much larger range — should NOT update qparams.
1636        let (_, _) = fq.forward(&[0.0, 1000.0]);
1637        let scale_after = fq.qparams.as_ref().unwrap().scale[0];
1638        assert!(
1639            (scale_after - cached_scale).abs() < 1e-10,
1640            "scale should not change when observer is disabled"
1641        );
1642    }
1643
1644    #[test]
1645    // reason: with fake_quant disabled the STE mask is filled with the exact
1646    // bit pattern 1.0 (no arithmetic), so equality is the right check.
1647    #[allow(clippy::float_cmp)]
1648    fn test_fake_quantize_disabled_is_identity() {
1649        let mut fq = FakeQuantize::new(QuantDtype::Int8);
1650        fq.fake_quant_enabled = false;
1651        let data = vec![1.234, 5.678, -9.012];
1652        let (output, mask) = fq.forward(&data);
1653        assert_eq!(output, data);
1654        assert!(mask.iter().all(|&m| m == 1.0));
1655    }
1656
1657    // ----- QatModel -----
1658
1659    #[test]
1660    fn test_qat_model_register_and_fq_weights() {
1661        let mut model = QatModel::new(QuantDtype::Int8);
1662        model.register_layer("fc1");
1663
1664        let weights = vec![0.1, 0.2, 0.3, 0.4];
1665        let (fq_weights, originals) = model.fake_quantize_weights("fc1", &weights).unwrap();
1666
1667        // Originals should be exact copies.
1668        assert_eq!(originals, weights);
1669        // Fake-quantized weights should be close to originals.
1670        for (i, (&fq, &orig)) in fq_weights.iter().zip(weights.iter()).enumerate() {
1671            assert!((fq - orig).abs() < 0.1, "weight {i}: fq={fq}, orig={orig}");
1672        }
1673    }
1674
1675    #[test]
1676    fn test_qat_model_activation_fq_per_layer() {
1677        let mut model = QatModel::new(QuantDtype::Int8);
1678        model.register_layer("layer1");
1679        model.register_layer("layer2");
1680
1681        // Both layers should have independent activation FakeQuantize.
1682        let (act1, _) = model
1683            .fake_quantize_activations("layer1", &[1.0, 2.0])
1684            .unwrap();
1685        let (act2, _) = model
1686            .fake_quantize_activations("layer2", &[10.0, 20.0])
1687            .unwrap();
1688        assert_eq!(act1.len(), 2);
1689        assert_eq!(act2.len(), 2);
1690    }
1691
1692    #[test]
1693    fn test_qat_model_unregistered_layer_errors() {
1694        let mut model = QatModel::new(QuantDtype::Int8);
1695        let result = model.fake_quantize_weights("nonexistent", &[1.0]);
1696        assert!(result.is_err());
1697    }
1698
1699    // ----- prepare_qat -----
1700
1701    #[test]
1702    fn test_prepare_qat_skips_bias() {
1703        let names = &["fc1.weight", "fc1.bias", "fc2.weight", "fc2.bias"];
1704        let model = prepare_qat(names, QuantDtype::Int8);
1705
1706        assert!(model.layers.contains_key("fc1"));
1707        assert!(model.layers.contains_key("fc2"));
1708        assert_eq!(model.layers.len(), 2);
1709    }
1710
1711    #[test]
1712    fn test_prepare_qat_bias_only_still_registers() {
1713        let names = &["fc1.bias"];
1714        let model = prepare_qat(names, QuantDtype::Int8);
1715        // Even bias-only parameters should get a layer registered.
1716        assert!(model.layers.contains_key("fc1"));
1717    }
1718
1719    // ----- cuda_rng -----
1720    //
1721    // The cuda_rng module exposes a process-global `Mutex<u64>` state plus
1722    // a fork/join stack. Both tests below mutate that state and read it
1723    // back; under cargo's default parallel test runner they race with each
1724    // other. Serialise via a local static mutex (same pattern as the
1725    // capture-lock used in #602 for the GPU-graph tests).
1726    fn cuda_rng_test_lock() -> std::sync::MutexGuard<'static, ()> {
1727        static LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
1728        LOCK.lock().unwrap_or_else(|p| p.into_inner())
1729    }
1730
1731    #[test]
1732    fn test_cuda_rng_fork_join() {
1733        let _g = cuda_rng_test_lock();
1734        let initial = cuda_rng::get_state();
1735        cuda_rng::fork_rng(0x12345678);
1736        assert_eq!(cuda_rng::get_state(), 0x12345678);
1737        cuda_rng::join_rng();
1738        assert_eq!(cuda_rng::get_state(), initial);
1739    }
1740
1741    #[test]
1742    fn test_cuda_rng_next_seed() {
1743        let _g = cuda_rng_test_lock();
1744        let s1 = cuda_rng::next_seed();
1745        let s2 = cuda_rng::next_seed();
1746        assert_ne!(s1, s2, "consecutive seeds should differ");
1747    }
1748}