Skip to main content

ruprim_host/quantization/
tensor.rs

1use alloc::vec::Vec;
2
3use ruda_core::tensor::{DType, QTensorPrimitive, TensorMetadata, quantization::{QuantStore, params_shape}};
4use ruda_core::tensor::{QuantScheme, Shape};
5
6use ruda_core::tensor::host::HostTensor;
7
8/// Quantized tensor for the Flex backend.
9///
10/// Stores quantized i8 values in the tensor and keeps scales separately
11/// for efficient dequantization without reparsing bytes.
12#[derive(Clone)]
13pub struct HostQTensor {
14    /// The underlying quantized data (stored as i8).
15    pub(crate) tensor: HostTensor,
16    /// Quantization scheme.
17    pub(crate) scheme: QuantScheme,
18    /// Per-tensor or per-block scale factors.
19    pub(crate) scales: Vec<f32>,
20}
21
22impl HostQTensor {
23    /// Create a new quantized tensor.
24    ///
25    /// The tensor must store i8 data and scales must match the quantization parameter shape.
26    pub fn new(tensor: HostTensor, scheme: QuantScheme, scales: Vec<f32>) -> Self {
27        assert_eq!(
28            tensor.dtype(),
29            DType::I8,
30            "quantized tensor must store i8 data, got {:?}",
31            tensor.dtype()
32        );
33        assert_eq!(
34            scales.len(),
35            params_shape(&tensor.shape(), scheme.level).num_elements(),
36            "quantized scale count must match the parameter shape"
37        );
38        Self {
39            tensor,
40            scheme,
41            scales,
42        }
43    }
44
45    /// Get the underlying tensor.
46    pub fn tensor(&self) -> &HostTensor {
47        &self.tensor
48    }
49
50    /// Get the quantization scales.
51    pub fn scales(&self) -> &[f32] {
52        &self.scales
53    }
54}
55
56impl QTensorPrimitive for HostQTensor {
57    fn scheme(&self) -> &QuantScheme {
58        &self.scheme
59    }
60
61    fn default_scheme() -> QuantScheme {
62        QuantScheme::default().with_store(QuantStore::Native)
63    }
64}
65
66impl TensorMetadata for HostQTensor {
67    fn dtype(&self) -> DType {
68        DType::QFloat(self.scheme)
69    }
70
71    fn shape(&self) -> Shape {
72        self.tensor.shape()
73    }
74
75    fn rank(&self) -> usize {
76        self.tensor.rank()
77    }
78}
79
80impl core::fmt::Debug for HostQTensor {
81    fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
82        f.debug_struct("FlexQTensor")
83            .field("tensor", &self.tensor)
84            .field("scheme", &self.scheme)
85            .field("scales", &self.scales)
86            .finish()
87    }
88}