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