#[cfg(test)]
mod dispatch_trace {
fn trace_one(label: &str, m: Option<usize>, k: Option<usize>, n: Option<usize>) {
let mut ops = crate::generic();
crate::wasm::plug(&mut ops);
let mmm = ops.mmm(tract_data::prelude::DatumType::F32, m, k, n).unwrap();
eprintln!(
"DFN3 {} (m={:?} k={:?} n={:?}) => {} [mr={}, nr={}]",
label,
m,
k,
n,
mmm.name(),
mmm.mr(),
mmm.nr()
);
}
#[test]
fn dfn3_shapes() {
trace_one("lsnr_fc-style m=1 k=512", Some(1), Some(512), Some(1));
trace_one("small m=16 k=96", Some(16), Some(96), Some(1));
trace_one("medium m=32 k=256", Some(32), Some(256), Some(1));
trace_one("GRU m=256 k=256", Some(256), Some(256), Some(1));
trace_one("post-rnn m=256 k=512", Some(256), Some(512), Some(1));
trace_one("frame-encoder m=64 k=96", Some(64), Some(96), Some(1));
trace_one("MM m=64 k=64 n=8", Some(64), Some(64), Some(8));
}
#[test]
fn band_edges() {
trace_one("band 4x1 lo m=1", Some(1), Some(64), Some(1));
trace_one("band 4x1 hi m=4", Some(4), Some(64), Some(1));
trace_one("band 8x1 lo m=5", Some(5), Some(64), Some(1));
trace_one("band 8x1 hi m=8", Some(8), Some(64), Some(1));
trace_one("band 16x1 lo m=9", Some(9), Some(64), Some(1));
trace_one("band 16x1 hi m=16", Some(16), Some(64), Some(1));
trace_one("band 32x1 lo m=17", Some(17), Some(64), Some(1));
trace_one("band 32x1 hi m=512", Some(512), Some(64), Some(1));
}
#[test]
fn dispatch_kernels_are_manually_optimized() {
use crate::mmm::ImplementationQuality::ManuallyOptimized;
let mut ops = crate::generic();
crate::wasm::plug(&mut ops);
for (label, m, k, n) in [
("GEMM m=64 k=64 n=8", 64, 64, 8),
("GEMM m=256 k=256 n=256", 256, 256, 256),
("GEMM m=1024 k=576 n=10", 1024, 576, 10),
("GEMV m=1 k=512 n=1", 1, 512, 1),
("GEMV m=256 k=256 n=1", 256, 256, 1),
] {
let mmm =
ops.mmm(tract_data::prelude::DatumType::F32, Some(m), Some(k), Some(n)).unwrap();
assert_eq!(
mmm.quality(),
ManuallyOptimized,
"{label}: dispatch returned {} tagged {:?} — strategize would \
discard it and reroute onto a GEMV kernel",
mmm.name(),
mmm.quality(),
);
}
}
}
use crate::mmm::{AsInputValue, FusedSpec};
use tract_data::internal::*;
fn pick(name: &str) -> Box<dyn crate::mmm::MatMatMul> {
let mut ops = crate::generic();
crate::wasm::plug(&mut ops);
for impl_ in ops.mmm_impls() {
if impl_.name() == name {
return impl_.clone();
}
}
panic!("kernel {name} not registered")
}
#[test]
fn numerical_consistency_16x1_vs_32x1() {
let m = 256usize;
let k = 256usize;
let mut a_data = vec![0f32; m * k];
for (i, x) in a_data.iter_mut().enumerate() {
*x = ((i % 13) as f32 - 6.0) * 0.1 + ((i / 17) % 11) as f32 * 0.07;
}
let mut b_data = vec![0f32; k];
for (i, x) in b_data.iter_mut().enumerate() {
*x = (i as f32).sin() * 0.5;
}
let a = Tensor::from_shape(&[m, k], &a_data).unwrap();
let b = Tensor::from_shape(&[k, 1], &b_data).unwrap();
let run = |name: &str| -> Vec<f32> {
let kernel = pick(name);
let packing = &kernel.packings()[0];
let pa = packing.0.prepare_one(&a, 1, 0).unwrap();
let pb = packing.1.prepare_one(&b, 0, 1).unwrap();
let mut c = Tensor::zero::<f32>(&[m, 1]).unwrap();
unsafe {
kernel
.run(
m,
1,
&[
FusedSpec::AddMatMul {
a: AsInputValue::Borrowed(&*pa),
b: AsInputValue::Borrowed(&*pb),
packing: 0,
},
FusedSpec::Store(kernel.c_view(Some(0), Some(0)).wrap(&c.view_mut())),
],
)
.unwrap();
}
c.try_as_plain().unwrap().as_slice::<f32>().unwrap().to_vec()
};
let c16 = run("wasm_f32_16x1");
let c32 = run("wasm_f32_32x1");
#[cfg(not(target_feature = "relaxed-simd"))]
{
for (i, (x16, x32)) in c16.iter().zip(c32.iter()).enumerate() {
assert!(
x16.to_bits() == x32.to_bits(),
"row {i}: 16x1={x16} (bits 0x{:x}) != 32x1={x32} (bits 0x{:x})",
x16.to_bits(),
x32.to_bits()
);
}
eprintln!("bit-identity OK over m={m} k={k} ({} rows)", m);
}
#[cfg(target_feature = "relaxed-simd")]
{
let mut max_abs = 0.0f32;
let mut max_rel = 0.0f32;
for (i, (x16, x32)) in c16.iter().zip(c32.iter()).enumerate() {
let abs = (x16 - x32).abs();
let scale = x16.abs().max(x32.abs()).max(1.0e-9);
let rel = abs / scale;
assert!(
rel < 1.0e-4,
"row {i}: relative drift {rel:e} too large; 16x1={x16} 32x1={x32}"
);
if abs > max_abs {
max_abs = abs;
}
if rel > max_rel {
max_rel = rel;
}
}
eprintln!(
"relaxed-simd consistency OK over m={m} k={k}: max abs={max_abs:.3e}, max rel={max_rel:.3e}"
);
}
}