use super::QuantLevel;
use crate::tensor::{DType, Shape, metadata::Metadata};
use serde::{Deserialize, Serialize};
#[derive(
Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default,
)]
pub enum QuantAcc {
#[default]
F32,
F16,
BF16,
}
#[derive(
Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default,
)]
pub enum QuantPropagation {
Propagate,
#[default]
Inhibit,
}
#[derive(Clone, Debug)]
pub struct QParams<S> {
pub scales: S,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct QParamTensor {
pub offset_start: usize,
pub offset_end: usize,
pub metadata: Metadata,
pub dtype: DType,
}
pub fn params_shape(data_shape: &Shape, level: QuantLevel) -> Shape {
match level {
QuantLevel::Tensor => Shape::new([1]),
QuantLevel::Block(block_size) => {
let mut params_shape = data_shape.clone();
let block_size = block_size.to_dim_vec(data_shape.num_dims());
for (shape, block_size) in params_shape.iter_mut().zip(block_size) {
*shape = (*shape).div_ceil(block_size as usize);
}
params_shape
}
}
}