1use alloc::vec::Vec;
4#[cfg(not(feature = "std"))]
5#[allow(unused_imports)]
6use num_traits::Float;
7
8use burn_backend::{
9 DType, ExecutionError, FloatDType, TensorData, TensorMetadata,
10 ops::{IntTensorOps, QTensorOps},
11 quantization::{
12 QuantLevel, QuantParam, QuantScheme, QuantStore, QuantizationParametersPrimitive,
13 QuantizedBytes, scale_to_param,
14 },
15 tensor::{Device, FloatTensor, IntTensor, QuantizedTensor},
16};
17use burn_std::{Bytes, Shape, Slice, bf16, f16};
18
19use super::float_storage_as_f32;
20use crate::{Flex, FlexQTensor, FlexTensor, Layout};
21
22impl QTensorOps<Flex> for Flex {
23 fn q_from_data(data: TensorData, _device: &Device<Flex>) -> QuantizedTensor<Flex> {
24 let scheme = match data.dtype {
25 DType::QFloat(scheme) => scheme,
26 _ => panic!("Expected quantized dtype, got {:?}", data.dtype),
27 };
28
29 let shape = data.shape.clone();
30 let num_elements = data.num_elements();
31
32 let q_bytes = QuantizedBytes {
33 bytes: data.into_bytes(),
34 scheme,
35 num_elements,
36 };
37
38 let (values, qparams) = q_bytes.into_vec_i8();
39 let tensor_data = TensorData::new(values, shape);
40 let tensor = FlexTensor::from_data(tensor_data);
41
42 let scheme = scheme.with_store(QuantStore::Native);
44
45 FlexQTensor::new(tensor, scheme, qparams.scales)
46 }
47
48 fn quantize_dynamic(tensor: FloatTensor<Flex>, scheme: &QuantScheme) -> QuantizedTensor<Flex> {
49 let shape = tensor.shape();
50 let tensor = tensor.to_contiguous();
51 let float_data = float_storage_as_f32(&tensor);
52 let (a, b) = scheme.value.range();
53 let range = b - a;
54
55 let (quantized, scales) = match scheme.level {
56 QuantLevel::BlockTensor { .. } => {
57 unimplemented!("two-level quantization is not supported yet")
58 }
59 QuantLevel::Tensor => {
60 let mut alpha: f32 = 0.0;
62 for &x in &*float_data {
63 let abs = x.abs();
64 if abs > alpha {
65 alpha = abs;
66 }
67 }
68 let scale = validated_scale(2.0 * alpha / range, scheme.param);
69 let inv_scale = 1.0 / scale;
70
71 let quantized = float_data
73 .iter()
74 .map(|&x| (x * inv_scale).round().clamp(a, b) as i8)
75 .collect::<Vec<i8>>();
76
77 (quantized, alloc::vec![scale])
78 }
79 QuantLevel::Block(block_size) => {
80 let block_elems = block_size.num_elements();
81 debug_assert!(
82 float_data.len().is_multiple_of(block_elems),
83 "tensor length {} not divisible by block size {}",
84 float_data.len(),
85 block_elems
86 );
87 let num_blocks = float_data.len() / block_elems;
88 let mut scales = Vec::with_capacity(num_blocks);
89 let mut quantized = Vec::with_capacity(float_data.len());
90
91 for block in float_data.chunks(block_elems) {
92 let mut alpha: f32 = 0.0;
94 for &x in block {
95 let abs = x.abs();
96 if abs > alpha {
97 alpha = abs;
98 }
99 }
100 let scale = validated_scale(2.0 * alpha / range, scheme.param);
101 let inv_scale = 1.0 / scale;
102 scales.push(scale);
103
104 for &x in block {
106 quantized.push((x * inv_scale).round().clamp(a, b) as i8);
107 }
108 }
109
110 (quantized, scales)
111 }
112 };
113
114 let bytes = Bytes::from_elems(quantized);
115 let layout = Layout::contiguous(shape);
116 let qt = FlexTensor::new(bytes, layout, DType::I8);
117
118 FlexQTensor::new(qt, scheme.with_store(QuantStore::Native), scales)
119 }
120
121 fn quantize(
122 tensor: FloatTensor<Flex>,
123 scheme: &QuantScheme,
124 qparams: QuantizationParametersPrimitive<Flex>,
125 ) -> QuantizedTensor<Flex> {
126 let shape = tensor.shape();
127 let tensor = tensor.to_contiguous();
128 let float_data = float_storage_as_f32(&tensor);
129
130 let scales_tensor = qparams.scales.to_contiguous();
135 let scales_data = float_storage_as_f32(&scales_tensor);
136 let scales: Vec<f32> = scales_data
137 .iter()
138 .copied()
139 .map(|s| validated_scale(s, scheme.param))
140 .collect();
141
142 let (a, b) = scheme.value.range();
143
144 let quantized = match scheme.level {
145 QuantLevel::BlockTensor { .. } => {
146 unimplemented!("two-level quantization is not supported yet")
147 }
148 QuantLevel::Tensor => {
149 let inv_scale = 1.0 / scales[0];
150 float_data
151 .iter()
152 .map(|&x| (x * inv_scale).round().clamp(a, b) as i8)
153 .collect::<Vec<i8>>()
154 }
155 QuantLevel::Block(block_size) => {
156 let block_elems = block_size.num_elements();
157 debug_assert!(
158 float_data.len().is_multiple_of(block_elems),
159 "tensor length {} not divisible by block size {}",
160 float_data.len(),
161 block_elems
162 );
163 let mut quantized = Vec::with_capacity(float_data.len());
164 for (block, &scale) in float_data.chunks(block_elems).zip(scales.iter()) {
165 let inv_scale = 1.0 / scale;
166 for &x in block {
167 quantized.push((x * inv_scale).round().clamp(a, b) as i8);
168 }
169 }
170 quantized
171 }
172 };
173
174 let bytes = Bytes::from_elems(quantized);
175 let layout = Layout::contiguous(shape);
176 let qt = FlexTensor::new(bytes, layout, DType::I8);
177
178 FlexQTensor::new(qt, scheme.with_store(QuantStore::Native), scales)
179 }
180
181 fn dequantize(tensor: QuantizedTensor<Flex>, dtype: FloatDType) -> FloatTensor<Flex> {
182 let shape = tensor.tensor.shape();
183 let qt = tensor.tensor.to_contiguous();
184 let q_data: &[i8] = qt.storage();
185
186 let dequantized = match tensor.scheme.level {
187 QuantLevel::BlockTensor { .. } => {
188 unimplemented!("two-level quantization is not supported yet")
189 }
190 QuantLevel::Tensor => {
191 let scale = tensor.scales[0];
192 q_data
193 .iter()
194 .map(|&x_q| scale * x_q as f32)
195 .collect::<Vec<f32>>()
196 }
197 QuantLevel::Block(block_size) => {
198 let block_elems = block_size.num_elements();
199 q_data
200 .chunks(block_elems)
201 .zip(tensor.scales.iter())
202 .flat_map(|(block, &scale)| block.iter().map(move |&x_q| scale * x_q as f32))
203 .collect::<Vec<f32>>()
204 }
205 };
206
207 let layout = Layout::contiguous(shape);
208 match dtype {
209 FloatDType::F32 | FloatDType::Flex32 => {
210 FlexTensor::new(Bytes::from_elems(dequantized), layout, DType::F32)
211 }
212 FloatDType::F64 => {
213 let data: Vec<f64> = dequantized.iter().map(|&v| v as f64).collect();
214 FlexTensor::new(Bytes::from_elems(data), layout, DType::F64)
215 }
216 FloatDType::F16 => {
217 let data: Vec<f16> = dequantized.iter().map(|&v| f16::from_f32(v)).collect();
218 FlexTensor::new(Bytes::from_elems(data), layout, DType::F16)
219 }
220 FloatDType::BF16 => {
221 let data: Vec<bf16> = dequantized.iter().map(|&v| bf16::from_f32(v)).collect();
222 FlexTensor::new(Bytes::from_elems(data), layout, DType::BF16)
223 }
224 }
225 }
226
227 fn q_to_device(tensor: QuantizedTensor<Flex>, _device: &Device<Flex>) -> QuantizedTensor<Flex> {
228 tensor
229 }
230
231 fn q_reshape(tensor: QuantizedTensor<Flex>, shape: Shape) -> QuantizedTensor<Flex> {
232 block_safe_layout_op(tensor, |t| t.reshape(shape))
233 }
234
235 async fn q_into_data(tensor: QuantizedTensor<Flex>) -> Result<TensorData, ExecutionError> {
236 let shape = tensor.tensor.shape();
237 let scheme = tensor.scheme;
238 let qt = tensor.tensor.to_contiguous();
239 let values: Vec<i8> = qt.storage::<i8>().to_vec();
240
241 Ok(TensorData::quantized(
242 values,
243 shape.to_vec(),
244 scheme,
245 &tensor.scales,
246 ))
247 }
248
249 fn q_swap_dims(
250 tensor: QuantizedTensor<Flex>,
251 dim1: usize,
252 dim2: usize,
253 ) -> QuantizedTensor<Flex> {
254 block_safe_layout_op(tensor, |t| t.transpose(dim1, dim2))
255 }
256
257 fn q_permute(tensor: QuantizedTensor<Flex>, axes: &[usize]) -> QuantizedTensor<Flex> {
258 block_safe_layout_op(tensor, |t| t.permute(axes))
259 }
260
261 fn q_flip(tensor: QuantizedTensor<Flex>, axes: &[usize]) -> QuantizedTensor<Flex> {
262 block_safe_layout_op(tensor, |t| crate::ops::flip::flip(t, axes))
263 }
264
265 fn q_expand(tensor: QuantizedTensor<Flex>, shape: Shape) -> QuantizedTensor<Flex> {
266 block_safe_layout_op(tensor, |t| crate::ops::expand::expand(t, shape))
267 }
268
269 fn q_select(
270 tensor: QuantizedTensor<Flex>,
271 dim: usize,
272 indices: IntTensor<Flex>,
273 ) -> QuantizedTensor<Flex> {
274 match tensor.scheme.level {
275 QuantLevel::BlockTensor { .. } => {
276 unimplemented!("two-level quantization is not supported yet")
277 }
278 QuantLevel::Tensor => FlexQTensor::new(
279 crate::ops::gather_scatter::select::<i8>(tensor.tensor, dim, indices),
280 tensor.scheme,
281 tensor.scales,
282 ),
283 QuantLevel::Block(_) => {
284 let scheme = tensor.scheme;
285 let float_tensor = Flex::dequantize(tensor, FloatDType::F32);
286 let result = crate::ops::gather_scatter::select::<f32>(float_tensor, dim, indices);
287 Flex::quantize_dynamic(result, &scheme)
288 }
289 }
290 }
291
292 fn q_slice(tensor: QuantizedTensor<Flex>, slices: &[Slice]) -> QuantizedTensor<Flex> {
293 block_safe_layout_op(tensor, |t| crate::ops::slice::slice(t, slices))
294 }
295
296 fn q_argmax(
297 tensor: QuantizedTensor<Flex>,
298 dim: usize,
299 out_dtype: burn_std::IntDType,
300 ) -> IntTensor<Flex> {
301 let result = crate::ops::reduce::argmax(tensor.tensor, dim);
302 if result.dtype() != DType::from(out_dtype) {
303 Flex::int_cast(result, out_dtype)
304 } else {
305 result
306 }
307 }
308
309 fn q_argmin(
310 tensor: QuantizedTensor<Flex>,
311 dim: usize,
312 out_dtype: burn_std::IntDType,
313 ) -> IntTensor<Flex> {
314 let result = crate::ops::reduce::argmin(tensor.tensor, dim);
315 if result.dtype() != DType::from(out_dtype) {
316 Flex::int_cast(result, out_dtype)
317 } else {
318 result
319 }
320 }
321
322 fn q_gather(
323 dim: usize,
324 tensor: QuantizedTensor<Flex>,
325 indices: IntTensor<Flex>,
326 ) -> QuantizedTensor<Flex> {
327 match tensor.scheme.level {
328 QuantLevel::BlockTensor { .. } => {
329 unimplemented!("two-level quantization is not supported yet")
330 }
331 QuantLevel::Tensor => FlexQTensor::new(
332 crate::ops::gather_scatter::gather::<i8>(tensor.tensor, dim, indices),
333 tensor.scheme,
334 tensor.scales,
335 ),
336 QuantLevel::Block(_) => {
337 let scheme = tensor.scheme;
338 let float_tensor = Flex::dequantize(tensor, FloatDType::F32);
339 let result = crate::ops::gather_scatter::gather::<f32>(float_tensor, dim, indices);
340 Flex::quantize_dynamic(result, &scheme)
341 }
342 }
343 }
344}
345
346fn block_safe_layout_op(
350 qtensor: FlexQTensor,
351 op: impl FnOnce(FlexTensor) -> FlexTensor,
352) -> FlexQTensor {
353 match qtensor.scheme.level {
354 QuantLevel::BlockTensor { .. } => {
355 unimplemented!("two-level quantization is not supported yet")
356 }
357 QuantLevel::Tensor => FlexQTensor::new(op(qtensor.tensor), qtensor.scheme, qtensor.scales),
358 QuantLevel::Block(_) => {
359 let scheme = qtensor.scheme;
360 let float_tensor = Flex::dequantize(qtensor, FloatDType::F32);
361 let result = op(float_tensor);
362 Flex::quantize_dynamic(result, &scheme)
363 }
364 }
365}
366
367fn validated_scale(scale: f32, param: QuantParam) -> f32 {
374 let scale = scale_to_param(scale, param);
375 if scale > 0.0 && scale.is_finite() {
376 scale
377 } else {
378 f32::MIN_POSITIVE
379 }
380}
381
382#[cfg(test)]
390mod tests {
391 use super::*;
392 use burn_backend::{TensorMetadata, quantization::QuantValue};
393
394 #[test]
395 fn test_quantize_dequantize_roundtrip() {
396 let values = vec![0.0f32, 1.0, 2.0, 3.0, 4.0, 5.0];
398 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [2, 3]));
399
400 let scheme = QuantScheme::default()
401 .with_value(QuantValue::Q8S)
402 .with_store(QuantStore::Native);
403
404 let scale: f32 = 2.0 * 5.0 / 254.0;
408 let scales_tensor = FlexTensor::from_data(TensorData::new(vec![scale], [1]));
409
410 let qparams = QuantizationParametersPrimitive {
411 scales: scales_tensor,
412 };
413
414 let qtensor = Flex::quantize(tensor, &scheme, qparams);
416 assert_eq!(qtensor.tensor.shape().to_vec(), vec![2, 3]);
417 assert_eq!(qtensor.tensor.dtype(), DType::I8);
418
419 let q_vals: &[i8] = qtensor.tensor.storage();
421 assert_eq!(q_vals[0], 0);
423 assert_eq!(q_vals[1], 25);
424 assert_eq!(q_vals[5], 127);
425
426 let result = Flex::dequantize(qtensor, FloatDType::F32);
428 assert_eq!(result.shape().to_vec(), vec![2, 3]);
429 assert_eq!(result.dtype(), DType::F32);
430
431 let result_vals: &[f32] = result.storage();
432 for (orig, deq) in values.iter().zip(result_vals.iter()) {
434 assert!((orig - deq).abs() < 0.05, "orig={orig}, dequantized={deq}");
435 }
436 }
437
438 #[test]
439 fn test_quantize_dequantize_negative_values() {
440 let values = vec![-3.0f32, -1.5, 0.0, 1.5, 3.0];
441 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [5]));
442
443 let scheme = QuantScheme::default()
444 .with_value(QuantValue::Q8S)
445 .with_store(QuantStore::Native);
446
447 let scale: f32 = 2.0 * 3.0 / 254.0;
448 let scales_tensor = FlexTensor::from_data(TensorData::new(vec![scale], [1]));
449
450 let qparams = QuantizationParametersPrimitive {
451 scales: scales_tensor,
452 };
453
454 let qtensor = Flex::quantize(tensor, &scheme, qparams);
455 let result = Flex::dequantize(qtensor, FloatDType::F32);
456 let result_vals: &[f32] = result.storage();
457
458 for (orig, deq) in values.iter().zip(result_vals.iter()) {
459 assert!((orig - deq).abs() < 0.05, "orig={orig}, dequantized={deq}");
460 }
461 }
462
463 #[test]
464 fn test_q_from_data_into_data_roundtrip() {
465 let values = vec![0i8, 25, 51, 76, 102, 127];
467 let scale = 0.03937008f32;
468 let scheme = QuantScheme::default()
469 .with_value(QuantValue::Q8S)
470 .with_store(QuantStore::Native);
471
472 let data = TensorData::quantized(values.clone(), [2, 3], scheme, &[scale]);
473
474 let qtensor = Flex::q_from_data(data, &Default::default());
476 assert_eq!(qtensor.tensor.shape().to_vec(), vec![2, 3]);
477 assert_eq!(qtensor.scales, vec![scale]);
478
479 let float_tensor = Flex::dequantize(qtensor, FloatDType::F32);
481 let result: &[f32] = float_tensor.storage();
482 assert!((result[0]).abs() < 0.01); assert!((result[5] - 5.0).abs() < 0.05); }
485
486 #[test]
487 fn test_quantize_zero_tensor() {
488 let values = vec![0.0f32; 4];
489 let tensor = FlexTensor::from_data(TensorData::new(values, [4]));
490
491 let scheme = QuantScheme::default()
492 .with_value(QuantValue::Q8S)
493 .with_store(QuantStore::Native);
494
495 let scales_tensor = FlexTensor::from_data(TensorData::new(vec![0.0f32], [1]));
497 let qparams = QuantizationParametersPrimitive {
498 scales: scales_tensor,
499 };
500
501 let qtensor = Flex::quantize(tensor, &scheme, qparams);
502 let q_vals: &[i8] = qtensor.tensor.storage();
503 assert_eq!(q_vals, &[0, 0, 0, 0]);
504 }
505
506 #[test]
510 fn test_quantize_dynamic_honors_param_precision() {
511 let values = vec![-3.0f32, -1.5, 0.0, 1.5, 3.0, 4.5];
513 let scheme = QuantScheme::default()
514 .with_value(QuantValue::Q8S)
515 .with_store(QuantStore::Native);
516
517 let quantize_with = |param| {
518 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [2, 3]));
519 Flex::quantize_dynamic(tensor, &scheme.with_param(param)).scales[0]
520 };
521
522 let exact = quantize_with(QuantParam::F32);
523 let coarse = quantize_with(QuantParam::UE4M3);
524
525 assert_ne!(
526 exact, coarse,
527 "UE4M3 scale should differ from the exact f32 scale"
528 );
529 assert_eq!(
530 coarse,
531 scale_to_param(exact, QuantParam::UE4M3),
532 "stored scale should be the exact scale rounded to the param"
533 );
534 }
535
536 #[test]
537 fn test_quantize_dynamic_roundtrip() {
538 let values = vec![-3.0f32, -1.5, 0.0, 1.5, 3.0, 4.5];
539 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [2, 3]));
540
541 let scheme = QuantScheme::default()
542 .with_value(QuantValue::Q8S)
543 .with_store(QuantStore::Native);
544
545 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
546 assert_eq!(qtensor.tensor.shape().to_vec(), vec![2, 3]);
547 assert_eq!(qtensor.scales.len(), 1);
548
549 let expected_scale: f32 = 2.0 * 4.5 / 254.0;
551 assert!(
552 (qtensor.scales[0] - expected_scale).abs() < 1e-6,
553 "scale={}, expected={}",
554 qtensor.scales[0],
555 expected_scale
556 );
557
558 let result = Flex::dequantize(qtensor, FloatDType::F32);
559 let result_vals: &[f32] = result.storage();
560 for (orig, deq) in values.iter().zip(result_vals.iter()) {
561 assert!((orig - deq).abs() < 0.1, "orig={orig}, dequantized={deq}");
562 }
563 }
564
565 #[test]
566 fn test_per_block_quantize_dequantize() {
567 use burn_std::quantization::BlockSize;
568
569 let values = vec![0.0f32, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0];
570 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [8]));
571
572 let block_size = BlockSize::new([4]);
573 let scheme = QuantScheme::default()
574 .with_value(QuantValue::Q8S)
575 .with_level(QuantLevel::Block(block_size))
576 .with_store(QuantStore::Native);
577
578 let scale_1: f32 = 2.0 * 3.0 / 254.0;
581 let scale_2: f32 = 2.0 * 7.0 / 254.0;
582 let scales_tensor = FlexTensor::from_data(TensorData::new(vec![scale_1, scale_2], [2]));
583
584 let qparams = QuantizationParametersPrimitive {
585 scales: scales_tensor,
586 };
587
588 let qtensor = Flex::quantize(tensor, &scheme, qparams);
589 assert_eq!(qtensor.scales.len(), 2);
590
591 let result = Flex::dequantize(qtensor, FloatDType::F32);
592 let result_vals: &[f32] = result.storage();
593
594 for (orig, deq) in values.iter().zip(result_vals.iter()) {
595 assert!((orig - deq).abs() < 0.1, "orig={orig}, dequantized={deq}");
596 }
597 }
598
599 #[test]
600 fn test_quantize_dynamic_block() {
601 use burn_std::quantization::BlockSize;
602
603 let values = vec![-2.0f32, -1.0, 0.0, 1.0, 4.0, 5.0, 6.0, 7.0];
604 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [8]));
605
606 let block_size = BlockSize::new([4]);
607 let scheme = QuantScheme::default()
608 .with_value(QuantValue::Q8S)
609 .with_level(QuantLevel::Block(block_size))
610 .with_store(QuantStore::Native);
611
612 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
613 assert_eq!(qtensor.scales.len(), 2);
614
615 let expected_scale_1: f32 = 2.0 * 2.0 / 254.0;
618 let expected_scale_2: f32 = 2.0 * 7.0 / 254.0;
619 assert!((qtensor.scales[0] - expected_scale_1).abs() < 1e-6);
620 assert!((qtensor.scales[1] - expected_scale_2).abs() < 1e-6);
621
622 let result = Flex::dequantize(qtensor, FloatDType::F32);
623 let result_vals: &[f32] = result.storage();
624 for (orig, deq) in values.iter().zip(result_vals.iter()) {
625 assert!((orig - deq).abs() < 0.1, "orig={orig}, dequantized={deq}");
626 }
627 }
628
629 #[test]
630 fn test_quantize_dynamic_q8f() {
631 let values = vec![-5.0f32, -2.5, 0.0, 2.5, 5.0, 7.5];
633 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [6]));
634
635 let scheme = QuantScheme::default()
636 .with_value(QuantValue::Q8F)
637 .with_store(QuantStore::Native);
638
639 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
640
641 let expected_scale: f32 = 2.0 * 7.5 / 255.0;
644 assert!(
645 (qtensor.scales[0] - expected_scale).abs() < 1e-6,
646 "scale={}, expected={}",
647 qtensor.scales[0],
648 expected_scale
649 );
650
651 let result = Flex::dequantize(qtensor, FloatDType::F32);
652 let result_vals: &[f32] = result.storage();
653 for (orig, deq) in values.iter().zip(result_vals.iter()) {
654 assert!((orig - deq).abs() < 0.1, "orig={orig}, dequantized={deq}");
655 }
656 }
657
658 #[test]
659 fn test_block_quantized_transpose_dequantize() {
660 use burn_std::quantization::BlockSize;
661
662 let values = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
664 let tensor = FlexTensor::from_data(TensorData::new(values, [2, 4]));
665
666 let block_size = BlockSize::new([4]);
667 let scheme = QuantScheme::default()
668 .with_value(QuantValue::Q8S)
669 .with_level(QuantLevel::Block(block_size))
670 .with_store(QuantStore::Native);
671
672 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
673
674 let transposed = Flex::q_swap_dims(qtensor, 0, 1);
676 assert_eq!(transposed.tensor.shape().to_vec(), vec![4, 2]);
677
678 let result = Flex::dequantize(transposed, FloatDType::F32);
679 let result_vals: &[f32] = result.storage();
680
681 let expected = [1.0f32, 5.0, 2.0, 6.0, 3.0, 7.0, 4.0, 8.0];
683 for (exp, deq) in expected.iter().zip(result_vals.iter()) {
684 assert!(
685 (exp - deq).abs() < 0.15,
686 "expected={exp}, dequantized={deq}"
687 );
688 }
689 }
690
691 #[test]
692 fn test_block_quantized_select() {
693 use burn_std::quantization::BlockSize;
694
695 let values = vec![1.0f32, 2.0, 3.0, 4.0, 10.0, 20.0, 30.0, 40.0];
697 let tensor = FlexTensor::from_data(TensorData::new(values, [2, 4]));
698
699 let block_size = BlockSize::new([4]);
700 let scheme = QuantScheme::default()
701 .with_value(QuantValue::Q8S)
702 .with_level(QuantLevel::Block(block_size))
703 .with_store(QuantStore::Native);
704
705 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
706
707 let indices = FlexTensor::from_data(TensorData::new(vec![1i64], [1]));
709 let selected = Flex::q_select(qtensor, 0, indices);
710 assert_eq!(selected.tensor.shape().to_vec(), vec![1, 4]);
711
712 let result = Flex::dequantize(selected, FloatDType::F32);
713 let result_vals: &[f32] = result.storage();
714 let expected = [10.0f32, 20.0, 30.0, 40.0];
715 for (exp, deq) in expected.iter().zip(result_vals.iter()) {
716 assert!((exp - deq).abs() < 0.5, "expected={exp}, dequantized={deq}");
717 }
718 }
719
720 #[test]
721 fn test_block_quantized_flip_dequantize() {
722 use burn_std::quantization::BlockSize;
723
724 let values = vec![1.0f32, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0];
725 let tensor = FlexTensor::from_data(TensorData::new(values, [2, 4]));
726
727 let block_size = BlockSize::new([4]);
728 let scheme = QuantScheme::default()
729 .with_value(QuantValue::Q8S)
730 .with_level(QuantLevel::Block(block_size))
731 .with_store(QuantStore::Native);
732
733 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
734
735 let flipped = Flex::q_flip(qtensor, &[0]);
737 assert_eq!(flipped.tensor.shape().to_vec(), vec![2, 4]);
738
739 let result = Flex::dequantize(flipped, FloatDType::F32);
740 let result_vals: &[f32] = result.storage();
741 let expected = [5.0f32, 6.0, 7.0, 8.0, 1.0, 2.0, 3.0, 4.0];
742 for (exp, deq) in expected.iter().zip(result_vals.iter()) {
743 assert!(
744 (exp - deq).abs() < 0.15,
745 "expected={exp}, dequantized={deq}"
746 );
747 }
748 }
749
750 #[test]
751 fn test_quantize_dynamic_f64_tensor() {
752 use burn_backend::quantization::QuantValue;
753
754 let values = vec![0.0f64, 1.0, 2.0, 3.0, 4.0, 5.0];
755 let tensor = FlexTensor::new(
756 Bytes::from_elems(values),
757 Layout::contiguous([6].into()),
758 DType::F64,
759 );
760 assert_eq!(tensor.dtype(), DType::F64);
761
762 let scheme = QuantScheme::default()
763 .with_value(QuantValue::Q8S)
764 .with_store(QuantStore::Native);
765
766 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
767 assert_eq!(qtensor.tensor.dtype(), DType::I8);
768
769 let result = Flex::dequantize(qtensor, FloatDType::F32);
771 let result_vals: &[f32] = result.storage();
772 let expected = [0.0f32, 1.0, 2.0, 3.0, 4.0, 5.0];
773 for (exp, deq) in expected.iter().zip(result_vals.iter()) {
774 assert!(
775 (exp - deq).abs() < 0.15,
776 "expected={exp}, dequantized={deq}"
777 );
778 }
779 }
780
781 #[test]
782 fn test_dequantize_f64() {
783 let values = vec![0.0f32, 1.0, 2.0, 3.0];
784 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [4]));
785
786 let scheme = QuantScheme::default()
787 .with_value(QuantValue::Q8S)
788 .with_store(QuantStore::Native);
789
790 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
791 let result = Flex::dequantize(qtensor, FloatDType::F64);
792 assert_eq!(result.dtype(), DType::F64);
793 let result_vals: &[f64] = result.storage();
794 for (orig, deq) in values.iter().zip(result_vals.iter()) {
795 assert!(
796 (*orig as f64 - deq).abs() < 0.05,
797 "orig={orig}, dequantized={deq}"
798 );
799 }
800 }
801
802 #[test]
803 fn test_dequantize_f16() {
804 let values = vec![0.0f32, 1.0, 2.0, 3.0];
805 let tensor = FlexTensor::from_data(TensorData::new(values.clone(), [4]));
806
807 let scheme = QuantScheme::default()
808 .with_value(QuantValue::Q8S)
809 .with_store(QuantStore::Native);
810
811 let qtensor = Flex::quantize_dynamic(tensor, &scheme);
812 let result = Flex::dequantize(qtensor, FloatDType::F16);
813 assert_eq!(result.dtype(), DType::F16);
814 let result_vals: &[f16] = result.storage();
815 for (orig, deq) in values.iter().zip(result_vals.iter()) {
816 assert!(
817 (*orig - f32::from(*deq)).abs() < 0.05,
818 "orig={orig}, dequantized={deq}"
819 );
820 }
821 }
822}