#![allow(clippy::all)]
use super::iq_grids::*;
use super::k_quants::{BlockQ8_0, QK_K};
use super::quant_format::quant_format;
use half::f16;
pub const QK1_0: usize = 128;
pub const QK_NVFP4: usize = 64;
pub const QK_NVFP4_SUB: usize = 16;
const IQ1S_DELTA: f32 = 0.125;
const KVALUES_MXFP4: [i8; 16] = [0, 1, 2, 3, 4, 6, 8, 12, 0, -1, -2, -3, -4, -6, -8, -12];
#[inline]
fn kmask(j: usize) -> u8 {
1u8 << j
}
fn ue4m3_to_fp32(x: u8) -> f32 {
if x == 0 || x == 0x7F {
return 0.0;
}
let exp = ((x >> 3) & 0xF) as i32;
let man = (x & 0x7) as i32;
let raw = if exp == 0 {
(man as f32) * (-9f32).exp2()
} else {
(1.0 + man as f32 / 8.0) * ((exp - 7) as f32).exp2()
};
raw * 0.5
}
quant_format! {
name: BlockQ1_0,
dtype: Q1_0,
block_elems: QK1_0,
byte_size: 18,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u8; QK1_0 / 8],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK1_0), "dequantize_row_q1_0: {k} % {QK1_0} != 0");
let nb = k / QK1_0;
for i in 0..nb {
let d = xs[i].d.to_f32();
let neg_d = -d;
for j in 0..QK1_0 {
let byte = xs[i].qs[j / 8];
let bit = (byte >> (j % 8)) & 1;
ys[i * QK1_0 + j] = if bit != 0 { d } else { neg_d };
}
}
},
}
quant_format! {
name: BlockTQ2_0,
dtype: TQ2_0,
block_elems: QK_K,
byte_size: 66,
vec_dot: BlockQ8_0,
fields: {
qs: [u8; QK_K / 4],
d: f16,
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_tq2_0: {k} % {QK_K} != 0");
let nb = k / QK_K;
let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
let qslen = QK_K / 4;
let mut j = 0usize;
while j < qslen {
for l in 0..4 {
for m in 0..32 {
let q = ((xs[i].qs[j + m] >> (l * 2)) & 3) as i32;
ys[yi] = (q - 1) as f32 * d;
yi += 1;
}
}
j += 32;
}
}
},
}
quant_format! {
name: BlockTQ1_0,
dtype: TQ1_0,
block_elems: QK_K,
byte_size: 54,
vec_dot: BlockQ8_0,
fields: {
qs: [u8; (QK_K - 4 * QK_K / 64) / 5],
qh: [u8; QK_K / 64],
d: f16,
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_tq1_0: {k} % {QK_K} != 0");
let nb = k / QK_K;
let pow3: [u8; 6] = [1, 3, 9, 27, 81, 243];
let qslen = (QK_K - 4 * QK_K / 64) / 5; let qhlen = QK_K / 64; let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
let main_end = qslen - qslen % 32; let mut j = 0usize;
while j < main_end {
for n in 0..5 {
for m in 0..32 {
let q = xs[i].qs[j + m].wrapping_mul(pow3[n]);
let xi = ((q as u16) * 3) >> 8;
ys[yi] = (xi as i32 - 1) as f32 * d;
yi += 1;
}
}
j += 32;
}
let mut j = main_end;
while j < qslen {
for n in 0..5 {
for m in 0..16 {
let q = xs[i].qs[j + m].wrapping_mul(pow3[n]);
let xi = ((q as u16) * 3) >> 8;
ys[yi] = (xi as i32 - 1) as f32 * d;
yi += 1;
}
}
j += 16;
}
for n in 0..4 {
for j in 0..qhlen {
let q = xs[i].qh[j].wrapping_mul(pow3[n]);
let xi = ((q as u16) * 3) >> 8;
ys[yi] = (xi as i32 - 1) as f32 * d;
yi += 1;
}
}
}
},
}
quant_format! {
name: BlockNVFP4,
dtype: NVFP4,
block_elems: QK_NVFP4,
byte_size: 36,
vec_dot: BlockQ8_0,
fields: {
d: [u8; QK_NVFP4 / QK_NVFP4_SUB],
qs: [u8; QK_NVFP4 / 2],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_NVFP4), "dequantize_row_nvfp4: {k} % {QK_NVFP4} != 0");
let nb = k / QK_NVFP4;
let n_sub = QK_NVFP4 / QK_NVFP4_SUB; for i in 0..nb {
for s in 0..n_sub {
let d = ue4m3_to_fp32(xs[i].d[s]);
let yb = i * QK_NVFP4 + s * QK_NVFP4_SUB;
for j in 0..QK_NVFP4_SUB / 2 {
let q = xs[i].qs[s * (QK_NVFP4_SUB / 2) + j];
let v0 = KVALUES_MXFP4[(q & 0x0F) as usize];
let v1 = KVALUES_MXFP4[(q >> 4) as usize];
ys[yb + j] = v0 as f32 * d;
ys[yb + j + QK_NVFP4_SUB / 2] = v1 as f32 * d;
}
}
}
},
}
quant_format! {
name: BlockIQ2xxs,
dtype: IQ2_XXS,
block_elems: QK_K,
byte_size: 66,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u16; QK_K / 8],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_iq2_xxs: {k} % {QK_K} != 0");
let nb = k / QK_K;
let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
for ib32 in 0..QK_K / 32 {
let q = &xs[i].qs[4 * ib32..4 * ib32 + 4];
let aux32_0 = q[0] as u32 | ((q[1] as u32) << 16);
let aux32_1 = q[2] as u32 | ((q[3] as u32) << 16);
let aux8 = aux32_0.to_le_bytes();
let db = d * (0.5 + (aux32_1 >> 28) as f32) * 0.25;
for l in 0..4 {
let entry = IQ2XXS_GRID[aux8[l] as usize];
let signs = KSIGNS_IQ2XS[((aux32_1 >> (7 * l)) & 127) as usize];
for j in 0..8 {
let g = (entry >> (8 * j)) as u8 as f32;
let sign = if signs & kmask(j) != 0 { -1.0 } else { 1.0 };
ys[yi + j] = db * g * sign;
}
yi += 8;
}
}
}
},
}
quant_format! {
name: BlockIQ2xs,
dtype: IQ2_XS,
block_elems: QK_K,
byte_size: 74,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u16; QK_K / 8],
scales: [u8; QK_K / 32],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_iq2_xs: {k} % {QK_K} != 0");
let nb = k / QK_K;
let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
for ib32 in 0..QK_K / 32 {
let sc = xs[i].scales[ib32];
let db = [
d * (0.5 + (sc & 0xf) as f32) * 0.25,
d * (0.5 + (sc >> 4) as f32) * 0.25,
];
for l in 0..4 {
let qv = xs[i].qs[4 * ib32 + l];
let entry = IQ2XS_GRID[(qv & 511) as usize];
let signs = KSIGNS_IQ2XS[(qv >> 9) as usize];
let dl = db[l / 2];
for j in 0..8 {
let g = (entry >> (8 * j)) as u8 as f32;
let sign = if signs & kmask(j) != 0 { -1.0 } else { 1.0 };
ys[yi + j] = dl * g * sign;
}
yi += 8;
}
}
}
},
}
quant_format! {
name: BlockIQ2s,
dtype: IQ2_S,
block_elems: QK_K,
byte_size: 82,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u8; QK_K / 4],
qh: [u8; QK_K / 32],
scales: [u8; QK_K / 32],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_iq2_s: {k} % {QK_K} != 0");
let nb = k / QK_K;
let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
let mut qs_off = 0usize; let signs_base = QK_K / 8; let mut signs_off = 0usize;
for ib32 in 0..QK_K / 32 {
let sc = xs[i].scales[ib32];
let db = [
d * (0.5 + (sc & 0xf) as f32) * 0.25,
d * (0.5 + (sc >> 4) as f32) * 0.25,
];
let qh = xs[i].qh[ib32];
for l in 0..4 {
let dl = db[l / 2];
let idx = xs[i].qs[qs_off + l] as usize
| (((qh as usize) << (8 - 2 * l)) & 0x300);
let entry = IQ2S_GRID[idx];
let signs = xs[i].qs[signs_base + signs_off + l];
for j in 0..8 {
let g = (entry >> (8 * j)) as u8 as f32;
let sign = if signs & kmask(j) != 0 { -1.0 } else { 1.0 };
ys[yi + j] = dl * g * sign;
}
yi += 8;
}
qs_off += 4;
signs_off += 4;
}
}
},
}
quant_format! {
name: BlockIQ3xxs,
dtype: IQ3_XXS,
block_elems: QK_K,
byte_size: 98,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u8; 3 * QK_K / 8],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_iq3_xxs: {k} % {QK_K} != 0");
let nb = k / QK_K;
let scales_base = QK_K / 4; let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
let mut qs_off = 0usize; for ib32 in 0..QK_K / 32 {
let sb = scales_base + 4 * ib32;
let aux32 = u32::from_le_bytes([
xs[i].qs[sb],
xs[i].qs[sb + 1],
xs[i].qs[sb + 2],
xs[i].qs[sb + 3],
]);
let db = d * (0.5 + (aux32 >> 28) as f32) * 0.5;
for l in 0..4 {
let signs = KSIGNS_IQ2XS[((aux32 >> (7 * l)) & 127) as usize];
let g1 = IQ3XXS_GRID[xs[i].qs[qs_off + 2 * l] as usize];
let g2 = IQ3XXS_GRID[xs[i].qs[qs_off + 2 * l + 1] as usize];
for j in 0..4 {
let v1 = (g1 >> (8 * j)) as u8 as f32;
let v2 = (g2 >> (8 * j)) as u8 as f32;
let s1 = if signs & kmask(j) != 0 { -1.0 } else { 1.0 };
let s2 = if signs & kmask(j + 4) != 0 { -1.0 } else { 1.0 };
ys[yi + j] = db * v1 * s1;
ys[yi + j + 4] = db * v2 * s2;
}
yi += 8;
}
qs_off += 8;
}
}
},
}
quant_format! {
name: BlockIQ3s,
dtype: IQ3_S,
block_elems: QK_K,
byte_size: 110,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u8; QK_K / 4],
qh: [u8; QK_K / 32],
signs: [u8; QK_K / 8],
scales: [u8; QK_K / 64],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_iq3_s: {k} % {QK_K} != 0");
let nb = k / QK_K;
let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
let mut qs_off = 0usize; let mut signs_off = 0usize; let mut qh_off = 0usize; let mut ib32 = 0usize;
while ib32 < QK_K / 32 {
let sc = xs[i].scales[ib32 / 2];
let db1 = d * (1.0 + 2.0 * (sc & 0xf) as f32);
let db2 = d * (1.0 + 2.0 * (sc >> 4) as f32);
let qh0 = xs[i].qh[qh_off] as usize;
for l in 0..4 {
let idx1 = xs[i].qs[qs_off + 2 * l] as usize | ((qh0 << (8 - 2 * l)) & 256);
let idx2 = xs[i].qs[qs_off + 2 * l + 1] as usize | ((qh0 << (7 - 2 * l)) & 256);
let g1 = IQ3S_GRID[idx1];
let g2 = IQ3S_GRID[idx2];
let signs = xs[i].signs[signs_off + l];
for j in 0..4 {
let v1 = (g1 >> (8 * j)) as u8 as f32;
let v2 = (g2 >> (8 * j)) as u8 as f32;
let s1 = if signs & kmask(j) != 0 { -1.0 } else { 1.0 };
let s2 = if signs & kmask(j + 4) != 0 { -1.0 } else { 1.0 };
ys[yi + j] = db1 * v1 * s1;
ys[yi + j + 4] = db1 * v2 * s2;
}
yi += 8;
}
qs_off += 8;
signs_off += 4;
let qh1 = xs[i].qh[qh_off + 1] as usize;
for l in 0..4 {
let idx1 = xs[i].qs[qs_off + 2 * l] as usize | ((qh1 << (8 - 2 * l)) & 256);
let idx2 = xs[i].qs[qs_off + 2 * l + 1] as usize | ((qh1 << (7 - 2 * l)) & 256);
let g1 = IQ3S_GRID[idx1];
let g2 = IQ3S_GRID[idx2];
let signs = xs[i].signs[signs_off + l];
for j in 0..4 {
let v1 = (g1 >> (8 * j)) as u8 as f32;
let v2 = (g2 >> (8 * j)) as u8 as f32;
let s1 = if signs & kmask(j) != 0 { -1.0 } else { 1.0 };
let s2 = if signs & kmask(j + 4) != 0 { -1.0 } else { 1.0 };
ys[yi + j] = db2 * v1 * s1;
ys[yi + j + 4] = db2 * v2 * s2;
}
yi += 8;
}
qh_off += 2;
qs_off += 8;
signs_off += 4;
ib32 += 2;
}
}
},
}
quant_format! {
name: BlockIQ1s,
dtype: IQ1_S,
block_elems: QK_K,
byte_size: 50,
vec_dot: BlockQ8_0,
fields: {
d: f16,
qs: [u8; QK_K / 8],
qh: [u16; QK_K / 32],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_iq1_s: {k} % {QK_K} != 0");
let nb = k / QK_K;
let mut yi = 0usize;
for i in 0..nb {
let d = xs[i].d.to_f32();
let mut qs_off = 0usize; for ib in 0..QK_K / 32 {
let qh = xs[i].qh[ib];
let dl = d * (2.0 * ((qh >> 12) & 7) as f32 + 1.0);
let delta = if qh & 0x8000 != 0 { -IQ1S_DELTA } else { IQ1S_DELTA };
for l in 0..4 {
let idx = xs[i].qs[qs_off + l] as usize | ((((qh >> (3 * l)) & 7) as usize) << 8);
let entry = IQ1S_GRID[idx];
for j in 0..8 {
let g = (entry >> (8 * j)) as u8 as i8 as f32;
ys[yi + j] = dl * (g + delta);
}
yi += 8;
}
qs_off += 4;
}
}
},
}
quant_format! {
name: BlockIQ1m,
dtype: IQ1_M,
block_elems: QK_K,
byte_size: 56,
vec_dot: BlockQ8_0,
fields: {
qs: [u8; QK_K / 8],
qh: [u8; QK_K / 16],
scales: [u8; QK_K / 32],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_K), "dequantize_row_iq1_m: {k} % {QK_K} != 0");
let nb = k / QK_K;
let mut yi = 0usize;
for i in 0..nb {
let sc: [u16; 4] = [
u16::from_le_bytes([xs[i].scales[0], xs[i].scales[1]]),
u16::from_le_bytes([xs[i].scales[2], xs[i].scales[3]]),
u16::from_le_bytes([xs[i].scales[4], xs[i].scales[5]]),
u16::from_le_bytes([xs[i].scales[6], xs[i].scales[7]]),
];
let scale_u16 = (sc[0] >> 12)
| ((sc[1] >> 8) & 0x00f0)
| ((sc[2] >> 4) & 0x0f00)
| (sc[3] & 0xf000);
let d = f16::from_bits(scale_u16).to_f32();
let mut qs_off = 0usize; let mut qh_off = 0usize; for ib in 0..QK_K / 32 {
let dl1 = d * (2.0 * ((sc[ib / 2] >> (6 * (ib % 2))) & 0x7) as f32 + 1.0);
let dl2 = d * (2.0 * ((sc[ib / 2] >> (6 * (ib % 2) + 3)) & 0x7) as f32 + 1.0);
let qh0 = xs[i].qh[qh_off] as usize;
let qh1 = xs[i].qh[qh_off + 1] as usize;
let idx = [
xs[i].qs[qs_off] as usize | ((qh0 << 8) & 0x700),
xs[i].qs[qs_off + 1] as usize | ((qh0 << 4) & 0x700),
xs[i].qs[qs_off + 2] as usize | ((qh1 << 8) & 0x700),
xs[i].qs[qs_off + 3] as usize | ((qh1 << 4) & 0x700),
];
let delta = [
if qh0 & 0x08 != 0 { -IQ1S_DELTA } else { IQ1S_DELTA },
if qh0 & 0x80 != 0 { -IQ1S_DELTA } else { IQ1S_DELTA },
if qh1 & 0x08 != 0 { -IQ1S_DELTA } else { IQ1S_DELTA },
if qh1 & 0x80 != 0 { -IQ1S_DELTA } else { IQ1S_DELTA },
];
for l in 0..2 {
let entry = IQ1S_GRID[idx[l]];
for j in 0..8 {
let g = (entry >> (8 * j)) as u8 as i8 as f32;
ys[yi + j] = dl1 * (g + delta[l]);
}
yi += 8;
}
for l in 2..4 {
let entry = IQ1S_GRID[idx[l]];
for j in 0..8 {
let g = (entry >> (8 * j)) as u8 as i8 as f32;
ys[yi + j] = dl2 * (g + delta[l]);
}
yi += 8;
}
qs_off += 4;
qh_off += 2;
}
}
},
}
pub const QK_ROCMFP4: usize = 32;
const KVALUES_ROCMFP4: [i8; 16] = [0, 1, 2, 3, 4, 6, 8, 10, 0, -1, -2, -3, -4, -6, -8, -10];
fn rocmfp4_half_to_f32(e: u8) -> f32 {
if e > 0x7e {
return 0.0;
}
ue4m3_to_fp32(e)
}
quant_format! {
name: BlockROCMFP4,
dtype: ROCMFP4,
block_elems: QK_ROCMFP4,
byte_size: 18,
vec_dot: BlockQ8_0,
fields: {
qs: [u8; 16],
e: [u8; 2],
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_ROCMFP4), "dequantize_row_rocmfp4: {k} % {QK_ROCMFP4} != 0");
let nb = k / QK_ROCMFP4;
for i in 0..nb {
let d0 = rocmfp4_half_to_f32(xs[i].e[0]);
let d1 = rocmfp4_half_to_f32(xs[i].e[1]);
for j in 0..16 {
let q = xs[i].qs[j];
ys[i * QK_ROCMFP4 + j] = KVALUES_ROCMFP4[(q & 0x0F) as usize] as f32 * d0;
ys[i * QK_ROCMFP4 + j + 16] = KVALUES_ROCMFP4[(q >> 4) as usize] as f32 * d1;
}
}
},
}
quant_format! {
name: BlockROCMFP4Fast,
dtype: ROCMFP4_FAST,
block_elems: QK_ROCMFP4,
byte_size: 17,
vec_dot: BlockQ8_0,
fields: {
qs: [u8; 16],
e: u8,
},
decode: |xs, ys| {
let k = ys.len();
debug_assert!(k.is_multiple_of(QK_ROCMFP4), "dequantize_row_rocmfp4_fast: {k} % {QK_ROCMFP4} != 0");
let nb = k / QK_ROCMFP4;
for i in 0..nb {
let d = rocmfp4_half_to_f32(xs[i].e);
for j in 0..16 {
let q = xs[i].qs[j];
ys[i * QK_ROCMFP4 + j] = KVALUES_ROCMFP4[(q & 0x0F) as usize] as f32 * d;
ys[i * QK_ROCMFP4 + j + 16] = KVALUES_ROCMFP4[(q >> 4) as usize] as f32 * d;
}
}
},
}
#[cfg(test)]
mod rocmfp4_tests {
use super::*;
use crate::quantized::k_quants::{GgmlType, QK8_0};
#[test]
fn rocmfp4_dual_matches_manual_decode() {
let mut b = BlockROCMFP4 {
qs: [0; 16],
e: [0; 2],
};
b.e = [0x50, 0x38];
for j in 0..16 {
b.qs[j] = j as u8 | (((15 - j) as u8) << 4);
}
let mut ys = [0f32; 32];
BlockROCMFP4::to_float(std::slice::from_ref(&b), &mut ys);
let d0 = rocmfp4_half_to_f32(0x50);
let d1 = rocmfp4_half_to_f32(0x38);
assert!((d0 - 4.0).abs() < 1e-9, "d0 {d0}");
assert!((d1 - 0.5).abs() < 1e-9, "d1 {d1}");
for j in 0..16 {
assert_eq!(ys[j], KVALUES_ROCMFP4[j] as f32 * d0, "low half {j}");
assert_eq!(
ys[j + 16],
KVALUES_ROCMFP4[15 - j] as f32 * d1,
"high half {j}"
);
}
}
#[test]
fn rocmfp4_fast_matches_manual_decode() {
let mut b = BlockROCMFP4Fast {
qs: [0; 16],
e: 0x20,
}; for j in 0..16 {
b.qs[j] = j as u8 | ((j as u8) << 4);
}
let mut ys = [0f32; 32];
BlockROCMFP4Fast::to_float(std::slice::from_ref(&b), &mut ys);
let d = rocmfp4_half_to_f32(0x20);
assert!((d - 0.0625).abs() < 1e-9, "d {d}");
for j in 0..32 {
assert_eq!(ys[j], KVALUES_ROCMFP4[j % 16] as f32 * d, "elem {j}");
}
}
#[test]
fn rocmfp4_invalid_scales_decode_zero() {
let b = BlockROCMFP4Fast {
qs: [0x0F; 16],
e: 0x7F,
};
let mut ys = [0f32; 32];
BlockROCMFP4Fast::to_float(std::slice::from_ref(&b), &mut ys);
assert!(ys.iter().all(|v| *v == 0.0));
}
#[test]
fn rocmfp4_vec_dot_close_to_f32() {
let nb = 8;
let n = nb * QK_ROCMFP4;
let xs: Vec<BlockROCMFP4Fast> = (0..nb)
.map(|i| {
let mut b = BlockROCMFP4Fast { qs: [0; 16], e: 0 };
b.e = 0x30 + (i % 4) as u8;
for j in 0..16 {
b.qs[j] = ((i * 5 + j) & 0xFF) as u8;
}
b
})
.collect();
let wf32: Vec<f32> = {
let mut v = vec![0f32; n];
for (i, b) in xs.iter().enumerate() {
BlockROCMFP4Fast::to_float(
std::slice::from_ref(b),
&mut v[i * QK_ROCMFP4..][..QK_ROCMFP4],
);
}
v
};
let yf32: Vec<f32> = (0..n).map(|i| (i % 7) as f32 * 0.25 - 1.0).collect();
let exact: f32 = wf32.iter().zip(&yf32).map(|(a, b)| a * b).sum();
let ys_q8: Vec<BlockQ8_0> = {
let nb8 = n / QK8_0;
let mut v = vec![
BlockQ8_0 {
d: f16::from_bits(0),
qs: [0; QK8_0]
};
nb8
];
for (i, b) in v.iter_mut().enumerate() {
BlockQ8_0::from_float(&yf32[i * QK8_0..][..QK8_0], std::slice::from_mut(b));
}
v
};
let got = BlockROCMFP4Fast::vec_dot(n, &xs, &ys_q8);
let err = ((got - exact).abs() / exact.abs()).min(1.0);
assert!(
err < 2e-2,
"vec_dot rel err {err} (got {got}, exact {exact})"
);
}
}