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