ruprim_host/quantization/
tensor.rs1use 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#[derive(Clone)]
13pub struct HostQTensor {
14 pub(crate) tensor: HostTensor,
16 pub(crate) scheme: QuantScheme,
18 pub(crate) scales: Vec<f32>,
20}
21
22impl HostQTensor {
23 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 pub fn tensor(&self) -> &HostTensor {
47 &self.tensor
48 }
49
50 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}