use alloc::vec::Vec;
use ruda_core::tensor::{DType, QTensorPrimitive, TensorMetadata, quantization::{QuantStore, params_shape}};
use ruda_core::tensor::{QuantScheme, Shape};
use ruda_core::tensor::host::HostTensor;
#[derive(Clone)]
pub struct HostQTensor {
pub(crate) tensor: HostTensor,
pub(crate) scheme: QuantScheme,
pub(crate) scales: Vec<f32>,
}
impl HostQTensor {
pub fn new(tensor: HostTensor, scheme: QuantScheme, scales: Vec<f32>) -> Self {
assert_eq!(
tensor.dtype(),
DType::I8,
"quantized tensor must store i8 data, got {:?}",
tensor.dtype()
);
assert_eq!(
scales.len(),
params_shape(&tensor.shape(), scheme.level).num_elements(),
"quantized scale count must match the parameter shape"
);
Self {
tensor,
scheme,
scales,
}
}
pub fn tensor(&self) -> &HostTensor {
&self.tensor
}
pub fn scales(&self) -> &[f32] {
&self.scales
}
}
impl QTensorPrimitive for HostQTensor {
fn scheme(&self) -> &QuantScheme {
&self.scheme
}
fn default_scheme() -> QuantScheme {
QuantScheme::default().with_store(QuantStore::Native)
}
}
impl TensorMetadata for HostQTensor {
fn dtype(&self) -> DType {
DType::QFloat(self.scheme)
}
fn shape(&self) -> Shape {
self.tensor.shape()
}
fn rank(&self) -> usize {
self.tensor.rank()
}
}
impl core::fmt::Debug for HostQTensor {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("FlexQTensor")
.field("tensor", &self.tensor)
.field("scheme", &self.scheme)
.field("scales", &self.scales)
.finish()
}
}