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        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
22/// The blocks over `shape`, which must be a whole number of blocks along every axis.
23fn 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
32/// The largest magnitude in each block of `values`, laid out as `blocks`.
33fn 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        // Use native storage since we've unpacked to i8
62        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                // Pass 2: quantize
111                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        // Extract and validate scales from the qparams tensor. The scales tensor
154        // shares its dtype with the float element type, which can be any of
155        // f32/f64/f16/bf16, so we normalise via float_storage_as_f32 instead of
156        // assuming f32 storage.
157        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
373/// Apply a layout operation to a quantized tensor. A block-quantized tensor is dequantized,
374/// moved, and requantized under `scheme`, which the caller has already rewritten to follow the
375/// move (a permuted tensor's blocks are permuted with it, as every other backend keeps them).
376fn 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
396/// Unrounded; callers round separately.
397fn 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
402/// Only an exactly zero scale is replaced: the dtype's subnormals carry real information for a
403/// small tensor, and the replacement goes back through the dtype because `f32::MIN_POSITIVE`
404/// would itself encode to zero in a narrow one.
405fn 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// Tests kept here exercise flex-specific behavior: quantization scheme
415// roundtrips, per-block / dynamic quantization, block-quantized layout
416// ops (transpose / select / flip dequantize), and f16/f64 dequantize
417// dtype paths. Plain layout-preservation / select / slice / argmax /
418// argmin / gather tests are covered generically in
419// crates/burn-backend-tests/tests/tensor/float/quantization/ops/extended/
420// so they run on every backend.
421#[cfg(test)]
422mod tests {
423    use super::*;
424    use burn_backend::{TensorMetadata, quantization::QuantValue};
425
426    #[test]
427    fn test_quantize_dequantize_roundtrip() {
428        // Create a float tensor
429        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        // Compute scale: symmetric, so scale = 2 * max(|min|, |max|) / (b - a)
437        // max_abs = 5.0, range = 127 - (-127) = 254
438        // scale = 2 * 5.0 / 254 = 0.03937008
439        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        // Quantize
448        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        // Check quantized values
453        let q_vals: &[i8] = qtensor.tensor.storage();
454        // 0 / 0.03937 = 0, 1 / 0.03937 = 25.4 -> 25, etc.
455        assert_eq!(q_vals[0], 0);
456        assert_eq!(q_vals[1], 25);
457        assert_eq!(q_vals[5], 127);
458
459        // Dequantize
460        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        // Values should be approximately equal (quantization introduces small errors)
466        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        // Create quantized TensorData the standard way
500        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        // Load into FlexQTensor
509        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        // Dequantize and check values
514        let float_tensor = Flex::dequantize(qtensor, FloatDType::F32);
515        let result: &[f32] = float_tensor.storage();
516        assert!((result[0]).abs() < 0.01); // 0 * scale ~ 0
517        assert!((result[5] - 5.0).abs() < 0.05); // 127 * scale ~ 5.0
518    }
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        // Scale of 0 should be handled gracefully
530        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    /// The scale a narrow dtype can represent is coarser than the exact one, so the stored scale
542    /// must reflect the dtype rather than staying at full `f32` precision. Before scales were
543    /// rounded, every scale dtype produced byte-identical results here.
544    #[test]
545    fn test_quantize_dynamic_honors_scale_dtype_precision() {
546        // 4.5 / 127 is not representable in e4m3, so rounding is observable.
547        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        // Scale should be 2 * 4.5 / 254
585        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        // Block 1: [0, 1, 2, 3] -> max_abs=3, scale = 6/254
614        // Block 2: [4, 5, 6, 7] -> max_abs=7, scale = 14/254
615        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        // Block 1: [-2, -1, 0, 1] -> alpha=2, scale = 4/254
652        // Block 2: [4, 5, 6, 7] -> alpha=7, scale = 14/254
653        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        // Q8F uses asymmetric range [-128, 127]
668        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        // Q8F range: [-128, 127], so range = 255
678        // alpha = 7.5, scale = 2 * 7.5 / 255
679        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        // 2x4 tensor, 2 blocks of 4
699        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        // Transpose to [4, 2], then dequantize
711        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        // Original [[1,2,3,4],[5,6,7,8]] transposed to [[1,5],[2,6],[3,7],[4,8]]
718        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        // 2x4 tensor, 2 blocks of 4
732        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        // Select row 1 -> [10, 20, 30, 40]
744        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        // Flip along axis 0: [[5,6,7,8],[1,2,3,4]]
772        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        // Dequantize and verify round-trip accuracy
806        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}