use alloc::vec;
use alloc::vec::Vec;
use core::{default::Default, ops::Deref};
use serde::{Deserialize, Serialize};
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub struct QuantScheme {
pub value: QuantValue,
pub store: QuantStore,
pub mode: QuantMode,
tensor: Option<ScaleDtype>,
block: Option<BlockScale>,
}
impl Default for QuantScheme {
fn default() -> Self {
Self {
value: QuantValue::Q8F,
store: QuantStore::PackedU32(0),
mode: QuantMode::Symmetric,
tensor: None,
block: None,
}
}
}
impl QuantScheme {
pub fn with_mode(mut self, mode: QuantMode) -> Self {
self.mode = mode;
self
}
pub fn with_value(mut self, value: QuantValue) -> Self {
self.value = value;
self
}
pub fn with_store(mut self, store: QuantStore) -> Self {
self.store = store;
self
}
pub fn per_tensor(mut self, dtype: ScaleDtype) -> Self {
self.tensor = Some(dtype);
self
}
pub fn per_block(mut self, block: impl AsRef<[u8]>, dtype: ScaleDtype) -> Self {
self.block = Some(BlockScale {
size: BlockSize::new(block),
dtype,
});
self
}
pub fn tensor_scale(&self) -> Option<ScaleDtype> {
if self.tensor.is_none() && self.block.is_none() {
return Some(ScaleDtype::F32);
}
self.tensor
}
pub fn block_scale(&self) -> Option<BlockScale> {
self.block
}
pub fn num_levels(&self) -> usize {
self.block_scale().is_some() as usize + self.tensor_scale().is_some() as usize
}
pub fn scale_dtype(&self) -> ScaleDtype {
self.block
.map(|block| block.dtype)
.or(self.tensor)
.unwrap_or(ScaleDtype::F32)
}
pub fn block_size(&self) -> Option<BlockSize> {
self.block.map(|block| block.size)
}
pub fn swap_block_dims(&mut self, rank: usize, dim0: usize, dim1: usize) {
let mut axes: Vec<usize> = (0..rank).collect();
axes.swap(dim0, dim1);
self.permute_block_dims(rank, &axes);
}
pub fn permute_block_dims(&mut self, rank: usize, axes: &[usize]) {
if let Some(block) = &mut self.block {
let dims = block.size.to_dim_vec(rank);
let permuted: Vec<u8> = axes.iter().map(|&axis| dims[axis]).collect();
block.size = BlockSize::new(permuted);
}
}
pub fn size_bits_stored(&self) -> usize {
self.store.size_bits(&self.value)
}
pub fn size_bits_value(&self) -> usize {
self.value.size_bits()
}
pub fn num_quants(&self) -> usize {
self.size_bits_stored() / self.value.size_bits()
}
pub fn native_packing(&self) -> usize {
self.value.native_packing()
}
pub fn packing_dim(&self) -> Option<usize> {
self.store.packing_dim()
}
pub fn swap_packing_dim(&mut self, dim0: usize, dim1: usize) {
if let QuantStore::PackedU32(packed_dim) | QuantStore::PackedNative(packed_dim) =
&mut self.store
{
if *packed_dim == dim0 {
*packed_dim = dim1;
} else if *packed_dim == dim1 {
*packed_dim = dim0;
}
}
}
}
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub struct BlockScale {
pub size: BlockSize,
pub dtype: ScaleDtype,
}
impl ScaleDtype {
pub fn max_representable(&self) -> f32 {
match self {
ScaleDtype::F32 => f32::MAX,
ScaleDtype::F16 => half::f16::MAX.to_f32(),
ScaleDtype::BF16 => half::bf16::MAX.to_f32(),
ScaleDtype::UE8M0 => f32::from_bits(0x7F00_0000), ScaleDtype::UE4M3 => 448.0,
}
}
pub fn round_up(&self, scale: f32) -> Option<f32> {
match self {
ScaleDtype::F32 => {
return Some(scale);
}
ScaleDtype::UE8M0 => {
return None;
}
_ => {}
}
if scale.is_nan() {
return Some(scale);
}
debug_assert!(scale >= 0.0, "a quantization scale is never negative");
let max = self.max_representable();
if scale >= max {
return Some(max);
}
let grid = self.f32_grid();
if let Some(subnormals) = grid.subnormals
&& scale < subnormals.min_normal
{
return Some(num_traits::Float::ceil(scale / subnormals.spacing) * subnormals.spacing);
}
Some(f32::from_bits(
(scale.to_bits() + grid.round_up_bias()) & grid.truncate_mask(),
))
}
pub fn f32_grid(&self) -> F32Grid {
const fn bit_step(mantissa_digits: u32) -> u32 {
1 << (f32::MANTISSA_DIGITS - mantissa_digits)
}
match self {
ScaleDtype::F16 => F32Grid {
bit_step: bit_step(half::f16::MANTISSA_DIGITS),
subnormals: Some(SubnormalRange {
min_normal: half::f16::MIN_POSITIVE.to_f32(),
spacing: half::f16::MIN_POSITIVE_SUBNORMAL.to_f32(),
}),
},
ScaleDtype::BF16 => F32Grid {
bit_step: bit_step(half::bf16::MANTISSA_DIGITS),
subnormals: None,
},
ScaleDtype::UE4M3 => F32Grid {
bit_step: bit_step(4),
subnormals: Some(SubnormalRange {
min_normal: 0.015625, spacing: 0.001953125, }),
},
ScaleDtype::F32 => {
unimplemented!("F32 is the grid, it has no narrower one to round onto")
}
ScaleDtype::UE8M0 => unimplemented!("UE8M0 scales are not yet supported"),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct F32Grid {
pub bit_step: u32,
pub subnormals: Option<SubnormalRange>,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct SubnormalRange {
pub min_normal: f32,
pub spacing: f32,
}
impl F32Grid {
pub fn truncate_mask(&self) -> u32 {
!(self.bit_step - 1)
}
pub fn round_up_bias(&self) -> u32 {
self.bit_step - 1
}
}
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum QuantValue {
Q8F,
E5M2,
E4M3,
Q4F,
E2M1,
Q2F,
Q8S,
Q4S,
Q2S,
}
impl QuantValue {
pub fn size_bits(&self) -> usize {
match self {
QuantValue::Q8F | QuantValue::Q8S | QuantValue::E4M3 | QuantValue::E5M2 => 8,
QuantValue::Q4F | QuantValue::Q4S | QuantValue::E2M1 => 4,
QuantValue::Q2F | QuantValue::Q2S => 2,
}
}
pub fn native_packing(&self) -> usize {
match self {
QuantValue::E2M1 => 2,
_ => 1,
}
}
pub fn range(&self) -> (f32, f32) {
match self {
QuantValue::Q8F => (i8::MIN as f32, i8::MAX as f32),
QuantValue::Q4F => (-8.0, 7.0),
QuantValue::Q2F => (-2.0, 1.0),
QuantValue::Q8S => (-i8::MAX as f32, i8::MAX as f32),
QuantValue::Q4S => (-7.0, 7.0),
QuantValue::Q2S => (-1.0, 1.0),
QuantValue::E4M3 => (-448.0, 448.0),
QuantValue::E5M2 => (-57344.0, 57344.0),
QuantValue::E2M1 => (-6.0, 6.0), }
}
pub fn is_symmetric(&self) -> bool {
match self {
Self::Q8F | Self::Q4F | Self::Q2F | Self::E4M3 | Self::E5M2 | Self::E2M1 => false,
Self::Q8S | Self::Q4S | Self::Q2S => true,
}
}
}
impl QuantStore {
pub fn size_bits(&self, value: &QuantValue) -> usize {
match self {
QuantStore::Native => value.size_bits(),
QuantStore::PackedNative(_) => value.size_bits() * value.native_packing(),
QuantStore::PackedU32(_) => 32,
}
}
fn packing_dim(&self) -> Option<usize> {
match self {
QuantStore::Native => None,
QuantStore::PackedNative(packing_dim) | QuantStore::PackedU32(packing_dim) => {
Some(*packing_dim)
}
}
}
}
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum QuantStore {
Native,
PackedNative(usize),
PackedU32(usize),
}
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum QuantMode {
Symmetric,
Lookup,
}
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)]
pub enum ScaleDtype {
F32,
F16,
BF16,
UE8M0,
UE4M3,
}
const MAX_DIMS: usize = 5;
#[derive(Clone, Copy, Hash, PartialEq, Eq, Serialize, Deserialize)]
pub struct BlockSize {
storage: [u8; MAX_DIMS],
len: u8,
}
impl PartialOrd for BlockSize {
fn partial_cmp(&self, other: &Self) -> Option<core::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for BlockSize {
fn cmp(&self, other: &Self) -> core::cmp::Ordering {
(self.len, self.as_slice()).cmp(&(other.len, other.as_slice()))
}
}
impl core::fmt::Debug for BlockSize {
fn fmt(&self, f: &mut core::fmt::Formatter) -> core::fmt::Result {
write!(f, "BlockSize({:?})", self.as_slice())
}
}
impl BlockSize {
pub const MAX_DIMS: usize = MAX_DIMS;
pub fn new(values: impl AsRef<[u8]>) -> Self {
Self::canonicalize(values.as_ref())
}
fn canonicalize(values: &[u8]) -> Self {
let skip = values
.iter()
.position(|&value| value != 1)
.unwrap_or(values.len());
let values = &values[skip..];
debug_assert!(
values.len() <= MAX_DIMS,
"Tried creating a block size larger than the cap"
);
let len = values.len().min(MAX_DIMS);
let mut storage = [1; MAX_DIMS];
storage[..len].copy_from_slice(&values[..len]);
Self {
storage,
len: len as u8,
}
}
pub fn as_slice(&self) -> &[u8] {
&self.storage[..self.len as usize]
}
pub fn to_vec(&self) -> Vec<u8> {
self.storage[..self.len as usize].to_vec()
}
pub fn as_dim<const N: usize>(&self) -> [u8; N] {
let data_len = N.min(self.len as usize);
let data_start = N - data_len;
let mut out = [1; N];
out[data_start..].copy_from_slice(&self.storage[..data_len]);
out
}
pub fn to_dim_vec(&self, len: usize) -> Vec<u8> {
let data_len = len.min(self.len as usize);
let data_start = len - data_len;
let mut out = vec![1; len];
out[data_start..].copy_from_slice(&self.storage[..data_len]);
out
}
pub fn num_blocks(&self, shape: &[usize]) -> Vec<usize> {
self.to_dim_vec(shape.len())
.into_iter()
.zip(shape)
.map(|(block, &dim)| dim.div_ceil(block as usize))
.collect()
}
pub fn iter(&self) -> impl Iterator<Item = &u8> {
self.as_slice().iter()
}
pub fn num_elements(&self) -> usize {
self.iter().map(|it| *it as usize).product()
}
}
impl Deref for BlockSize {
type Target = [u8];
fn deref(&self) -> &Self::Target {
self.as_slice()
}
}
impl<T: AsRef<[u8]>> From<T> for BlockSize {
fn from(value: T) -> Self {
BlockSize::new(value)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn blocks_remain_rank_relative() {
assert_ne!(BlockSize::new([32]), BlockSize::new([32, 32]));
assert_eq!(BlockSize::new([32]).to_dim_vec(2), vec![1, 32]);
assert_eq!(BlockSize::new([32, 32]).to_dim_vec(2), vec![32, 32]);
}
#[test]
fn leading_unit_dimensions_canonicalize_away() {
assert_eq!(BlockSize::new([1, 32]), BlockSize::new([32]));
}
#[test]
fn leading_unit_dimensions_beyond_the_cap_still_canonicalize() {
assert_eq!(
BlockSize::new([1, 1, 8, 4, 2, 3]),
BlockSize::new([8, 4, 2, 3])
);
}
#[test]
fn there_is_one_block_per_scale() {
assert_eq!(BlockSize::new([32]).num_blocks(&[8, 64]), vec![8, 2]);
assert_eq!(BlockSize::new([4, 32]).num_blocks(&[8, 64]), vec![2, 2]);
assert_eq!(BlockSize::new([32]).num_blocks(&[4, 8, 64]), vec![4, 8, 2]);
}
#[test]
fn a_partial_block_still_takes_a_scale() {
assert_eq!(BlockSize::new([32]).num_blocks(&[8, 70]), vec![8, 3]);
}
#[test]
fn the_default_scheme_resolves_to_per_tensor_f32() {
let scheme = QuantScheme::default();
assert_eq!(scheme.tensor_scale(), Some(ScaleDtype::F32));
assert_eq!(scheme.block_scale(), None);
assert_eq!(scheme.scale_dtype(), ScaleDtype::F32);
assert_eq!(scheme.block_size(), None);
assert_eq!(scheme.num_levels(), 1);
}
#[test]
fn a_block_level_stands_alone() {
let scheme = QuantScheme::default().per_block([32], ScaleDtype::F16);
assert_eq!(scheme.tensor_scale(), None);
assert_eq!(scheme.scale_dtype(), ScaleDtype::F16);
assert_eq!(scheme.block_size(), Some(BlockSize::new([32])));
assert_eq!(scheme.num_levels(), 1);
}
#[test]
fn both_levels_nest_the_block_inside_the_tensor() {
let scheme = QuantScheme::default()
.per_block([16], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::F32);
assert_eq!(scheme.scale_dtype(), ScaleDtype::UE4M3);
assert_eq!(scheme.tensor_scale(), Some(ScaleDtype::F32));
assert_eq!(scheme.num_levels(), 2);
}
#[test]
fn levels_set_in_any_order_are_the_same_scheme() {
assert_eq!(
QuantScheme::default()
.per_block([16], ScaleDtype::UE4M3)
.per_tensor(ScaleDtype::F32),
QuantScheme::default()
.per_tensor(ScaleDtype::F32)
.per_block([16], ScaleDtype::UE4M3),
);
}
#[test]
fn swapping_dims_rewrites_the_block_and_leaves_the_tensor_level_alone() {
let mut scheme = QuantScheme::default()
.per_block([4, 32], ScaleDtype::F16)
.per_tensor(ScaleDtype::F32);
scheme.swap_block_dims(2, 0, 1);
assert_eq!(
scheme,
QuantScheme::default()
.per_block([32, 4], ScaleDtype::F16)
.per_tensor(ScaleDtype::F32)
);
let mut per_tensor = QuantScheme::default();
per_tensor.swap_block_dims(2, 0, 1);
assert_eq!(per_tensor, QuantScheme::default());
}
#[test]
fn swapping_dims_canonicalizes_the_block() {
let mut scheme = QuantScheme::default().per_block([32, 1], ScaleDtype::F32);
scheme.swap_block_dims(2, 0, 1);
assert_eq!(scheme.block_size(), Some(BlockSize::new([32])));
}
#[test]
fn permuting_dims_rewrites_the_block() {
let mut scheme = QuantScheme::default().per_block([1, 4, 32], ScaleDtype::F16);
scheme.permute_block_dims(3, &[2, 0, 1]);
assert_eq!(scheme.block_size(), Some(BlockSize::new([32, 1, 4])));
}
#[test]
fn round_up_never_lands_below_the_scale() {
for dtype in [ScaleDtype::F16, ScaleDtype::BF16, ScaleDtype::UE4M3] {
for exp in -12..8 {
for step in 1..17 {
let scale = (step as f32 / 16.0) * 2f32.powi(exp);
let up = dtype.round_up(scale).unwrap();
assert!(
up >= scale,
"{dtype:?}: {up} is below {scale}, which clips the block maximum"
);
}
}
}
}
#[test]
fn round_up_saturates_rather_than_stepping_off_the_top() {
for dtype in [ScaleDtype::F16, ScaleDtype::BF16, ScaleDtype::UE4M3] {
let max = dtype.max_representable();
assert_eq!(dtype.round_up(max).unwrap(), max);
assert!(dtype.round_up(max * 2.0).unwrap().is_finite());
}
}
#[test]
fn round_up_answers_for_every_param() {
for dtype in [
ScaleDtype::F32,
ScaleDtype::F16,
ScaleDtype::BF16,
ScaleDtype::UE8M0,
ScaleDtype::UE4M3,
] {
assert_eq!(
dtype.round_up(0.3).is_some(),
dtype != ScaleDtype::UE8M0,
"{dtype:?}"
);
}
}
#[test]
fn round_up_is_the_identity_for_f32() {
for scale in [1.0e-30, 0.1, 1.0, 12345.678, f32::MAX] {
assert_eq!(ScaleDtype::F32.round_up(scale).unwrap(), scale);
}
}
#[cfg(feature = "fp8")]
mod storage_types {
use super::*;
#[test]
fn round_up_is_the_nearest_representable_value_not_below() {
for dtype in [ScaleDtype::F16, ScaleDtype::BF16, ScaleDtype::UE4M3] {
for exp in -8..6 {
let scale = 1.7 * 2f32.powi(exp);
let up = dtype.round_up(scale).unwrap();
assert_eq!(
up,
dtype.round_up(up).unwrap(),
"{dtype:?}: not idempotent at {scale}"
);
assert!(
step(dtype, up, -1) < scale,
"{dtype:?}: {up} overshoots {scale} by at least a step"
);
}
}
}
#[test]
fn f32_grid_matches_the_storage_types() {
for dtype in [ScaleDtype::F16, ScaleDtype::BF16, ScaleDtype::UE4M3] {
let grid = dtype.f32_grid();
if let Some(subnormals) = grid.subnormals {
assert_eq!(
subnormals.min_normal,
min_normal(dtype),
"{dtype:?}: minimum normal"
);
assert_eq!(
subnormals.spacing,
step(dtype, 0.0, 1),
"{dtype:?}: subnormal spacing"
);
}
let mut value = min_normal(dtype);
let max = dtype.max_representable();
while value < max {
let stepped = f32::from_bits(value.to_bits() + grid.bit_step);
assert_eq!(
stepped,
step(dtype, value, 1),
"{dtype:?}: step above {value}"
);
value = stepped;
}
assert_eq!(
value, max,
"{dtype:?}: the grid has to land exactly on the maximum"
);
}
}
#[test]
fn max_representable_matches_the_e4m3_type() {
assert_eq!(
ScaleDtype::UE4M3.max_representable(),
crate::e4m3::MAX.to_f32()
);
}
#[test]
fn max_representable_matches_the_e8m0_type() {
assert_eq!(
ScaleDtype::UE8M0.max_representable(),
crate::ue8m0::MAX.to_f32()
);
}
fn step(dtype: ScaleDtype, value: f32, offset: i32) -> f32 {
match dtype {
ScaleDtype::F16 => half::f16::from_bits(
(half::f16::from_f32(value).to_bits() as i32 + offset) as u16,
)
.to_f32(),
ScaleDtype::BF16 => half::bf16::from_bits(
(half::bf16::from_f32(value).to_bits() as i32 + offset) as u16,
)
.to_f32(),
ScaleDtype::UE4M3 => crate::e4m3::from_bits(
(crate::e4m3::from_f32(value).to_bits() as i32 + offset) as u8,
)
.to_f32(),
ScaleDtype::F32 | ScaleDtype::UE8M0 => unreachable!(),
}
}
fn min_normal(dtype: ScaleDtype) -> f32 {
match dtype {
ScaleDtype::F16 => half::f16::MIN_POSITIVE.to_f32(),
ScaleDtype::BF16 => half::bf16::MIN_POSITIVE.to_f32(),
ScaleDtype::UE4M3 => crate::e4m3::MIN_POSITIVE.to_f32(),
ScaleDtype::F32 | ScaleDtype::UE8M0 => unreachable!(),
}
}
}
}