burn_tensor/tensor/
quantization.rs1use crate::{
2 Tensor,
3 ops::{BridgeKind, BridgeTensor},
4};
5use burn_backend::quantization;
6
7use burn_dispatch::Dispatch;
9pub use burn_std::quantization::*;
10
11pub type QuantizationParameters = QParams<Tensor<1>>;
13
14#[derive(Clone, Debug)]
16pub struct CalibrationRange {
17 pub min: Tensor<1>,
19 pub max: Tensor<1>,
21}
22
23pub 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
44pub 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}