1use torsh_core::{Result, TorshError};
35use torsh_tensor::Tensor;
36
37#[derive(Debug, Clone, Copy, PartialEq, Eq)]
39pub enum QuantizationMethod {
40 Symmetric,
42 Affine,
44 PerChannel,
46}
47
48#[derive(Debug, Clone)]
50pub struct QuantizationParams {
51 pub scale: f32,
53 pub zero_point: i32,
55 pub bits: usize,
57 pub method: QuantizationMethod,
59}
60
61#[derive(Debug, Clone)]
63pub struct QuantizedTensor {
64 pub data: Vec<i8>,
66 pub shape: Vec<usize>,
68 pub params: QuantizationParams,
70}
71
72pub fn quantize_matrix(
86 tensor: &Tensor,
87 bits: usize,
88 method: QuantizationMethod,
89) -> Result<(QuantizedTensor, QuantizationParams)> {
90 if tensor.shape().ndim() != 2 {
92 return Err(TorshError::InvalidArgument(
93 "Quantization requires 2D tensor".to_string(),
94 ));
95 }
96
97 if bits != 8 && bits != 16 {
98 return Err(TorshError::InvalidArgument(
99 "Only 8-bit and 16-bit quantization supported".to_string(),
100 ));
101 }
102
103 let params = calibrate_quantization(tensor, bits, method)?;
105
106 let shape_binding = tensor.shape();
108 let shape = shape_binding.dims();
109 let (rows, cols) = (shape[0], shape[1]);
110
111 let mut quantized_data = Vec::with_capacity(rows * cols);
112 for i in 0..rows {
113 for j in 0..cols {
114 let val = tensor.get(&[i, j])?;
115 let q_val = ((val / params.scale) + params.zero_point as f32).round() as i8;
116 quantized_data.push(q_val);
117 }
118 }
119
120 let quantized = QuantizedTensor {
121 data: quantized_data,
122 shape: vec![rows, cols],
123 params: params.clone(),
124 };
125
126 Ok((quantized, params))
127}
128
129pub fn quantize_matrix_per_channel(
142 tensor: &Tensor,
143 bits: usize,
144) -> Result<(QuantizedTensor, QuantizationParams)> {
145 quantize_matrix(tensor, bits, QuantizationMethod::PerChannel)
147}
148
149pub fn dequantize_matrix(
163 quantized: &QuantizedTensor,
164 params: &QuantizationParams,
165) -> Result<Tensor> {
166 let shape = &quantized.shape;
168 if shape.len() != 2 {
169 return Err(TorshError::InvalidArgument(
170 "Dequantization requires 2D shape".to_string(),
171 ));
172 }
173
174 let (rows, cols) = (shape[0], shape[1]);
175 let mut dequantized_data = Vec::with_capacity(rows * cols);
176
177 for &q_val in &quantized.data {
178 let val = (q_val as f32 - params.zero_point as f32) * params.scale;
179 dequantized_data.push(val);
180 }
181
182 Tensor::from_data(
183 dequantized_data,
184 vec![rows, cols],
185 torsh_core::DeviceType::Cpu,
186 )
187}
188
189pub fn quantized_matmul(
204 a: &QuantizedTensor,
205 a_params: &QuantizationParams,
206 b: &QuantizedTensor,
207 b_params: &QuantizationParams,
208) -> Result<Tensor> {
209 if a.shape.len() != 2 || b.shape.len() != 2 {
211 return Err(TorshError::InvalidArgument(
212 "Quantized matmul requires 2D tensors".to_string(),
213 ));
214 }
215
216 if a.shape[1] != b.shape[0] {
217 return Err(TorshError::InvalidArgument(format!(
218 "Incompatible dimensions for quantized matmul: {}x{} and {}x{}",
219 a.shape[0], a.shape[1], b.shape[0], b.shape[1]
220 )));
221 }
222
223 let a_deq = dequantize_matrix(a, a_params)?;
225 let b_deq = dequantize_matrix(b, b_params)?;
226
227 a_deq.matmul(&b_deq)
229}
230
231pub fn calibrate_quantization(
246 tensor: &Tensor,
247 bits: usize,
248 method: QuantizationMethod,
249) -> Result<QuantizationParams> {
250 let shape_binding = tensor.shape();
252 let shape = shape_binding.dims();
253 let mut min_val = f32::INFINITY;
254 let mut max_val = f32::NEG_INFINITY;
255
256 if shape.len() == 1 {
257 for i in 0..shape[0] {
258 let val = tensor.get(&[i])?;
259 min_val = min_val.min(val);
260 max_val = max_val.max(val);
261 }
262 } else if shape.len() == 2 {
263 for i in 0..shape[0] {
264 for j in 0..shape[1] {
265 let val = tensor.get(&[i, j])?;
266 min_val = min_val.min(val);
267 max_val = max_val.max(val);
268 }
269 }
270 } else {
271 return Err(TorshError::InvalidArgument(
272 "Calibration only supports 1D and 2D tensors".to_string(),
273 ));
274 }
275
276 let (scale, zero_point) = match method {
278 QuantizationMethod::Symmetric => {
279 let max_abs = max_val.abs().max(min_val.abs());
281 let qmax = (1 << (bits - 1)) - 1;
282 let scale = max_abs / qmax as f32;
283 (scale, 0)
284 }
285 QuantizationMethod::Affine | QuantizationMethod::PerChannel => {
286 let qmin = -(1 << (bits - 1));
288 let qmax = (1 << (bits - 1)) - 1;
289 let scale = (max_val - min_val) / (qmax - qmin) as f32;
290 let zero_point = qmin - (min_val / scale).round() as i32;
291 (scale, zero_point)
292 }
293 };
294
295 Ok(QuantizationParams {
296 scale,
297 zero_point,
298 bits,
299 method,
300 })
301}
302
303#[cfg(test)]
304mod tests {
305 use super::*;
306
307 #[test]
308 fn test_quantization_method_equality() {
309 assert_eq!(QuantizationMethod::Symmetric, QuantizationMethod::Symmetric);
310 assert_ne!(QuantizationMethod::Symmetric, QuantizationMethod::Affine);
311 }
312
313 #[test]
314 fn test_calibrate_quantization_symmetric() -> Result<()> {
315 let data = vec![-2.0f32, -1.0, 0.0, 1.0, 2.0];
316 let tensor = Tensor::from_data(data, vec![5], torsh_core::DeviceType::Cpu)?;
317
318 let params = calibrate_quantization(&tensor, 8, QuantizationMethod::Symmetric)?;
319
320 assert_eq!(params.bits, 8);
321 assert_eq!(params.zero_point, 0);
322 assert!(params.scale > 0.0);
323 assert_eq!(params.method, QuantizationMethod::Symmetric);
324
325 Ok(())
326 }
327
328 #[test]
329 fn test_calibrate_quantization_affine() -> Result<()> {
330 let data = vec![1.0f32, 2.0, 3.0, 4.0, 5.0];
331 let tensor = Tensor::from_data(data, vec![5], torsh_core::DeviceType::Cpu)?;
332
333 let params = calibrate_quantization(&tensor, 8, QuantizationMethod::Affine)?;
334
335 assert_eq!(params.bits, 8);
336 assert!(params.scale > 0.0);
337 assert_eq!(params.method, QuantizationMethod::Affine);
338
339 Ok(())
340 }
341
342 #[test]
343 fn test_quantize_dequantize_roundtrip() -> Result<()> {
344 let data = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0];
345 let tensor = Tensor::from_data(data.clone(), vec![2, 3], torsh_core::DeviceType::Cpu)?;
346
347 let (quantized, params) = quantize_matrix(&tensor, 8, QuantizationMethod::Symmetric)?;
349
350 let dequantized = dequantize_matrix(&quantized, ¶ms)?;
352
353 assert_eq!(dequantized.shape().dims(), &[2, 3]);
355
356 for i in 0..2 {
358 for j in 0..3 {
359 let original = tensor.get(&[i, j])?;
360 let recovered = dequantized.get(&[i, j])?;
361 let error = (original - recovered).abs();
362 assert!(error < 1.0, "Error too large: {error} at [{i}, {j}]");
363 }
364 }
365
366 Ok(())
367 }
368
369 #[test]
370 fn test_quantized_matmul_basic() -> Result<()> {
371 let a = Tensor::from_data(
373 vec![1.0f32, 2.0, 3.0, 4.0],
374 vec![2, 2],
375 torsh_core::DeviceType::Cpu,
376 )?;
377 let b = Tensor::from_data(
378 vec![5.0f32, 6.0, 7.0, 8.0],
379 vec![2, 2],
380 torsh_core::DeviceType::Cpu,
381 )?;
382
383 let (a_q, a_params) = quantize_matrix(&a, 8, QuantizationMethod::Symmetric)?;
385 let (b_q, b_params) = quantize_matrix(&b, 8, QuantizationMethod::Symmetric)?;
386
387 let c_q = quantized_matmul(&a_q, &a_params, &b_q, &b_params)?;
389
390 let c_expected = a.matmul(&b)?;
392
393 assert_eq!(c_q.shape().dims(), &[2, 2]);
395
396 for i in 0..2 {
398 for j in 0..2 {
399 let expected = c_expected.get(&[i, j])?;
400 let actual = c_q.get(&[i, j])?;
401 let rel_error = ((expected - actual) / expected).abs();
402 assert!(rel_error < 0.5, "Relative error too large: {rel_error}");
403 }
404 }
405
406 Ok(())
407 }
408
409 #[test]
410 fn test_dimension_validation() {
411 let tensor =
413 Tensor::from_data(vec![1.0f32; 8], vec![2, 2, 2], torsh_core::DeviceType::Cpu).unwrap();
414
415 let result = calibrate_quantization(&tensor, 8, QuantizationMethod::Symmetric);
416 assert!(result.is_err());
417 }
418}