ruda_tensor/tensor/quantization/
scheme.rs1pub 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
8pub 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
99pub 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 let (a, b) = scheme.value.range();
114
115 let min_abs = B::float_abs(min);
117 let max_abs = B::float_abs(max);
118
119 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}