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