ruda_tensor/api/
quantization.rs1use crate::api::{Tensor, TensorPrimitive, backend::Backend};
2use crate::tensor::quantization;
3use crate::{Shape, TensorMetadata};
4
5pub use crate::{QTensorPrimitive, quantization::*};
7
8pub type QuantizationParameters<B> = QParams<Tensor<B, 1>>;
10
11#[derive(Clone, Debug)]
13pub struct CalibrationRange<B: Backend> {
14 pub min: Tensor<B, 1>,
16 pub max: Tensor<B, 1>,
18}
19
20pub fn compute_range<B: Backend, const D: usize>(
22 scheme: &QuantScheme,
23 tensor: &Tensor<B, D>,
24 calibration: &Calibration,
25) -> CalibrationRange<B> {
26 let (min, max) = match &tensor.primitive {
27 TensorPrimitive::Float(tensor) => {
28 quantization::compute_range::<B>(scheme, tensor.clone(), calibration)
29 }
30 TensorPrimitive::QFloat(_) => unreachable!(),
31 };
32
33 let min_shape = Shape::new([min.shape().num_elements()]);
34 let max_shape = Shape::new([max.shape().num_elements()]);
35 CalibrationRange {
36 min: Tensor::from_primitive(TensorPrimitive::Float(B::float_reshape(min, min_shape))),
37 max: Tensor::from_primitive(TensorPrimitive::Float(B::float_reshape(max, max_shape))),
38 }
39}
40
41pub fn compute_q_params<B: Backend>(
43 scheme: &QuantScheme,
44 range: CalibrationRange<B>,
45) -> QuantizationParameters<B> {
46 match (range.min.primitive, range.max.primitive) {
47 (TensorPrimitive::Float(min), TensorPrimitive::Float(max)) => {
48 let qparams = quantization::compute_q_params::<B>(scheme, min, max);
49 QuantizationParameters {
50 scales: Tensor::from_primitive(TensorPrimitive::Float(qparams.scales)),
51 }
52 }
53 _ => unreachable!(),
54 }
55}