use crate::{
Tensor,
ops::{BridgeKind, BridgeTensor},
};
use burn_backend::quantization;
use burn_dispatch::Dispatch;
pub use burn_std::quantization::*;
pub type QuantizationParameters = QParams<Tensor<1>>;
#[derive(Clone, Debug)]
pub struct CalibrationRange {
pub min: Tensor<1>,
pub max: Tensor<1>,
}
pub fn compute_range<const D: usize>(
scheme: &QuantScheme,
tensor: &Tensor<D>,
calibration: &Calibration,
) -> CalibrationRange {
let (kind, inner) = tensor.primitive.as_parts();
let (min, max) = match kind {
BridgeKind::Float => {
quantization::compute_range::<Dispatch>(scheme, inner.clone(), calibration)
}
BridgeKind::QFloat => unreachable!(),
_ => panic!("Should be Float primitive kind"),
};
CalibrationRange {
min: Tensor::new(BridgeTensor::float(min)),
max: Tensor::new(BridgeTensor::float(max)),
}
}
pub fn compute_q_params(scheme: &QuantScheme, range: CalibrationRange) -> QuantizationParameters {
let (min_kind, min) = range.min.primitive.into_parts();
let (max_kind, max) = range.max.primitive.into_parts();
match (min_kind, max_kind) {
(BridgeKind::Float, BridgeKind::Float) => {
let qparams = quantization::compute_q_params::<Dispatch>(scheme, min, max);
QuantizationParameters {
scales: Tensor::new(BridgeTensor::float(qparams.scales)),
}
}
_ => unreachable!(),
}
}