#[cfg(test)]
mod dispatch_trace {
fn trace_one(label: &str, m: Option<usize>, k: Option<usize>, n: Option<usize>) {
let ops = crate::MmmDispatch::native();
let mmm = ops.preferred_kernel(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_wasm_kernels() {
let ops = crate::MmmDispatch::native();
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
.preferred_kernel(tract_data::prelude::DatumType::F32, Some(m), Some(k), Some(n))
.unwrap();
assert!(
mmm.arch().is_some(),
"{label}: dispatch returned the generic {} — strategize would \
discard it and reroute onto a GEMV kernel",
mmm.name(),
);
}
}
}
use crate::mmm::{AsInputValue, FusedSpec};
use tract_data::internal::*;
fn pick(name: &str) -> Box<dyn crate::mmm::MatMatMul> {
let ops = crate::MmmDispatch::native();
for impl_ in ops.runnable() {
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}"
);
}
}
#[cfg(test)]
fn check_madd_pairing<K: crate::mmm::MatMatMulKer<Acc = f32>>(ker: &K) {
use crate::mmm::{FusedKerSpec, OutputStoreKer};
if !ker.runnable() {
return;
}
let (mr, nr) = (ker.mr(), ker.nr());
let v = 1f32 + 2f32.powi(-12);
let (pack_a, pack_b) = &ker.packings()[0];
let k = pack_a.k_alignment().max(pack_b.k_alignment());
let mut a_data = vec![0f32; mr * k];
let mut b_data = vec![0f32; k * nr];
for i in 0..mr {
a_data[i * k] = v;
}
b_data[..nr].copy_from_slice(&vec![v; nr]);
let a = Tensor::from_shape(&[mr, k], &a_data).unwrap();
let b = Tensor::from_shape(&[k, nr], &b_data).unwrap();
let pa = pack_a.prepare_one(&a, 1, 0).unwrap();
let pb = pack_b.prepare_one(&b, 0, 1).unwrap();
let rows = vec![v; mr];
let cols = vec![v; nr];
let run = |op: FusedKerSpec<f32>| -> Vec<f32> {
let out = vec![0f32; mr * nr];
let item = std::mem::size_of::<f32>();
let store = OutputStoreKer {
ptr: out.as_ptr() as *mut u8,
row_byte_stride: (item * nr) as isize,
col_byte_stride: item as isize,
item_size: item,
};
let ops = [
FusedKerSpec::Clear,
FusedKerSpec::ScalarAdd(-1.0),
op,
FusedKerSpec::Store(store),
FusedKerSpec::Done,
];
assert_eq!(ker.kernel(&ops), 0);
out
};
let from_row_col = run(FusedKerSpec::AddRowColProducts(rows.as_ptr(), cols.as_ptr()));
let from_mat_mul = run(FusedKerSpec::AddMatMul {
k,
pa: pa.panel_bytes(0, None).unwrap(),
pb: pb.panel_bytes(0, None).unwrap(),
packing: 0,
});
for (i, (rc, mm)) in from_row_col.iter().zip(from_mat_mul.iter()).enumerate() {
assert_eq!(
rc.to_bits(),
mm.to_bits(),
"{}: cell {i} is {rc:e} from AddRowColProducts but {mm:e} from AddMatMul — \
the two arms disagree on whether the multiply-add is fused",
ker.name()
);
}
}
#[test]
fn add_row_col_products_and_add_mat_mul_agree_on_fusion() {
check_madd_pairing(&*crate::wasm::wasm_f32_4x4);
check_madd_pairing(&*crate::wasm::wasm_f32_4x1);
check_madd_pairing(&*crate::wasm::wasm_f32_8x1);
check_madd_pairing(&*crate::wasm::wasm_f32_16x1);
check_madd_pairing(&*crate::wasm::wasm_f32_32x1);
check_madd_pairing(&*crate::wasm::wasm_f32_8x8);
check_madd_pairing(&*crate::wasm::wasm_f32_4x16);
}
#[test]
fn dispatch_never_returns_wasm_f32_4x4() {
let ops = crate::MmmDispatch::native();
for m in [1usize, 3, 4, 5, 8, 9, 16, 17, 32, 64, 256, 1024] {
for n in [1usize, 2, 4, 8, 10, 64, 256] {
for k in [1usize, 64, 576] {
let mmm = ops
.preferred_kernel(
tract_data::prelude::DatumType::F32,
Some(m),
Some(k),
Some(n),
)
.unwrap();
assert_ne!(
mmm.name(),
"wasm_f32_4x4",
"m={m} k={k} n={n} dispatched to wasm_f32_4x4, which is registered \
NEVER_PREFERRED and has no dispatch band"
);
}
}
}
}
#[cfg(test)]
mod fastenhancer_gate {
fn pick(m: usize, k: usize, n: Option<usize>) -> String {
let ops = crate::MmmDispatch::native();
let mmm =
ops.preferred_kernel(tract_data::prelude::DatumType::F32, Some(m), Some(k), n).unwrap();
eprintln!("m={m} k={k} n={n:?} -> {} (mr={} nr={})", mmm.name(), mmm.mr(), mmm.nr());
mmm.name().to_string()
}
#[test]
fn gate_picks_4x16_only_on_multiples_of_16() {
for (m, k, n) in [(64, 192, 64), (64, 128, 64), (64, 64, 32), (64, 64, 48)] {
assert_eq!(
pick(m, k, Some(n)),
"wasm_f32_4x16",
"m={m} k={k} n={n}: n fills the 16-wide tile, so the wide kernel should win"
);
}
for (m, k, n) in [(64, 64, 8), (64, 64, 12), (64, 64, 24), (36, 36, 36)] {
assert_eq!(
pick(m, k, Some(n)),
"wasm_f32_8x8",
"m={m} k={k} n={n}: a 16-wide tile wastes columns here — n=8 runs at 0.58x, \
and n=24 pays that on its second tile"
);
}
}
#[test]
fn unknown_n_stays_on_8x8() {
assert_eq!(pick(64, 64, None), "wasm_f32_8x8");
}
}