Skip to main content

burn_flex/ops/
qtensor.rs

1//! Quantized tensor operations for the Flex backend.
2
3use 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        // Use native storage since we've unpacked to i8
43        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                // Pass 1: find alpha = max(|min|, |max|)
61                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                // Pass 2: quantize
72                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                    // Find alpha for this block
93                    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                    // Quantize this block
105                    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        // Extract and validate scales from the qparams tensor. The scales tensor
131        // shares its dtype with the float element type, which can be any of
132        // f32/f64/f16/bf16, so we normalise via float_storage_as_f32 instead of
133        // assuming f32 storage.
134        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
346/// Apply a layout operation to a quantized tensor.
347/// For block-quantized tensors, dequantizes and requantizes to preserve
348/// correct scale-to-block mapping.
349fn 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
367/// Round the scale to the scheme's param dtype, then ensure it is finite and nonzero to avoid
368/// division by zero or NaN propagation.
369///
370/// Rounding comes first, and only an exactly zero scale is replaced: subnormals of the param are
371/// representable and carry real information for a small tensor, so substituting them would throw
372/// away what rounding up preserved.
373fn 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// Tests kept here exercise flex-specific behavior: quantization scheme
383// roundtrips, per-block / dynamic quantization, block-quantized layout
384// ops (transpose / select / flip dequantize), and f16/f64 dequantize
385// dtype paths. Plain layout-preservation / select / slice / argmax /
386// argmin / gather tests are covered generically in
387// crates/burn-backend-tests/tests/tensor/float/quantization/ops/extended/
388// so they run on every backend.
389#[cfg(test)]
390mod tests {
391    use super::*;
392    use burn_backend::{TensorMetadata, quantization::QuantValue};
393
394    #[test]
395    fn test_quantize_dequantize_roundtrip() {
396        // Create a float tensor
397        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        // Compute scale: symmetric, so scale = 2 * max(|min|, |max|) / (b - a)
405        // max_abs = 5.0, range = 127 - (-127) = 254
406        // scale = 2 * 5.0 / 254 = 0.03937008
407        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        // Quantize
415        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        // Check quantized values
420        let q_vals: &[i8] = qtensor.tensor.storage();
421        // 0 / 0.03937 = 0, 1 / 0.03937 = 25.4 -> 25, etc.
422        assert_eq!(q_vals[0], 0);
423        assert_eq!(q_vals[1], 25);
424        assert_eq!(q_vals[5], 127);
425
426        // Dequantize
427        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        // Values should be approximately equal (quantization introduces small errors)
433        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        // Create quantized TensorData the standard way
466        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        // Load into FlexQTensor
475        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        // Dequantize and check values
480        let float_tensor = Flex::dequantize(qtensor, FloatDType::F32);
481        let result: &[f32] = float_tensor.storage();
482        assert!((result[0]).abs() < 0.01); // 0 * scale ~ 0
483        assert!((result[5] - 5.0).abs() < 0.05); // 127 * scale ~ 5.0
484    }
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        // Scale of 0 should be handled gracefully
496        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    /// The scale a narrow param can represent is coarser than the exact one, so the stored scale
507    /// must reflect the param rather than staying at full `f32` precision. Before scales were
508    /// rounded, every param produced byte-identical results here.
509    #[test]
510    fn test_quantize_dynamic_honors_param_precision() {
511        // 4.5 / 127 is not representable in e4m3, so rounding is observable.
512        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        // Scale should be 2 * 4.5 / 254
550        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        // Block 1: [0, 1, 2, 3] -> max_abs=3, scale = 6/254
579        // Block 2: [4, 5, 6, 7] -> max_abs=7, scale = 14/254
580        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        // Block 1: [-2, -1, 0, 1] -> alpha=2, scale = 4/254
616        // Block 2: [4, 5, 6, 7] -> alpha=7, scale = 14/254
617        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        // Q8F uses asymmetric range [-128, 127]
632        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        // Q8F range: [-128, 127], so range = 255
642        // alpha = 7.5, scale = 2 * 7.5 / 255
643        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        // 2x4 tensor, 2 blocks of 4
663        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        // Transpose to [4, 2], then dequantize
675        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        // Original [[1,2,3,4],[5,6,7,8]] transposed to [[1,5],[2,6],[3,7],[4,8]]
682        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        // 2x4 tensor, 2 blocks of 4
696        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        // Select row 1 -> [10, 20, 30, 40]
708        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        // Flip along axis 0: [[5,6,7,8],[1,2,3,4]]
736        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        // Dequantize and verify round-trip accuracy
770        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}