use memra_engine::Engine;
fn e4m3_to_f32(x: u8) -> f32 {
let mag = (x & 0x7F) as u32;
if mag == 0x7F {
return 0.0;
}
let exp = ((mag >> 3) & 0xF) as i32;
let man = (mag & 0x7) as f32;
let raw = if exp == 0 {
(man * 0.125) * 0.015625 } else {
(1.0 + man * 0.125) * (2f32).powi(exp - 7)
};
if x & 0x80 != 0 { -raw } else { raw }
}
fn f32_to_e4m3(v: f32) -> u8 {
if v.is_nan() {
return 0x7F;
}
let sign = if v < 0.0 || (v == 0.0 && v.is_sign_negative()) {
0x80u8
} else {
0x00u8
};
let a = v.abs();
if a >= 448.0 {
return sign | 0x7E; }
let mut best = 0u8;
let mut best_err = f32::INFINITY;
for code in 0u8..=0x7E {
let c = e4m3_to_f32(code);
let err = (c - a).abs();
if err < best_err {
best_err = err;
best = code;
} else if err == best_err {
if code & 1 == 0 {
best = code;
}
}
}
sign | best
}
#[allow(clippy::too_many_arguments)]
fn host_ref(
w: &[u8], scales: &[f32], act_q: &[u8], act_d: &[f32], in_f: usize,
out_f: usize,
m: usize,
in_f_pad: usize,
) -> Vec<f32> {
let scols = in_f.div_ceil(128);
let mut y = vec![0f32; m * out_f];
let k_iter_end = in_f.div_ceil(128) * 128;
for i in 0..out_f {
let srow = i / 128;
for j in 0..m {
let mut sum = 0f32;
let mut kb = 0usize;
while kb < k_iter_end {
let s_blk = scales[srow * scols + (kb / 128).min(scols - 1)];
let db = if kb < in_f_pad {
act_d[j * (in_f_pad / 128) + kb / 128]
} else {
0.0
};
let mut c = 0f32;
for k01q in 0..4usize {
let g0 = kb + 32 * k01q;
for t in 0..32usize {
let g = g0 + t;
let wv = if g < in_f {
e4m3_to_f32(w[i * in_f + g])
} else {
0.0
};
let av = if g < in_f_pad {
e4m3_to_f32(act_q[j * in_f_pad + g])
} else {
0.0
};
c += wv * av;
}
}
sum += (s_blk * db) * c;
kb += 128;
}
y[j * out_f + i] = sum;
}
}
y
}
fn quantize_act_ref(x: &[f32], in_f: usize, m: usize) -> (usize, Vec<u8>, Vec<f32>) {
let in_f_pad = in_f.div_ceil(512) * 512;
let mut q = vec![0u8; m * in_f_pad];
let mut d = vec![0f32; m * (in_f_pad / 128)];
for j in 0..m {
for b in 0..(in_f_pad / 128) {
let mut amax = 0f32;
for t in 0..128 {
let g = b * 128 + t;
let v = if g < in_f { x[j * in_f + g] } else { 0.0 };
amax = amax.max(v.abs());
}
let (dv, dinv) = if amax == 0.0 {
(0.0, 0.0)
} else {
(amax / 448.0, 448.0 / amax)
};
d[j * (in_f_pad / 128) + b] = dv;
for t in 0..128 {
let g = b * 128 + t;
let v = if g < in_f { x[j * in_f + g] } else { 0.0 };
q[j * in_f_pad + g] = f32_to_e4m3(v * dinv);
}
}
}
(in_f_pad, q, d)
}
struct Rng(u64);
impl Rng {
fn next(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E3779B97F4A7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^ (z >> 31)
}
fn u8(&mut self) -> u8 {
(self.next() >> 33) as u8
}
fn f01(&mut self) -> f32 {
((self.next() >> 40) as f32) / (1u32 << 24) as f32
}
}
const INT_CODES: [u8; 13] = [
0x00, 0x38, 0xB8, 0x40, 0xC0, 0x44, 0xC4, 0x48, 0xC8, 0x4C, 0xCC, 0x50, 0xD0, ];
struct ArmResult {
max_abs: f32,
rms_rel: f32,
ulp_mismatch: usize,
n: usize,
}
fn compare(got: &[f32], want: &[f32]) -> ArmResult {
let mut r = ArmResult {
max_abs: 0.0,
rms_rel: 0.0,
ulp_mismatch: 0,
n: got.len(),
};
let mut sq = 0f64;
for (g, w) in got.iter().zip(want.iter()) {
if g.to_bits() != w.to_bits() {
r.ulp_mismatch += 1;
}
let a = (g - w).abs();
if a > r.max_abs {
r.max_abs = a;
}
sq += (*w as f64) * (*w as f64);
}
let rms = (sq / got.len().max(1) as f64).sqrt() as f32;
r.rms_rel = if rms > 0.0 {
r.max_abs / rms
} else {
r.max_abs
};
r
}
#[allow(clippy::type_complexity)]
fn run_shape(
e: &Engine,
in_f: usize,
out_f: usize,
m: usize,
exact_arm: bool,
seed: u64,
) -> Result<(ArmResult, u32, usize), Box<dyn std::error::Error>> {
let srows = (out_f + 127) / 128;
let scols = (in_f + 127) / 128;
let mut rng = Rng(seed);
let mut w = vec![0u8; out_f * in_f];
if exact_arm {
for b in w.iter_mut() {
*b = INT_CODES[(rng.u8() as usize) % INT_CODES.len()];
}
} else {
for b in w.iter_mut() {
let c = rng.u8();
*b = if c & 0x7F == 0x7F { c & 0xBF } else { c };
}
}
let mut scales = vec![0f32; srows * scols];
if exact_arm {
for s in scales.iter_mut() {
*s = (2f32).powi(((rng.u8() % 8) as i32) - 4);
}
} else {
for s in scales.iter_mut() {
*s = 0.002 + 0.5 * rng.f01();
}
}
let mut x = vec![0f32; m * in_f];
if exact_arm {
for j in 0..m {
let nb = in_f / 128;
for b in 0..nb {
for t in 0..128 {
let v = (rng.u8() % 5) as f32; let sgn = if rng.u8() & 1 == 0 { 1.0 } else { -1.0 };
x[j * in_f + b * 128 + t] = sgn * v;
}
x[j * in_f + b * 128 + (b * 37 + j) % 128] = 448.0;
}
}
} else {
for v in x.iter_mut() {
*v = 2.0 * rng.f01() - 1.0;
}
}
let (in_f_pad, aq, ad) = quantize_act_ref(&x, in_f, m);
let want = host_ref(&w, &scales, &aq, &ad, in_f, out_f, m, in_f_pad);
let mut seen = [false; 256];
for b in w.iter().chain(aq.iter()) {
seen[*b as usize] = true;
}
let codes = seen.iter().filter(|s| **s).count();
let wd = e.htod_bytes(&w)?;
let sd = e.htod(&scales)?;
let xd = e.htod(&x)?;
let nan = e.fp8_blk_nan_count(&wd)?;
let got = e.dtoh(&e.qmatvec_mmq_fp8_blk(&wd, &sd, &xd, m, in_f, out_f)?)?;
Ok((compare(&got, &want), nan, codes))
}
fn main() -> Result<(), Box<dyn std::error::Error>> {
let e = Engine::new(0)?;
let args: Vec<String> = std::env::args().skip(1).collect();
let shapes: Vec<(usize, usize, usize)> = if args.len() == 3 {
vec![(args[0].parse()?, args[1].parse()?, args[2].parse()?)]
} else {
vec![
(128, 128, 128), (256, 256, 64), (512, 384, 128), (5120, 1536, 512), (128, 200, 128), (384, 320, 96), (144, 128, 128), (272, 136, 40), (5120, 1536, 1), ]
};
let mut fails = 0usize;
let mut codes_max = 0usize;
println!(
"{:<22} {:>6} {:>12} {:>12} {:>15} {:>5} {:>6}",
"shape(in,out,m)", "arm", "max_abs", "rms_rel", "bit_mismatch", "nan", "codes"
);
for (in_f, out_f, m) in shapes {
let (r, nan, c1) = run_shape(
&e,
in_f,
out_f,
m,
true,
0xC0FFEE_u64 ^ (in_f * 7919) as u64,
)?;
let ok1 = r.ulp_mismatch == 0 && nan == 0;
println!(
"{:<22} {:>6} {:>12.3e} {:>12.3e} {:>9}/{:<5} {:>5} {:>6} {}",
format!("{in_f},{out_f},{m}"),
"EXACT",
r.max_abs,
r.rms_rel,
r.ulp_mismatch,
r.n,
nan,
c1,
if ok1 { "PASS" } else { "FAIL" }
);
if !ok1 {
fails += 1;
}
let (r2, nan2, c2) = run_shape(
&e,
in_f,
out_f,
m,
false,
0xBADC0DE_u64 ^ (out_f * 104729) as u64,
)?;
let ok2 = r2.rms_rel < 1e-5 && nan2 == 0;
println!(
"{:<22} {:>6} {:>12.3e} {:>12.3e} {:>9}/{:<5} {:>5} {:>6} {}",
format!("{in_f},{out_f},{m}"),
"RAND",
r2.max_abs,
r2.rms_rel,
r2.ulp_mismatch,
r2.n,
nan2,
c2,
if ok2 { "PASS" } else { "FAIL" }
);
if !ok2 {
fails += 1;
}
codes_max = codes_max.max(c1).max(c2);
}
println!();
let codes_ok = codes_max >= 254;
println!(
"e4m3 code coverage: {codes_max}/254 legal codes exercised (both NaN magnitudes 0x7F/0xFF excluded by the dispatch precondition) {}",
if codes_ok { "PASS" } else { "FAIL" }
);
if !codes_ok {
fails += 1;
}
println!();
if fails == 0 {
println!(
"=== fp8-mmq-check ALL GREEN (EXACT arm bit-identical, RAND arm < 1e-5 of RMS, 254/254 codes) ==="
);
Ok(())
} else {
println!("=== fp8-mmq-check {fails} FAILURES ===");
std::process::exit(1);
}
}