feat(batch12): online quant calibration, draft distillation loss, gradient noise scale
Documentation / Build User Guide (push) Successful in 12s
CI / Build (ubuntu-latest) (push) Failing after 1m3s
CI / Format Check (push) Failing after 19s
CI / Clippy Check (push) Failing after 18s
Documentation / Build API Documentation (push) Failing after 35s
CI / Build CPU-Only (Explicit) (push) Failing after 1m17s
Performance Benchmarks / Run Benchmarks (push) Successful in 2m10s
CI / Build (macos-latest) (push) Failing after 57s
CI / Test (macos-latest) (push) Has been skipped
CI / Test (ubuntu-latest) (push) Has been skipped
CI / Python Bindings (maturin) (macos-latest) (push) Has been skipped
CI / Python Bindings (maturin) (ubuntu-latest) (push) Has been skipped
CI / WASM Build + Size Check (push) Has been skipped
CI / Distributed Training Tests (push) Has been skipped
CI / CI Success (push) Failing after 1s

- Online quantization calibration (rtx-compress): OnlineCalibrator with MaxAbs,
  EmaMaxAbs{momentum}, Percentile{percentile,bins} methods; streaming observe();
  quantize_int8/dequantize_int8; int8_maxabs/int8_ema/fp8_maxabs convenience ctors;
  ModelCalibrator tracks all tensors; 17 tests + 2 doctests
- Draft distillation loss (rtx-transformers): KL(p_target‖p_draft) + CE hard-label
  with temperature scaling; log_softmax/softmax/kl_divergence/token_acceptance_estimate
  primitives; DistillAccumulator for epoch-level tracking; normalize_by_length;
  13 tests + 6 doctests
- Gradient noise scale (rtx-transformers): GradientNoiseScale with McCandlish 2018
  two-point B_noise estimator + Welford single-pass mode; EMA smoothing; should_increase/
  decrease_batch signals; GnsTracker with bounded history + trend detection;
  16 tests + 2 doctests

