pub use cubecl_common::quant::scheme::{
BlockScale, BlockSize, QuantMode, QuantScheme, QuantStore, QuantValue, ScaleDtype,
};
pub const QPARAM_ALIGN: usize = core::mem::align_of::<f32>();
use alloc::vec::Vec;
use core::any::TypeId;
use cubecl_common::e4m3;
use num_traits::PrimInt;
use serde::{Deserialize, Serialize};
use crate::{DType, Metadata, Shape, bytes::Bytes};
#[derive(new, Debug, Clone, Copy, PartialEq, Eq, Default, Serialize, Deserialize)]
pub struct QuantConfig {
pub scheme: QuantScheme,
pub propagation: QuantPropagation,
}
#[derive(
Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize, Default,
)]
pub enum QuantAcc {
#[default]
F32,
F16,
BF16,
}
pub enum Calibration {
MinMax,
AbsMean,
}
#[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,
pub global: Option<S>,
}
#[derive(Clone, Debug, PartialEq)]
pub struct DecodedScales {
pub block: Vec<f32>,
pub global: Option<f32>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct QParamTensor {
pub offset_start: usize,
pub offset_end: usize,
pub metadata: Metadata,
pub dtype: DType,
}
pub fn quantizable(scheme: &QuantScheme) -> bool {
if scheme.scale_dtype().round_up(1.0).is_none() {
return false;
}
match (scheme.block_scale(), global_scale_dtype(scheme)) {
(Some(block), Some(global)) => {
block.dtype.max_representable() <= crate::f16::MAX.to_f32() && global == ScaleDtype::F32
}
_ => true,
}
}
pub fn global_scale_dtype(scheme: &QuantScheme) -> Option<ScaleDtype> {
scheme.block_scale().and(scheme.tensor_scale())
}
pub fn params_shape(data_shape: &Shape, scheme: &QuantScheme) -> Shape {
match scheme.block_size() {
None => Shape::new([1]),
Some(block_size) => Shape::from(block_size.num_blocks(data_shape.as_slice())),
}
}
#[derive(Debug, Clone)]
pub struct BlockLayout {
shape: Shape,
block: Vec<u8>,
blocks: Shape,
}
impl BlockLayout {
pub fn new(shape: &Shape, block: &BlockSize) -> Self {
Self {
shape: shape.clone(),
block: block.to_dim_vec(shape.num_dims()),
blocks: Shape::from(block.num_blocks(shape.as_slice())),
}
}
pub fn num_blocks(&self) -> usize {
self.blocks.num_elements()
}
pub fn divides(&self) -> bool {
self.shape
.iter()
.zip(&self.block)
.all(|(&dim, &extent)| dim.is_multiple_of(extent as usize))
}
pub fn block_of(&self, mut index: usize) -> usize {
let mut block = 0;
let mut stride = 1;
for dim in (0..self.shape.num_dims()).rev() {
let coordinate = index % self.shape[dim];
index /= self.shape[dim];
block += coordinate / self.block[dim] as usize * stride;
stride *= self.blocks[dim];
}
block
}
}
pub struct QuantizedBytes {
pub bytes: Bytes,
pub scheme: QuantScheme,
pub shape: Shape,
}
impl QuantizedBytes {
pub fn new<E: bytemuck::CheckedBitPattern + bytemuck::NoUninit>(
value: Vec<E>,
shape: impl Into<Shape>,
scheme: QuantScheme,
scales: &[f32],
global: Option<f32>,
) -> Self {
let shape = shape.into();
assert_eq!(
value.len(),
shape.num_elements(),
"{} quantized values do not fill a tensor of shape {shape:?}",
value.len()
);
if TypeId::of::<E>() != TypeId::of::<i8>() {
panic!("Invalid quantized type");
}
let i8s: Vec<i8> = bytemuck::allocation::cast_vec(value);
let mut bytes = Bytes::from_elems(i8s);
let scales = match scheme.block_size() {
None => &scales[..1],
Some(_) => scales,
};
let scale_bytes = encode_scales(scales, scheme.scale_dtype());
bytes.extend_from_byte_slice_aligned(scale_bytes.as_slice(), QPARAM_ALIGN);
match (global_scale_dtype(&scheme), global) {
(Some(dtype), Some(global)) => {
assert_eq!(
dtype,
ScaleDtype::F32,
"a two-level scheme stores its per-tensor scale as f32, got {scheme:?}"
);
let global_bytes = encode_scales(&[global], dtype);
bytes.extend_from_byte_slice_aligned(global_bytes.as_slice(), QPARAM_ALIGN);
}
(Some(_), None) => panic!("{scheme:?} requires a per-tensor scale"),
(None, Some(_)) => panic!("{scheme:?} does not take a per-tensor scale"),
(None, None) => {}
}
Self {
bytes,
scheme,
shape,
}
}
pub fn num_elements(&self) -> usize {
self.shape.num_elements()
}
pub fn into_vec_i8(self) -> (Vec<i8>, DecodedScales) {
let scheme = self.scheme;
let (values, (qparams, num_params)) = self.split_values_off();
let global_bytes = global_scale_size(&scheme);
let block_end = qparams
.len()
.checked_sub(global_bytes)
.expect("quantized parameter buffer is shorter than the scheme's global scale");
let block_start = block_end
.checked_sub(scale_size(scheme.scale_dtype()) * num_params)
.expect("quantized parameter buffer is shorter than the scheme's block scales");
let block = decode_scales(&qparams[block_start..block_end], scheme.scale_dtype());
let global =
global_scale_dtype(&scheme).map(|dtype| decode_scales(&qparams[block_end..], dtype)[0]);
(values, DecodedScales { block, global })
}
fn split_i8_values(self, scale_bytes: usize) -> (Vec<i8>, Vec<u8>) {
let mut values = read_bytes_to_i8(self.bytes);
let values_end = values
.len()
.checked_sub(scale_bytes)
.expect("quantized tensor data is shorter than its scheme's parameters");
let qparams = values.split_off(values_end);
(values, bytemuck::cast_vec(qparams))
}
fn split_values_off(self) -> (Vec<i8>, (Vec<u8>, usize)) {
let num_params = params_shape(&self.shape, &self.scheme).num_elements();
let scale_bytes =
scale_size(self.scheme.scale_dtype()) * num_params + global_scale_size(&self.scheme);
if let QuantStore::PackedU32(packed_dim) = self.scheme.store {
assert_eq!(
packed_dim, 0,
"Packing must be on innermost dimension for splitting off values"
);
}
let (values, qparams) = match self.scheme.store {
QuantStore::Native => self.split_i8_values(scale_bytes),
QuantStore::PackedU32(_) => match self.scheme.value {
QuantValue::Q8F | QuantValue::Q8S => self.split_i8_values(scale_bytes),
QuantValue::Q4F | QuantValue::Q4S | QuantValue::Q2F | QuantValue::Q2S => {
let split_at =
self.bytes.len().checked_sub(scale_bytes).expect(
"quantized tensor data is shorter than its scheme's parameters",
);
let qparams = self.bytes[split_at..].to_vec();
let values = bytemuck::cast_slice::<_, u32>(&self.bytes[..split_at]);
let values = unpack_q_to_i8s(values, self.num_elements(), &self.scheme.value);
(values, qparams)
}
QuantValue::E4M3 | QuantValue::E5M2 | QuantValue::E2M1 => {
unimplemented!("Not yet supported")
}
},
QuantStore::PackedNative(_) => unimplemented!("Not yet supported"),
};
(values, (qparams, num_params))
}
}
pub fn scale_to_dtype(scale: f32, dtype: ScaleDtype) -> f32 {
dtype
.round_up(scale)
.expect("UE8M0 scales are not yet supported")
}
pub fn global_scale_size(scheme: &QuantScheme) -> usize {
global_scale_dtype(scheme).map_or(0, scale_size)
}
fn storage_elements(scheme: &QuantScheme, shape: &Shape) -> usize {
let num_quants = scheme.num_quants();
match scheme.store {
QuantStore::PackedU32(packed_dim) | QuantStore::PackedNative(packed_dim)
if num_quants > 1 && !shape.is_empty() =>
{
let packed_dim = shape.num_dims() - packed_dim - 1;
let mut storage = shape.clone();
storage[packed_dim] = storage[packed_dim].div_ceil(num_quants);
storage.num_elements()
}
_ => shape.num_elements().div_ceil(num_quants),
}
}
pub fn quantized_data_len(scheme: &QuantScheme, shape: &Shape) -> usize {
let value_bytes = storage_elements(scheme, shape) * scheme.size_bits_stored().div_ceil(8);
let num_params = params_shape(shape, scheme).num_elements();
let scale_bytes = num_params * scale_size(scheme.scale_dtype());
value_bytes + scale_bytes + global_scale_size(scheme)
}
pub fn scale_size(dtype: ScaleDtype) -> usize {
match dtype {
ScaleDtype::F32 => 4,
ScaleDtype::F16 | ScaleDtype::BF16 => 2,
ScaleDtype::UE8M0 | ScaleDtype::UE4M3 => 1,
}
}
fn decode_scales(bytes: &[u8], dtype: ScaleDtype) -> Vec<f32> {
match dtype {
ScaleDtype::F32 => bytes
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_ne_bytes([c[0], c[1], c[2], c[3]]))
.collect(),
ScaleDtype::F16 => bytes
.as_chunks::<2>()
.0
.iter()
.map(|c| crate::f16::from_ne_bytes([c[0], c[1]]).to_f32())
.collect(),
ScaleDtype::BF16 => bytes
.as_chunks::<2>()
.0
.iter()
.map(|c| crate::bf16::from_ne_bytes([c[0], c[1]]).to_f32())
.collect(),
ScaleDtype::UE4M3 => bytes.iter().map(|b| e4m3::from_bits(*b).to_f32()).collect(),
ScaleDtype::UE8M0 => unimplemented!("UE8M0 scales are not yet supported"),
}
}
fn encode_scales(scales: &[f32], dtype: ScaleDtype) -> Vec<u8> {
match dtype {
ScaleDtype::F32 => scales.iter().flat_map(|s| s.to_ne_bytes()).collect(),
ScaleDtype::F16 => scales
.iter()
.flat_map(|s| crate::f16::from_f32(*s).to_ne_bytes())
.collect(),
ScaleDtype::BF16 => scales
.iter()
.flat_map(|s| crate::bf16::from_f32(*s).to_ne_bytes())
.collect(),
ScaleDtype::UE4M3 => scales
.iter()
.map(|s| e4m3::from_f32(*s).to_bits())
.collect(),
ScaleDtype::UE8M0 => unimplemented!("UE8M0 scales are not yet supported"),
}
}
fn read_bytes_to_i8(bytes: Bytes) -> Vec<i8> {
match bytes.try_into_vec::<i8>() {
Ok(val) => val,
Err(bytes) => unsafe { core::mem::transmute::<Vec<u8>, Vec<i8>>(bytes.to_vec()) },
}
}
pub fn pack_i8s_to_u32s(values: Vec<i8>) -> Vec<u32> {
#[cfg(target_endian = "big")]
{
values
.chunks(4)
.map(|x| {
x.iter()
.enumerate()
.fold(0u32, |acc, (i, x)| acc | (*x as u32 & 0xFF) << (i * 8))
})
.collect()
}
#[cfg(target_endian = "little")]
{
let mut values = values;
let remainder = values.len() % 4;
if remainder != 0 {
values.extend(core::iter::repeat_n(0, 4 - remainder));
}
let len = values.len() / 4;
let capacity = values.capacity() / 4;
let mut values = core::mem::ManuallyDrop::new(values);
let ptr = values.as_mut_ptr() as *mut u32;
unsafe { Vec::from_raw_parts(ptr, len, capacity) }
}
}
pub(crate) fn unpack_q_to_i8s<Q: PrimInt>(
values: &[Q],
numel: usize,
value: &QuantValue,
) -> Vec<i8> {
let size_store = size_of::<Q>() * 8;
let size_quant = value.size_bits();
let num_quants = size_store / size_quant;
let mask = Q::from((1 << size_quant) - 1).unwrap();
let sign_shift = 8 - size_quant; values
.iter()
.enumerate()
.flat_map(|(i, &packed)| {
let n = core::cmp::min(num_quants, numel - i * num_quants);
(0..n).map(move |i| {
let raw = (packed >> (i * size_quant) & mask).to_u8().unwrap();
((raw << sign_shift) as i8) >> sign_shift
})
})
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
#[test]
fn should_pack_i8s_to_u32() {
let packed = pack_i8s_to_u32s(vec![-128, 2, -3, 127]);
assert_eq!(packed, vec![2147287680]);
}
#[test]
fn should_pack_i8s_to_u32_padded() {
let packed = pack_i8s_to_u32s(vec![-128, 2, -3, 127, 55]);
let packed_padded = pack_i8s_to_u32s(vec![-128, 2, -3, 127, 55, 0, 0, 0]);
assert_eq!(packed, vec![2147287680, 55]);
assert_eq!(packed, packed_padded);
}
#[test]
fn should_unpack_u32s_to_i8s() {
let unpacked = unpack_q_to_i8s(&[2147287680u32], 4, &QuantValue::Q8S);
assert_eq!(unpacked, vec![-128, 2, -3, 127]);
}
#[test]
fn should_unpack_u32s_to_i8s_padded() {
let unpacked = unpack_q_to_i8s(&[55u32], 1, &QuantValue::Q8S);
assert_eq!(unpacked, vec![55]);
}
#[test]
fn should_unpack_u32s_to_i8s_arange() {
let unpacked = unpack_q_to_i8s(
&[
0u32, 286331136, 286331153, 572657937, 572662306, 857874978, 858993459, 858993459,
1145324612, 1145324612, 1431655748, 1431655765, 1717982549, 1717986918, 2003199590,
2004318071,
],
128,
&QuantValue::Q4S,
);
assert_eq!(
unpacked,
vec![
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1,
2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3, 3, 3,
3, 3, 3, 3, 3, 3, 3, 3, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 4, 5, 5,
5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 5, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6, 6,
6, 6, 6, 6, 6, 6, 7, 7, 7, 7, 7, 7, 7, 7, 7, 7
]
);
}
#[test]
fn should_pack_unpack_quantization_parameters_per_tensor_symmetric() {
let scale = 0.03937008;
let values = vec![0i8, 25, 51, 76, 102, 127];
let q_bytes = QuantizedBytes::new(
values.clone(),
[2, 3],
QuantScheme::default()
.with_value(QuantValue::Q8S)
.with_store(QuantStore::Native),
&[scale],
None,
);
let (q_values, qparams) = q_bytes.into_vec_i8();
assert_eq!(qparams.block, vec![scale]);
assert_eq!(q_values, values);
}
#[test]
fn scale_to_dtype_survives_the_codec() {
let scales = [0.5f32, 0.3, 1.0 / 3.0, 500.0, 1e-3, 7.7e-4];
for dtype in [
ScaleDtype::F32,
ScaleDtype::F16,
ScaleDtype::BF16,
ScaleDtype::UE4M3,
] {
let rounded: Vec<f32> = scales.iter().map(|s| scale_to_dtype(*s, dtype)).collect();
let via_codec = decode_scales(&encode_scales(&rounded, dtype), dtype);
assert_eq!(
rounded, via_codec,
"the codec moves a scale {dtype:?} can already represent"
);
for (scale, rounded) in scales.iter().zip(&rounded).filter(|(s, _)| **s < 500.0) {
assert!(
rounded >= scale,
"{scale} rounded down to {rounded} for {dtype:?}"
);
}
}
}
#[test]
fn should_pack_unpack_two_level_scales() {
let block_scales = [0.5f32, 0.125];
let global = 3.0f32;
let values = vec![0i8, 25, 51, 76, 102, 127, -128, -1];
let scheme = QuantScheme::default()
.with_value(QuantValue::Q8S)
.with_store(QuantStore::Native)
.per_block([4], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::F32);
let q_bytes = QuantizedBytes::new(values.clone(), [8], scheme, &block_scales, Some(global));
assert_eq!(q_bytes.bytes.len(), 8 + 2 + 4);
let (q_values, scales) = q_bytes.into_vec_i8();
assert_eq!(q_values, values);
assert_eq!(scales.block, block_scales);
assert_eq!(scales.global, Some(global));
}
#[test]
#[should_panic(expected = "requires a per-tensor scale")]
fn two_level_scheme_without_a_global_scale_is_rejected() {
let scheme = QuantScheme::default()
.with_value(QuantValue::Q8S)
.with_store(QuantStore::Native)
.per_block([4], ScaleDtype::F32)
.per_tensor(ScaleDtype::F32);
QuantizedBytes::new(vec![0i8; 8], [8], scheme, &[0.5, 0.125], None);
}
#[test]
#[should_panic(expected = "stores its per-tensor scale as f32")]
fn a_narrower_per_tensor_scale_is_rejected() {
let scheme = QuantScheme::default()
.with_value(QuantValue::Q8S)
.with_store(QuantStore::Native)
.per_block([4], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::F16);
QuantizedBytes::new(vec![0i8; 8], [8], scheme, &[0.5, 0.125], Some(3.0));
}
#[test]
fn quantizable_declines_what_no_backend_can_store() {
assert!(quantizable(&QuantScheme::default()));
assert!(quantizable(
&QuantScheme::default().per_block([4], ScaleDtype::F16)
));
assert!(quantizable(
&QuantScheme::default()
.per_block([4], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::F32)
));
assert!(!quantizable(
&QuantScheme::default().per_block([4], ScaleDtype::UE8M0)
));
assert!(!quantizable(
&QuantScheme::default().per_tensor(ScaleDtype::UE8M0)
));
assert!(!quantizable(
&QuantScheme::default()
.per_block([4], ScaleDtype::F32)
.per_tensor(ScaleDtype::F32)
));
assert!(!quantizable(
&QuantScheme::default()
.per_block([4], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::BF16)
));
}
#[test]
fn encoded_scale_width_matches_scale_size() {
let scales = [0.5f32, 0.25, 0.125];
for dtype in [
ScaleDtype::F32,
ScaleDtype::F16,
ScaleDtype::BF16,
ScaleDtype::UE4M3,
] {
assert_eq!(
encode_scales(&scales, dtype).len(),
scale_size(dtype) * scales.len(),
"encoded width disagrees with scale_size for {dtype:?}"
);
}
}
#[test]
fn should_pack_unpack_ue4m3_block_scales() {
let scales = [0.5f32, 0.125];
let values = vec![0i8, 25, 51, 76, 102, 127, -128, -1];
let q_bytes = QuantizedBytes::new(
values.clone(),
[8],
QuantScheme::default()
.with_value(QuantValue::Q8S)
.with_store(QuantStore::Native)
.per_block([4], ScaleDtype::UE4M3),
&scales,
None,
);
let (q_values, qparams) = q_bytes.into_vec_i8();
assert_eq!(qparams.block, scales);
assert_eq!(q_values, values);
}
}