use crate::{DType, Device, Result, Tensor};
#[inline]
fn pow2(e: i32) -> f32 {
if e > 127 {
f32::INFINITY
} else if e >= -126 {
f32::from_bits(((e + 127) as u32) << 23)
} else if e >= -149 {
f32::from_bits(1u32 << (e + 149))
} else {
0.0
}
}
#[inline]
fn block_scale(amax: f32, divisor: f32) -> f32 {
pow2((amax / divisor).log2().ceil() as i32)
}
#[inline]
fn e4m3_grid_value(code: u32) -> f32 {
const EXP_SCALE: [f32; 16] = [
0.0, 0.015625, 0.03125, 0.0625, 0.125, 0.25, 0.5, 1.0, 2.0, 4.0, 8.0, 16.0, 32.0, 64.0,
128.0, 256.0,
];
let exp = ((code >> 3) & 0x0f) as usize;
let mant = (code & 0x07) as f32;
if exp == 0 {
mant * 0.001953125 } else {
(1.0 + mant * 0.125) * EXP_SCALE[exp]
}
}
#[inline]
pub fn e4m3_nearest(x: f32) -> f32 {
let sign = if x < 0.0 { -1.0 } else { 1.0 };
let ax = x.abs().min(448.0);
let (mut lo, mut hi) = (0i32, 126i32);
while lo < hi {
let mid = (lo + hi + 1) >> 1;
if e4m3_grid_value(mid as u32) <= ax {
lo = mid;
} else {
hi = mid - 1;
}
}
let mut best = lo;
if best < 126 {
let best_diff = (ax - e4m3_grid_value(best as u32)).abs();
let next_diff = (ax - e4m3_grid_value((best + 1) as u32)).abs();
if next_diff < best_diff
|| (next_diff == best_diff && ((best + 1) & 1) == 0 && (best & 1) != 0)
{
best += 1;
}
}
sign * e4m3_grid_value(best as u32)
}
const E2M1_VALUES: [f32; 8] = [0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0];
#[inline]
pub fn e2m1_nearest(x: f32) -> f32 {
let sign = if x < 0.0 { -1.0 } else { 1.0 };
let ax = x.abs().min(6.0);
let mut best = 0usize;
let mut best_diff = (ax - E2M1_VALUES[0]).abs();
for i in 1..8 {
let diff = (ax - E2M1_VALUES[i]).abs();
if diff < best_diff || (diff == best_diff && (i & 1) == 0 && (best & 1) != 0) {
best = i;
best_diff = diff;
}
}
sign * E2M1_VALUES[best]
}
const HADAMARD128_SCALE: f32 = 0.088_388_347_648_318_45;
pub fn hadamard128_inplace(x: &mut [f32]) -> Result<()> {
if x.len() != 128 {
crate::bail!("DSV4 Hadamard-128 requires a 128-wide row, got {}", x.len());
}
let mut stride = 1usize;
while stride < 128 {
let mut base = 0usize;
while base < 128 {
for i in 0..stride {
let a = x[base + i];
let b = x[base + stride + i];
x[base + i] = a + b;
x[base + stride + i] = a - b;
}
base += 2 * stride;
}
stride <<= 1;
}
for v in x.iter_mut() {
*v *= HADAMARD128_SCALE;
}
Ok(())
}
fn fp4_act_quantize_blocks(x: &mut [f32]) {
for block in x.chunks_mut(32) {
let mut amax = 0.0f32;
for &v in block.iter() {
let av = v.abs();
if av > amax {
amax = av;
}
}
if amax < 7.052_966_104_933_725e-38 {
amax = 7.052_966_104_933_725e-38;
}
let scale = block_scale(amax, 6.0);
for v in block.iter_mut() {
let mut t = *v / scale;
if t > 6.0 {
t = 6.0;
} else if t < -6.0 {
t = -6.0;
}
*v = e2m1_nearest(t) * scale;
}
}
}
fn fp8_kv_quantize_row(row: &mut [f32], n_nope: usize) {
for block in row[..n_nope].chunks_mut(64) {
let mut amax = 0.0f32;
for &v in block.iter() {
let av = v.abs();
if av > amax {
amax = av;
}
}
if amax < 1.0e-4 {
amax = 1.0e-4;
}
let scale = block_scale(amax, 448.0);
for v in block.iter_mut() {
let mut t = *v / scale;
if t > 448.0 {
t = 448.0;
} else if t < -448.0 {
t = -448.0;
}
*v = e4m3_nearest(t) * scale;
}
}
}
pub fn fp8_kv_quantize_rows_inplace(x: &mut [f32], head_dim: usize, n_rot: usize) -> Result<()> {
if head_dim == 0 {
crate::bail!("DSV4 FP8 KV round-trip: head_dim must be non-zero");
}
if n_rot > head_dim {
crate::bail!("DSV4 FP8 KV round-trip: n_rot {n_rot} exceeds head_dim {head_dim}");
}
let n_nope = head_dim - n_rot;
if !n_nope.is_multiple_of(64) {
crate::bail!("DSV4 FP8 KV round-trip: nope width {n_nope} must be a multiple of 64");
}
if !x.len().is_multiple_of(head_dim) {
crate::bail!(
"DSV4 FP8 KV round-trip: buffer len {} is not a multiple of head_dim {head_dim}",
x.len()
);
}
for row in x.chunks_mut(head_dim) {
fp8_kv_quantize_row(row, n_nope);
}
Ok(())
}
pub fn indexer_qat_rows_inplace(x: &mut [f32], head_dim: usize) -> Result<()> {
if head_dim != 128 {
crate::bail!("DSV4 indexer QAT expects 128-wide rows, got {head_dim}");
}
if !x.len().is_multiple_of(128) {
crate::bail!(
"DSV4 indexer QAT: buffer len {} is not a multiple of 128",
x.len()
);
}
for row in x.chunks_mut(128) {
hadamard128_inplace(row)?;
fp4_act_quantize_blocks(row);
}
Ok(())
}
fn apply_rows<F>(t: &Tensor, kernel: F) -> Result<Tensor>
where
F: FnOnce(&mut [f32], usize) -> Result<()>,
{
let dims: Vec<usize> = t.dims().to_vec();
let head_dim = match dims.last() {
Some(&d) => d,
None => crate::bail!("DSV4 QAT: scalar tensor has no row dimension"),
};
let device = t.device().clone();
let mut data = t
.to_dtype(DType::F32)?
.contiguous()?
.flatten_all()?
.to_vec1::<f32>()?;
kernel(&mut data, head_dim)?;
let out = Tensor::from_vec(data, dims, &Device::Cpu)?;
if device.is_cpu() {
Ok(out)
} else {
out.to_device(&device)
}
}
pub fn fp8_kv_quantize(t: &Tensor, n_rot: usize) -> Result<Tensor> {
apply_rows(t, |data, head_dim| {
fp8_kv_quantize_rows_inplace(data, head_dim, n_rot)
})
}
pub fn indexer_qat(t: &Tensor) -> Result<Tensor> {
apply_rows(t, indexer_qat_rows_inplace)
}
#[cfg(test)]
mod tests {
use super::*;
use float8::F8E4M3;
#[test]
fn pow2_matches_ldexp() {
for e in -149i32..=127 {
let expected = (2f64.powi(e)) as f32;
assert_eq!(pow2(e), expected, "pow2({e})");
}
assert_eq!(pow2(128), f32::INFINITY);
assert_eq!(pow2(-150), 0.0);
}
#[test]
fn e4m3_grid_matches_float8_oracle() {
for code in 0u32..=126 {
let ours = e4m3_grid_value(code);
let oracle = F8E4M3::from_bits(code as u8).to_f32();
assert_eq!(ours, oracle, "E4M3 grid code {code}");
}
assert_eq!(e4m3_grid_value(0), 0.0);
assert_eq!(e4m3_grid_value(1), 0.001953125); assert_eq!(e4m3_grid_value(126), 448.0); }
#[test]
fn e4m3_nearest_grid_points_are_fixed() {
for code in 0u32..=126 {
let v = e4m3_grid_value(code);
assert_eq!(e4m3_nearest(v), v, "+grid {code}");
assert_eq!(e4m3_nearest(-v), -v, "-grid {code}");
assert_eq!(F8E4M3::from_f32(v).to_f32(), v, "oracle grid {code}");
}
}
#[test]
fn e4m3_nearest_saturates_and_rounds_half_even() {
assert_eq!(e4m3_nearest(500.0), 448.0);
assert_eq!(e4m3_nearest(-1.0e9), -448.0);
assert_eq!(e4m3_nearest(0.0), 0.0);
assert_eq!(e4m3_nearest(272.0), 256.0);
assert_eq!(e4m3_nearest(304.0), 320.0);
}
#[test]
fn e2m1_nearest_known_and_half_even() {
for (i, &v) in E2M1_VALUES.iter().enumerate() {
assert_eq!(e2m1_nearest(v), v, "grid {i}");
assert_eq!(e2m1_nearest(-v), -v, "neg grid {i}");
}
assert_eq!(e2m1_nearest(10.0), 6.0); assert_eq!(e2m1_nearest(-10.0), -6.0);
assert_eq!(e2m1_nearest(0.25), 0.0); assert_eq!(e2m1_nearest(1.25), 1.0); assert_eq!(e2m1_nearest(1.75), 2.0); assert_eq!(e2m1_nearest(2.5), 2.0); assert_eq!(e2m1_nearest(5.0), 4.0); }
#[test]
fn hadamard128_is_its_own_inverse() {
let mut x = [0f32; 128];
let mut s: u32 = 0x1234_5678;
for v in x.iter_mut() {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
*v = (s >> 8) as f32 / (1u32 << 24) as f32 - 0.5;
}
let orig = x;
hadamard128_inplace(&mut x).unwrap();
hadamard128_inplace(&mut x).unwrap();
for i in 0..128 {
assert!((x[i] - orig[i]).abs() < 1e-5, "involution at {i}");
}
}
#[test]
fn hadamard128_dc_concentrates_in_first_bin() {
let mut x = [1f32; 128];
hadamard128_inplace(&mut x).unwrap();
assert!((x[0] - 128f32.sqrt()).abs() < 1e-4, "DC bin = {}", x[0]);
for &v in x.iter().skip(1) {
assert!(v.abs() < 1e-4);
}
}
#[test]
fn hadamard128_wrong_len_errors() {
let mut x = [0f32; 64];
assert!(hadamard128_inplace(&mut x).is_err());
}
#[test]
fn fp8_kv_roundtrip_exact_on_grid_values() {
let mut row = vec![2.0f32; 64];
fp8_kv_quantize_rows_inplace(&mut row, 64, 0).unwrap();
for &v in &row {
assert_eq!(v, 2.0);
}
}
#[test]
fn fp8_kv_roundtrip_is_idempotent() {
let mut row = vec![0f32; 64];
let mut s: u32 = 0xC0FF_EE11;
for v in row.iter_mut() {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
*v = ((s >> 7) as f32 / (1u32 << 25) as f32 - 0.5) * 3.0;
}
let mut once = row.clone();
fp8_kv_quantize_rows_inplace(&mut once, 64, 0).unwrap();
let mut twice = once.clone();
fp8_kv_quantize_rows_inplace(&mut twice, 64, 0).unwrap();
assert_eq!(once, twice, "FP8 KV round-trip must be idempotent");
for i in 0..64 {
let rel = (once[i] - row[i]).abs() / row[i].abs().max(1e-6);
assert!(rel < 0.15, "rel err {rel} at {i}");
}
}
#[test]
fn fp8_kv_leaves_rope_part_untouched() {
let mut row = vec![0.123_456f32; 128];
for (i, v) in row.iter_mut().enumerate() {
*v = i as f32 * 0.01 + 0.001; }
let rope_before: Vec<f32> = row[64..].to_vec();
let nope_before: Vec<f32> = row[..64].to_vec();
fp8_kv_quantize_rows_inplace(&mut row, 128, 64).unwrap();
assert_eq!(&row[64..], rope_before.as_slice(), "RoPE part changed");
assert_ne!(
&row[..64],
nope_before.as_slice(),
"nope part not quantized"
);
}
#[test]
fn fp4_act_roundtrip_is_idempotent() {
let mut row = vec![0f32; 128];
let mut s: u32 = 0xBEEF_1234;
for v in row.iter_mut() {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
*v = ((s >> 6) as f32 / (1u32 << 26) as f32 - 0.5) * 4.0;
}
let mut once = row.clone();
fp4_act_quantize_blocks(&mut once);
let mut twice = once.clone();
fp4_act_quantize_blocks(&mut twice);
assert_eq!(once, twice, "FP4 act round-trip must be idempotent");
}
#[test]
fn indexer_qat_equals_hadamard_then_fp4() {
let mut row = vec![0f32; 128];
let mut s: u32 = 0x0BAD_F00D;
for v in row.iter_mut() {
s = s.wrapping_mul(1_664_525).wrapping_add(1_013_904_223);
*v = (s >> 8) as f32 / (1u32 << 24) as f32 - 0.5;
}
let mut expected = row.clone();
hadamard128_inplace(&mut expected).unwrap();
fp4_act_quantize_blocks(&mut expected);
let mut got = row.clone();
indexer_qat_rows_inplace(&mut got, 128).unwrap();
assert_eq!(got, expected, "indexer QAT must be hadamard-then-fp4");
}
#[test]
fn indexer_qat_rejects_wrong_width() {
let mut x = vec![0f32; 256];
assert!(indexer_qat_rows_inplace(&mut x, 64).is_err());
assert!(indexer_qat_rows_inplace(&mut x, 128).is_ok());
}
#[test]
fn fp8_kv_rejects_unaligned_nope() {
let mut x = vec![0f32; 100];
assert!(fp8_kv_quantize_rows_inplace(&mut x, 100, 0).is_err());
}
#[test]
fn tensor_fp8_matches_slice() {
let rows = 3usize;
let head_dim = 128usize;
let n = rows * head_dim;
let raw: Vec<f32> = (0..n).map(|i| (i as f32 * 0.013).sin() * 5.0).collect();
let mut slice_out = raw.clone();
fp8_kv_quantize_rows_inplace(&mut slice_out, head_dim, 64).unwrap();
let t = Tensor::from_vec(raw, (rows, head_dim), &Device::Cpu).unwrap();
let out = fp8_kv_quantize(&t, 64).unwrap();
assert_eq!(out.dims(), &[rows, head_dim]);
let tensor_out = out.flatten_all().unwrap().to_vec1::<f32>().unwrap();
assert_eq!(tensor_out, slice_out);
}
#[test]
fn tensor_indexer_matches_slice() {
let rows = 2usize;
let head_dim = 128usize;
let n = rows * head_dim;
let raw: Vec<f32> = (0..n).map(|i| (i as f32 * 0.007).cos()).collect();
let mut slice_out = raw.clone();
indexer_qat_rows_inplace(&mut slice_out, head_dim).unwrap();
let t = Tensor::from_vec(raw, (rows, head_dim), &Device::Cpu).unwrap();
let out = indexer_qat(&t).unwrap();
assert_eq!(out.dims(), &[rows, head_dim]);
let tensor_out = out.flatten_all().unwrap().to_vec1::<f32>().unwrap();
assert_eq!(tensor_out, slice_out);
}
}