Co-Authored-By: Claude Sonnet 4.6 <[email protected]>
This commit is contained in:
Omar Sobh
2026-06-27 05:37:19 +00:00
co-authored by Claude Sonnet 4.6
parent 3ec300af6b
commit 924c237096
5 changed files with 1972 additions and 0 deletions
@@ -0,0 +1,611 @@
//! Online quantization calibration.
//!
//! Tracks running tensor statistics (EMA of amax, running percentile via
//! histogram) during actual inference or early training steps, eliminating the
//! need for a separate offline calibration pass.
//!
//! This follows the approach used by TensorRT-LLM and NVIDIA Transformer
//! Engine for INT8/FP8 deployment.
//!
//! # Example
//!
//! ```rust
//! use rtx_compress::quantization::online_calib::{CalibMethod, OnlineCalibrator, ModelCalibrator};
//!
//! // Single-tensor INT8 calibration via EMA
//! let mut calib = OnlineCalibrator::int8_ema("layer0.weight", 0.999);
//! calib.observe(&[0.1, -0.5, 0.3, 0.9]);
//! calib.observe(&[0.2, -0.4, 0.8, 0.1]);
//! let scale = calib.scale().expect("at least one observation required");
//!
//! // Model-wide calibration
//! let mut model_calib = ModelCalibrator::new(CalibMethod::MaxAbs, 127.0);
//! model_calib.observe("fc1.weight", &[0.1, -0.9, 0.5]);
//! model_calib.observe("fc2.weight", &[-0.3, 0.7, 0.2]);
//! let scales = model_calib.export_scales();
//! ```
use std::collections::HashMap;
// ─── CalibMethod ────────────────────────────────────────────────────────────
/// Method for computing the quantization scale from statistics.
#[derive(Debug, Clone, PartialEq)]
pub enum CalibMethod {
/// Track running maximum absolute value; `scale = amax / quant_max`.
MaxAbs,
/// Exponential moving average of amax; `scale = ema_amax / quant_max`.
///
/// `momentum` is the decay factor applied to the historical EMA value
/// (typical: 0.999). A new observation `a` updates the EMA as:
/// `ema = momentum * ema + (1 - momentum) * a`.
EmaMaxAbs {
/// Decay factor in `(0, 1)`. Higher values give more weight to
/// historical observations. Typical: 0.999.
momentum: f32,
},
/// Running percentile via a fixed-width histogram; avoids outlier
/// sensitivity.
///
/// Absolute values are bucketed into `bins` bins over `[0, running_max]`.
/// The effective amax is the upper edge of the bin at which the cumulative
/// count first reaches `percentile / 100` of total observations.
Percentile {
/// Target percentile in `(0, 100]`. Typical: 99.9.
percentile: f32,
/// Number of histogram bins. Typical: 512.
bins: usize,
},
}
// ─── OnlineCalibrator ────────────────────────────────────────────────────────
/// Online calibrator for a single tensor (one weight or activation).
///
/// Call [`observe`](OnlineCalibrator::observe) once per forward pass (or
/// inference step) to accumulate statistics, then call
/// [`scale`](OnlineCalibrator::scale) to obtain the quantization scale that
/// maps the observed range into `[-quant_max, quant_max]`.
#[derive(Debug, Clone)]
pub struct OnlineCalibrator {
/// Human-readable name, e.g. `"model.layers.0.self_attn.q_proj"`.
pub name: String,
/// Calibration method used to derive the scale.
pub method: CalibMethod,
/// Quantization max value: 127.0 for INT8, 448.0 for FP8 E4M3.
pub quant_max: f32,
/// Total number of `observe()` calls made so far.
pub num_observations: usize,
// --- internal state --------------------------------------------------------
/// Running absolute maximum across all observations (MaxAbs / Percentile).
running_amax: f32,
/// Current EMA of amax (EmaMaxAbs only).
ema_amax: f32,
/// Histogram bin counts (Percentile only). Length == `bins`.
histogram: Vec<u64>,
/// Lower edge of the histogram range (always 0.0 for one-sided abs).
hist_min: f32,
/// Upper edge of the histogram range (== running_amax when initialised).
hist_max: f32,
}
impl OnlineCalibrator {
// ── Constructors ──────────────────────────────────────────────────────────
/// Create a new calibrator with explicit parameters.
///
/// # Panics
///
/// Panics if `quant_max <= 0.0`.
pub fn new(name: impl Into<String>, method: CalibMethod, quant_max: f32) -> Self {
assert!(quant_max > 0.0, "quant_max must be positive");
let bins = match &method {
CalibMethod::Percentile { bins, .. } => *bins,
_ => 0,
};
Self {
name: name.into(),
method,
quant_max,
num_observations: 0,
running_amax: 0.0,
ema_amax: 0.0,
histogram: vec![0u64; bins],
hist_min: 0.0,
hist_max: 0.0,
}
}
/// Convenience: INT8 MaxAbs calibrator (`quant_max = 127.0`).
pub fn int8_maxabs(name: impl Into<String>) -> Self {
Self::new(name, CalibMethod::MaxAbs, 127.0)
}
/// Convenience: INT8 EMA calibrator (`quant_max = 127.0`).
pub fn int8_ema(name: impl Into<String>, momentum: f32) -> Self {
Self::new(name, CalibMethod::EmaMaxAbs { momentum }, 127.0)
}
/// Convenience: FP8 E4M3 MaxAbs calibrator (`quant_max = 448.0`).
pub fn fp8_maxabs(name: impl Into<String>) -> Self {
Self::new(name, CalibMethod::MaxAbs, 448.0)
}
// ── Core API ──────────────────────────────────────────────────────────────
/// Observe a batch of tensor values and update running statistics.
///
/// Absolute values are used internally; signs are discarded because
/// symmetric quantization uses a single scale.
///
/// Calling with an empty slice is a no-op.
pub fn observe(&mut self, data: &[f32]) {
if data.is_empty() {
return;
}
// Compute absolute maximum for this batch.
let batch_amax = data.iter().map(|v| v.abs()).fold(0.0f32, f32::max);
self.num_observations += 1;
match &self.method.clone() {
CalibMethod::MaxAbs => {
self.running_amax = self.running_amax.max(batch_amax);
}
CalibMethod::EmaMaxAbs { momentum } => {
let m = *momentum;
if self.num_observations == 1 {
// Bootstrap: first observation initialises both fields.
self.ema_amax = batch_amax;
self.running_amax = batch_amax;
} else {
self.ema_amax = m * self.ema_amax + (1.0 - m) * batch_amax;
self.running_amax = self.running_amax.max(batch_amax);
}
}
CalibMethod::Percentile { bins, .. } => {
let bins = *bins;
self.running_amax = self.running_amax.max(batch_amax);
if self.num_observations == 1 {
// Initialise histogram bounds on first call.
self.hist_min = 0.0;
self.hist_max = if batch_amax > 0.0 {
batch_amax
} else {
1.0 // guard against all-zero first batch
};
self.histogram = vec![0u64; bins];
}
// If new data exceeds current bounds, reset and rebuild with
// the expanded range. We lose prior counts but maintain a
// valid histogram shape going forward.
if batch_amax > self.hist_max {
self.hist_max = batch_amax;
self.histogram = vec![0u64; bins];
}
// Bucket all absolute values from this batch.
let range = self.hist_max - self.hist_min; // hist_min == 0
for &v in data {
let abs_v = v.abs();
if abs_v <= self.hist_max {
let idx = if range > 0.0 {
((abs_v / range) * bins as f32) as usize
} else {
0
};
let idx = idx.min(bins - 1);
self.histogram[idx] += 1;
}
// Values > hist_max were already handled by the range
// expansion above, so this branch should not be reached.
}
}
}
}
/// Compute the current scale from accumulated statistics.
///
/// Returns `None` if no observations have been made yet.
///
/// The scale `s` satisfies: `quantized = clamp(round(x / s), -quant_max, quant_max)`.
pub fn scale(&self) -> Option<f32> {
if self.num_observations == 0 {
return None;
}
let effective_amax = match &self.method {
CalibMethod::MaxAbs => self.running_amax,
CalibMethod::EmaMaxAbs { .. } => self.ema_amax,
CalibMethod::Percentile { percentile, bins } => {
self.percentile_amax(*percentile, *bins)
}
};
if effective_amax <= 0.0 {
// All zeros observed — return a unit scale to avoid division by zero.
return Some(1.0 / self.quant_max);
}
Some(effective_amax / self.quant_max)
}
/// Compute quantized zero-point.
///
/// Always returns 0 because symmetric quantization is used.
#[inline]
pub fn zero_point(&self) -> i32 {
0
}
/// Apply the calibrated scale to quantize `data` into INT8 range `[-128, 127]`.
///
/// Returns `None` if no observations have been made.
pub fn quantize_int8(&self, data: &[f32]) -> Option<Vec<i8>> {
let s = self.scale()?;
if s == 0.0 {
return Some(vec![0i8; data.len()]);
}
Some(
data.iter()
.map(|&v| {
let q = (v / s).round();
q.clamp(-128.0, 127.0) as i8
})
.collect(),
)
}
/// Dequantize INT8 back to f32 using the current scale.
///
/// `scale()` must have returned `Some` (i.e. at least one observation)
/// for the result to be meaningful; otherwise the scale defaults to
/// `1.0 / quant_max`.
pub fn dequantize_int8(&self, data: &[i8]) -> Vec<f32> {
let s = self.scale().unwrap_or(1.0 / self.quant_max);
data.iter().map(|&q| q as f32 * s).collect()
}
/// Reset all accumulated statistics.
///
/// After calling `reset()`, `scale()` returns `None` again.
pub fn reset(&mut self) {
self.num_observations = 0;
self.running_amax = 0.0;
self.ema_amax = 0.0;
for b in &mut self.histogram {
*b = 0;
}
self.hist_min = 0.0;
self.hist_max = 0.0;
}
/// Raw running amax (before EMA smoothing if applicable).
#[inline]
pub fn current_amax(&self) -> f32 {
self.running_amax
}
// ── Private helpers ───────────────────────────────────────────────────────
/// Walk the histogram to find the upper edge of the bin at which the
/// cumulative count first reaches `percentile`% of total observations.
fn percentile_amax(&self, percentile: f32, bins: usize) -> f32 {
let total: u64 = self.histogram.iter().sum();
if total == 0 || bins == 0 {
return self.running_amax;
}
// Clamp percentile to (0, 100].
let target_frac = (percentile / 100.0).clamp(0.0, 1.0);
let target_count = (target_frac * total as f32).ceil() as u64;
let mut cumulative: u64 = 0;
let bin_width = (self.hist_max - self.hist_min) / bins as f32;
for (i, &count) in self.histogram.iter().enumerate() {
cumulative += count;
if cumulative >= target_count {
// Upper edge of bin i.
return self.hist_min + (i + 1) as f32 * bin_width;
}
}
// All bins exhausted — return hist_max.
self.hist_max
}
}
// ─── ModelCalibrator ─────────────────────────────────────────────────────────
/// Manages online calibrators for all named tensors in a model.
///
/// On the first call to [`observe`](ModelCalibrator::observe) for a given
/// tensor name, a new [`OnlineCalibrator`] is created automatically with the
/// shared `method` and `quant_max`.
///
/// # Example
///
/// ```rust
/// use rtx_compress::quantization::online_calib::{CalibMethod, ModelCalibrator};
///
/// let mut mc = ModelCalibrator::new(CalibMethod::MaxAbs, 127.0);
/// mc.observe("layer0.weight", &[0.5, -0.9, 0.3]);
/// mc.observe("layer1.weight", &[0.1, -0.2, 0.8]);
///
/// assert_eq!(mc.num_tensors(), 2);
/// let scales = mc.export_scales();
/// assert!(scales.contains_key("layer0.weight"));
/// ```
#[derive(Debug)]
pub struct ModelCalibrator {
calibrators: HashMap<String, OnlineCalibrator>,
method: CalibMethod,
quant_max: f32,
}
impl ModelCalibrator {
/// Create a new model-level calibrator.
///
/// `method` and `quant_max` are shared across all per-tensor calibrators
/// created by this instance.
///
/// # Panics
///
/// Panics if `quant_max <= 0.0`.
pub fn new(method: CalibMethod, quant_max: f32) -> Self {
assert!(quant_max > 0.0, "quant_max must be positive");
Self {
calibrators: HashMap::new(),
method,
quant_max,
}
}
/// Observe a named tensor, creating a calibrator on first call.
pub fn observe(&mut self, name: &str, data: &[f32]) {
let method = self.method.clone();
let quant_max = self.quant_max;
let calib = self
.calibrators
.entry(name.to_string())
.or_insert_with(|| OnlineCalibrator::new(name, method, quant_max));
calib.observe(data);
}
/// Get the calibrated scale for a named tensor.
///
/// Returns `None` if the tensor has never been observed.
pub fn scale_for(&self, name: &str) -> Option<f32> {
self.calibrators.get(name)?.scale()
}
/// Number of distinct tensor names being tracked.
pub fn num_tensors(&self) -> usize {
self.calibrators.len()
}
/// Names of all tensors for which at least one observation has been made.
pub fn tensor_names(&self) -> Vec<&str> {
self.calibrators.keys().map(String::as_str).collect()
}
/// Export all calibrated scales as a `HashMap<name, scale>`.
///
/// Tensors that have been observed but whose `scale()` returns `None`
/// (should not occur in practice) are omitted.
pub fn export_scales(&self) -> HashMap<String, f32> {
self.calibrators
.iter()
.filter_map(|(name, calib)| calib.scale().map(|s| (name.clone(), s)))
.collect()
}
}
// ─── Tests ───────────────────────────────────────────────────────────────────
#[cfg(test)]
mod tests {
use super::*;
// Floating-point comparison helper.
fn approx_eq(a: f32, b: f32, tol: f32) -> bool {
(a - b).abs() <= tol
}
// ── OnlineCalibrator ─────────────────────────────────────────────────────
#[test]
fn test_maxabs_single_observe() {
let mut c = OnlineCalibrator::int8_maxabs("t");
c.observe(&[1.0, 2.0, 3.0]);
let s = c.scale().expect("should have scale after observe");
assert!(approx_eq(s, 3.0 / 127.0, 1e-6));
}
#[test]
fn test_maxabs_accumulates_across_calls() {
let mut c = OnlineCalibrator::int8_maxabs("t");
c.observe(&[1.0, 2.0]);
c.observe(&[5.0, 0.5]);
// Overall max is 5.0
let s = c.scale().unwrap();
assert!(approx_eq(s, 5.0 / 127.0, 1e-6));
}
#[test]
fn test_ema_smooths_spike() {
// Large spike on first observation, then many small values.
// EMA with high momentum should end up much less than the spike.
let mut c = OnlineCalibrator::int8_ema("t", 0.999);
c.observe(&[100.0]); // spike — bootstraps ema_amax = 100
for _ in 0..500 {
c.observe(&[1.0]);
}
let s = c.scale().unwrap();
// After 500 steps with momentum 0.999, EMA decays substantially.
// 100 * 0.999^500 ≈ 60.6, but then we add 0.001 * 1 each step,
// converging toward 1.0. The resulting scale should be << 100/127.
assert!(s < 100.0 / 127.0, "EMA scale {s} should be < spike/quant_max");
}
#[test]
fn test_percentile_ignores_outlier() {
// 999 values near 1.0 and one extreme outlier at 1000.0.
// At 99.9th percentile the outlier should be excluded.
let mut c = OnlineCalibrator::new(
"t",
CalibMethod::Percentile {
percentile: 99.9,
bins: 512,
},
127.0,
);
let mut data: Vec<f32> = std::iter::repeat(1.0f32).take(999).collect();
data.push(1000.0); // single outlier
c.observe(&data);
let s = c.scale().unwrap();
// The effective amax should be close to 1.0, not 1000.0.
// Allow generous tolerance due to bin discretisation.
assert!(
s < 10.0 / 127.0,
"percentile scale {s} should be near 1/127, not 1000/127"
);
}
#[test]
fn test_scale_none_before_observe() {
let c = OnlineCalibrator::int8_maxabs("t");
assert_eq!(c.scale(), None);
}
#[test]
fn test_quantize_int8_range() {
let mut c = OnlineCalibrator::int8_maxabs("t");
c.observe(&[-5.0, 5.0, 0.0]);
let large: Vec<f32> = (-200i32..=200).map(|i| i as f32 * 0.1).collect();
let q = c.quantize_int8(&large).unwrap();
for &v in &q {
assert!(v >= i8::MIN && v <= i8::MAX, "quantized value {v} out of i8 range");
}
}
#[test]
fn test_quantize_dequantize_roundtrip() {
let mut c = OnlineCalibrator::int8_maxabs("t");
let data: Vec<f32> = (0..=10).map(|i| i as f32 * 0.1).collect();
c.observe(&data);
let scale = c.scale().unwrap();
let q = c.quantize_int8(&data).unwrap();
let dq = c.dequantize_int8(&q);
// Round-trip error should be at most one LSB == scale.
for (&orig, &reconstructed) in data.iter().zip(dq.iter()) {
assert!(
(orig - reconstructed).abs() <= scale + 1e-5,
"round-trip error too large: orig={orig}, reconstructed={reconstructed}, scale={scale}"
);
}
}
#[test]
fn test_reset_clears_state() {
let mut c = OnlineCalibrator::int8_maxabs("t");
c.observe(&[1.0, 2.0, 3.0]);
assert!(c.scale().is_some());
c.reset();
assert_eq!(c.scale(), None);
assert_eq!(c.num_observations, 0);
assert_eq!(c.current_amax(), 0.0);
}
#[test]
fn test_int8_convenience_ctor() {
let c = OnlineCalibrator::int8_maxabs("t");
assert_eq!(c.quant_max, 127.0);
}
#[test]
fn test_fp8_convenience_ctor() {
let c = OnlineCalibrator::fp8_maxabs("t");
assert_eq!(c.quant_max, 448.0);
}
#[test]
fn test_zero_point_always_zero() {
let mut c = OnlineCalibrator::int8_maxabs("t");
c.observe(&[1.0]);
assert_eq!(c.zero_point(), 0);
}
#[test]
fn test_ema_bootstraps_on_first_observation() {
// After exactly one observation, EMA should equal that observation.
let mut c = OnlineCalibrator::int8_ema("t", 0.9);
c.observe(&[7.0]);
// ema_amax bootstrapped to 7.0 → scale = 7/127
let s = c.scale().unwrap();
assert!(approx_eq(s, 7.0 / 127.0, 1e-5));
}
// ── ModelCalibrator ──────────────────────────────────────────────────────
#[test]
fn test_model_calibrator_observe_multiple() {
let mut mc = ModelCalibrator::new(CalibMethod::MaxAbs, 127.0);
mc.observe("a", &[1.0]);
mc.observe("b", &[2.0]);
mc.observe("c", &[3.0]);
assert_eq!(mc.num_tensors(), 3);
}
#[test]
fn test_model_calibrator_export_scales() {
let mut mc = ModelCalibrator::new(CalibMethod::MaxAbs, 127.0);
mc.observe("w1", &[0.5, -0.9]);
mc.observe("w2", &[0.1, -0.3]);
let scales = mc.export_scales();
assert!(scales.contains_key("w1"));
assert!(scales.contains_key("w2"));
assert!(approx_eq(scales["w1"], 0.9 / 127.0, 1e-6));
assert!(approx_eq(scales["w2"], 0.3 / 127.0, 1e-6));
}
#[test]
fn test_model_calibrator_creates_on_first_observe() {
let mut mc = ModelCalibrator::new(CalibMethod::MaxAbs, 127.0);
assert_eq!(mc.num_tensors(), 0);
mc.observe("new_tensor", &[1.0]);
assert_eq!(mc.num_tensors(), 1);
assert!(mc.scale_for("new_tensor").is_some());
}
#[test]
fn test_model_calibrator_unknown_tensor_returns_none() {
let mc = ModelCalibrator::new(CalibMethod::MaxAbs, 127.0);
assert_eq!(mc.scale_for("does_not_exist"), None);
}
#[test]
fn test_percentile_multiple_batches() {
// Observe the same uniform distribution across two batches.
// With 99.9th percentile the scale should still be near max/quant_max.
let mut c = OnlineCalibrator::new(
"t",
CalibMethod::Percentile {
percentile: 99.9,
bins: 512,
},
127.0,
);
let batch: Vec<f32> = (1..=100).map(|i| i as f32 * 0.01).collect();
c.observe(&batch);
c.observe(&batch);
let s = c.scale().unwrap();
// Effective amax should be close to 1.0 (upper end of the batch).
assert!(s > 0.0 && s <= 2.0 / 127.0, "scale {s} out of expected range");
}
}