Skip to main content

ruprim_host/quantization/
conversion.rs

1use super::*;
2use ruda_core::tensor::quantization::params_shape;
3
4pub fn quantize_dynamic(tensor: HostTensor, scheme: &QuantScheme) -> HostQTensor {
5    let shape = tensor.shape();
6    let tensor = tensor.to_contiguous();
7    let float_data = float_storage_as_f32(&tensor);
8    let (a, b) = scheme.value.range();
9    let range = b - a;
10
11    let (quantized, scales) = match scheme.level {
12        QuantLevel::Tensor => {
13            // Pass 1: find alpha = max(|min|, |max|)
14            let mut alpha: f32 = 0.0;
15            for &x in &*float_data {
16                let abs = x.abs();
17                if abs > alpha {
18                    alpha = abs;
19                }
20            }
21            let scale = validated_scale(2.0 * alpha / range);
22            let inv_scale = 1.0 / scale;
23
24            // Pass 2: quantize
25            let quantized = float_data
26                .iter()
27                .map(|&x| (x * inv_scale).round().clamp(a, b) as i8)
28                .collect::<Vec<i8>>();
29
30            (quantized, alloc::vec![scale])
31        }
32        QuantLevel::Block(block_size) => {
33            let block_dims = block_size.to_dim_vec(shape.rank());
34            let params_shape = params_shape(&shape, scheme.level);
35            let mut alphas = alloc::vec![0.0f32; params_shape.num_elements()];
36            for (index, &x) in float_data.iter().enumerate() {
37                let block = block_param_index(index, &shape, &block_dims, &params_shape);
38                let abs = x.abs();
39                if abs > alphas[block] {
40                    alphas[block] = abs;
41                }
42            }
43            let scales = alphas.into_iter()
44                .map(|alpha| validated_scale(2.0 * alpha / range))
45                .collect::<Vec<_>>();
46            let inv_scales = scales.iter().map(|scale| 1.0 / scale).collect::<Vec<_>>();
47            let quantized = float_data.iter().enumerate()
48                .map(|(index, &x)| {
49                    let block = block_param_index(index, &shape, &block_dims, &params_shape);
50                    (x * inv_scales[block]).round().clamp(a, b) as i8
51                })
52                .collect();
53            (quantized, scales)
54        }
55    };
56
57    let bytes = Bytes::from_elems(quantized);
58    let layout = Layout::contiguous(shape);
59    let qt = HostTensor::new(bytes, layout, DType::I8);
60
61    HostQTensor::new(qt, scheme.with_store(QuantStore::Native), scales)
62}
63
64pub fn quantize(
65    tensor: HostTensor,
66    scheme: &QuantScheme,
67    qparams: QParams<HostTensor>,
68) -> HostQTensor {
69    let shape = tensor.shape();
70    let tensor = tensor.to_contiguous();
71    let float_data = float_storage_as_f32(&tensor);
72
73    // Extract and validate scales from the qparams tensor. The scales tensor
74    // shares its dtype with the float element type, which can be any of
75    // f32/f64/f16/bf16, so we normalise via float_storage_as_f32 instead of
76    // assuming f32 storage.
77    let scales_tensor = qparams.scales.to_contiguous();
78    let scales_data = float_storage_as_f32(&scales_tensor);
79    let scales: Vec<f32> = scales_data.iter().copied().map(validated_scale).collect();
80    assert_eq!(
81        scales.len(), params_shape(&shape, scheme.level).num_elements(),
82        "quantized scale count must match the parameter shape"
83    );
84
85    let (a, b) = scheme.value.range();
86
87    let quantized = match scheme.level {
88        QuantLevel::Tensor => {
89            let inv_scale = 1.0 / scales[0];
90            float_data
91                .iter()
92                .map(|&x| (x * inv_scale).round().clamp(a, b) as i8)
93                .collect::<Vec<i8>>()
94        }
95        QuantLevel::Block(block_size) => {
96            let block_dims = block_size.to_dim_vec(shape.rank());
97            let params_shape = params_shape(&shape, scheme.level);
98            let inv_scales = scales.iter().map(|scale| 1.0 / scale).collect::<Vec<_>>();
99            float_data.iter().enumerate()
100                .map(|(index, &x)| {
101                    let block = block_param_index(index, &shape, &block_dims, &params_shape);
102                    (x * inv_scales[block]).round().clamp(a, b) as i8
103                })
104                .collect::<Vec<_>>()
105        }
106    };
107
108    let bytes = Bytes::from_elems(quantized);
109    let layout = Layout::contiguous(shape);
110    let qt = HostTensor::new(bytes, layout, DType::I8);
111
112    HostQTensor::new(qt, scheme.with_store(QuantStore::Native), scales)
113}
114
115pub fn dequantize(tensor: HostQTensor, dtype: FloatDType) -> HostTensor {
116    let shape = tensor.tensor.shape();
117    let qt = tensor.tensor.to_contiguous();
118    let q_data: &[i8] = qt.storage();
119
120    let dequantized = match tensor.scheme.level {
121        QuantLevel::Tensor => {
122            let scale = tensor.scales[0];
123            q_data
124                .iter()
125                .map(|&x_q| scale * x_q as f32)
126                .collect::<Vec<f32>>()
127        }
128        QuantLevel::Block(block_size) => {
129            let block_dims = block_size.to_dim_vec(shape.rank());
130            let params_shape = params_shape(&shape, tensor.scheme.level);
131            q_data
132                .iter().enumerate()
133                .map(|(index, &x_q)| {
134                    let block = block_param_index(index, &shape, &block_dims, &params_shape);
135                    tensor.scales[block] * x_q as f32
136                })
137                .collect::<Vec<f32>>()
138        }
139    };
140
141    let layout = Layout::contiguous(shape);
142    match dtype {
143        FloatDType::F32 | FloatDType::Flex32 => {
144            HostTensor::new(Bytes::from_elems(dequantized), layout, DType::F32)
145        }
146        FloatDType::F64 => {
147            let data: Vec<f64> = dequantized.iter().map(|&v| v as f64).collect();
148            HostTensor::new(Bytes::from_elems(data), layout, DType::F64)
149        }
150        FloatDType::F16 => {
151            let data: Vec<f16> = dequantized.iter().map(|&v| f16::from_f32(v)).collect();
152            HostTensor::new(Bytes::from_elems(data), layout, DType::F16)
153        }
154        FloatDType::BF16 => {
155            let data: Vec<bf16> = dequantized.iter().map(|&v| bf16::from_f32(v)).collect();
156            HostTensor::new(Bytes::from_elems(data), layout, DType::BF16)
157        }
158    }
159}
160
161fn block_param_index(mut index: usize, shape: &Shape, block_dims: &[u8], params_shape: &Shape) -> usize {
162    let mut parameter = 0;
163    let mut stride = 1;
164    for axis in (0..shape.rank()).rev() {
165        let coordinate = index % shape[axis];
166        index /= shape[axis];
167        parameter += coordinate / block_dims[axis] as usize * stride;
168        stride *= params_shape[axis];
169    }
170    parameter
171}
172
173/// Ensure scale is finite and nonzero to avoid division by zero or NaN propagation.
174fn validated_scale(scale: f32) -> f32 {
175    if scale.is_normal() {
176        scale
177    } else {
178        f32::MIN_POSITIVE
179    }
180}
181