Skip to main content

ruda_tensor/tensor/quantization/
scheme.rs

1pub use ruda_core::tensor::{QPARAM_ALIGN, params_shape};
2use ruda_core::tensor::{QuantLevel, QuantMode, QuantScheme, Shape, Slice};
3use alloc::vec;
4
5use super::{Calibration, QuantizationParametersPrimitive};
6use crate::{Backend, TensorMetadata, get_device_settings};
7
8/// Compute the quantization range mapping.
9pub fn compute_range<B: Backend>(
10    scheme: &QuantScheme,
11    tensor: B::FloatTensorPrimitive,
12    calibration: &Calibration,
13) -> (B::FloatTensorPrimitive, B::FloatTensorPrimitive) {
14    match calibration {
15        Calibration::MinMax => match scheme.level {
16            QuantLevel::Tensor => (B::float_min(tensor.clone()), B::float_max(tensor)),
17            QuantLevel::Block(block_size) => {
18                let shape = tensor.shape();
19                let block_dims = block_size.to_dim_vec(shape.rank());
20                assert!(!block_dims.contains(&0), "Quantization block dimensions must be nonzero");
21                let params_shape = params_shape(&shape, scheme.level);
22                if shape.num_elements() == 0 {
23                    let device = B::float_device(&tensor);
24                    let dtype = tensor.dtype().into();
25                    return (
26                        B::float_empty(params_shape.clone(), &device, dtype),
27                        B::float_empty(params_shape, &device, dtype),
28                    );
29                }
30
31                let mut contiguous_blocks = true;
32                let mut inner_axes_complete = true;
33                for (&dim, &block) in shape.iter().zip(&block_dims).rev() {
34                    let block = block as usize;
35                    contiguous_blocks &= dim.is_multiple_of(block)
36                        && (block == 1 || inner_axes_complete);
37                    inner_axes_complete &= dim == block;
38                }
39                if contiguous_blocks {
40                    let num_blocks = params_shape.num_elements();
41                    let block_elems = shape.num_elements() / num_blocks;
42                    let blocks = B::float_reshape(tensor, Shape::new([num_blocks, block_elems]));
43                    return (
44                        B::float_reshape(B::float_min_dim(blocks.clone(), 1), params_shape.clone()),
45                        B::float_reshape(B::float_max_dim(blocks, 1), params_shape),
46                    );
47                }
48
49                let mut min = tensor.clone();
50                let mut max = tensor;
51                for (axis, &block) in block_dims.iter().enumerate().rev() {
52                    if block > 1 {
53                        min = reduce_block_axis::<B>(min, axis, block as usize, B::float_min_dim);
54                        max = reduce_block_axis::<B>(max, axis, block as usize, B::float_max_dim);
55                    }
56                }
57                (min, max)
58            }
59        },
60    }
61}
62
63fn reduce_block_axis<B: Backend>(
64    tensor: B::FloatTensorPrimitive,
65    axis: usize,
66    block: usize,
67    reduce: fn(B::FloatTensorPrimitive, usize) -> B::FloatTensorPrimitive,
68) -> B::FloatTensorPrimitive {
69    let shape = tensor.shape();
70    if shape[axis] <= block {
71        return reduce(tensor, axis);
72    }
73    let full_blocks = shape[axis] / block;
74    let full_len = full_blocks * block;
75    let has_tail = full_len != shape[axis];
76    let mut slices = vec![Slice::from(..); shape.rank()];
77    let head = if has_tail {
78        slices[axis] = Slice::from(0..full_len);
79        B::float_slice(tensor.clone(), &slices)
80    } else {
81        tensor.clone()
82    };
83    let mut grouped_shape = shape.clone();
84    grouped_shape[axis] = full_blocks;
85    grouped_shape.insert(axis + 1, block);
86    let head = reduce(B::float_reshape(head, grouped_shape), axis + 1);
87    let mut output_shape = shape;
88    output_shape[axis] = full_blocks;
89    let head = B::float_reshape(head, output_shape);
90    if has_tail {
91        slices[axis] = Slice::from(full_len..);
92        let tail = reduce(B::float_slice(tensor, &slices), axis);
93        B::float_cat(vec![head, tail], axis)
94    } else {
95        head
96    }
97}
98
99/// Compute the quantization parameters.
100pub fn compute_q_params<B: Backend>(
101    scheme: &QuantScheme,
102    min: B::FloatTensorPrimitive,
103    max: B::FloatTensorPrimitive,
104) -> QuantizationParametersPrimitive<B> {
105    match scheme {
106        QuantScheme {
107            level: QuantLevel::Tensor | QuantLevel::Block(_),
108            mode: QuantMode::Symmetric,
109            ..
110        } => {
111            let bool_dtype = get_device_settings::<B>(&B::float_device(&min)).bool_dtype;
112            // Quantized range `[a, b]`
113            let (a, b) = scheme.value.range();
114
115            // Compute scale to convert an input value in range `[-alpha, alpha]`
116            let min_abs = B::float_abs(min);
117            let max_abs = B::float_abs(max);
118
119            // `min_abs.max_pair(max_abs)`
120            let mask = B::float_lower(min_abs.clone(), max_abs.clone(), bool_dtype);
121            let values_range =
122                B::float_mul_scalar(B::float_mask_where(min_abs, mask, max_abs), 2f32.into());
123
124            QuantizationParametersPrimitive {
125                scales: B::float_div_scalar(values_range, (b - a).into()),
126            }
127        }
128    }
129}