Skip to main content

burn_tensor/tensor/
quantization.rs

1use crate::{
2    Tensor,
3    ops::{BridgeKind, BridgeTensor},
4};
5use burn_backend::quantization;
6
7// User-facing quantization data types come from burn-std.
8use burn_dispatch::Dispatch;
9pub use burn_std::quantization::*;
10
11/// The tensor quantization parameters.
12pub type QuantizationParameters = QParams<Tensor<1>>;
13
14/// The observed input calibration range.
15#[derive(Clone, Debug)]
16pub struct CalibrationRange {
17    /// Minimum observed value(s).
18    pub min: Tensor<1>,
19    /// Maximum observed value(s).
20    pub max: Tensor<1>,
21}
22
23/// Compute the quantization range mapping.
24pub fn compute_range<const D: usize>(
25    scheme: &QuantScheme,
26    tensor: &Tensor<D>,
27    calibration: &Calibration,
28) -> CalibrationRange {
29    let (kind, inner) = tensor.primitive.as_parts();
30    let (min, max) = match kind {
31        BridgeKind::Float => {
32            quantization::compute_range::<Dispatch>(scheme, inner.clone(), calibration)
33        }
34        BridgeKind::QFloat => unreachable!(),
35        _ => panic!("Should be Float primitive kind"),
36    };
37
38    CalibrationRange {
39        min: Tensor::new(BridgeTensor::float(min)),
40        max: Tensor::new(BridgeTensor::float(max)),
41    }
42}
43
44/// Compute the quantization parameters.
45pub fn compute_q_params(scheme: &QuantScheme, range: CalibrationRange) -> QuantizationParameters {
46    let (min_kind, min) = range.min.primitive.into_parts();
47    let (max_kind, max) = range.max.primitive.into_parts();
48    match (min_kind, max_kind) {
49        (BridgeKind::Float, BridgeKind::Float) => {
50            let qparams = quantization::compute_q_params::<Dispatch>(scheme, min, max);
51            QuantizationParameters {
52                scales: Tensor::new(BridgeTensor::float(qparams.scales)),
53            }
54        }
55        _ => unreachable!(),
56    }
57}