#![cfg(all(
feature = "half",
feature = "std",
not(miri),
not(target_family = "wasm")
))]
use gemmkit::{MatMut, MatRef, Parallelism, bf16, f16, gemm, tuning};
static KNOB_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
fn knob_guard() -> std::sync::MutexGuard<'static, ()> {
KNOB_LOCK.lock().unwrap_or_else(|e| e.into_inner())
}
fn fill<N: gemmkit::NarrowFloat>(n: usize, seed: u64) -> Vec<N> {
let mut s = seed | 1;
(0..n)
.map(|_| {
s ^= s << 13;
s ^= s >> 7;
s ^= s << 17;
N::narrow((s >> 40) as f32 / (1u64 << 24) as f32 - 0.5)
})
.collect()
}
#[allow(clippy::too_many_arguments)]
fn run<N: gemmkit::NarrowFloat + gemmkit::GemmScalar + Copy>(
engage: bool,
m: usize,
k: usize,
n: usize,
a: &[N],
b: &[N],
c0: &[N],
rsc: isize,
csc: isize,
alpha: N,
beta: N,
par: Parallelism,
) -> Vec<u16> {
tuning::set_deep_kc_bytes(if engage { 1 } else { usize::MAX });
let mut c = c0.to_vec();
gemm(
alpha,
MatRef::from_col_major(a, m, k),
MatRef::from_col_major(b, k, n),
beta,
MatMut::new(&mut c, m, n, rsc, csc),
par,
);
c.iter().map(|v| bits16(*v)).collect()
}
fn bits16<N: gemmkit::NarrowFloat>(v: N) -> u16 {
unsafe { core::mem::transmute_copy::<N, u16>(&v) }
}
fn max_abs<N: gemmkit::NarrowFloat>(bits: &[u16], to: fn(u16) -> N) -> f32 {
bits.iter()
.map(|&b| to(b).widen().abs())
.fold(0.0, f32::max)
}
fn f16_from_bits(b: u16) -> f16 {
f16::from_bits(b)
}
fn bf16_from_bits(b: u16) -> bf16 {
bf16::from_bits(b)
}
fn bit_identity_case<N: gemmkit::NarrowFloat + gemmkit::GemmScalar + Copy>(label: &str) {
let _lock = knob_guard();
let restore = tuning::deep_kc_bytes();
let shapes: &[(usize, usize, isize, isize)] = &[
(48, 48, 1, 48),
(40, 50, 1, 40),
(130, 33, 1, 130),
(37, 41, 2, 2 * 37),
];
let k = 4096usize;
for &(m, n, rsc, csc) in shapes {
let a: Vec<N> = fill(m * k, 0x1234 ^ m as u64);
let b: Vec<N> = fill(k * n, 0x9abc ^ n as u64);
let cbacking = (m as isize * rsc.abs() + n as isize * csc.abs()) as usize + 8;
let c0: Vec<N> = fill(cbacking, 0x55 ^ (m * n) as u64);
for &beta_f in &[0.0f32, 1.0] {
let (alpha, beta) = (N::narrow(1.25), N::narrow(beta_f));
let single = run(
false,
m,
k,
n,
&a,
&b,
&c0,
rsc,
csc,
alpha,
beta,
Parallelism::Serial,
);
let deep_ser = run(
true,
m,
k,
n,
&a,
&b,
&c0,
rsc,
csc,
alpha,
beta,
Parallelism::Serial,
);
let deep_par = run(
true,
m,
k,
n,
&a,
&b,
&c0,
rsc,
csc,
alpha,
beta,
Parallelism::Rayon(0),
);
assert_eq!(
deep_ser, single,
"{label}: deep-k must be bit-identical to the single panel (m={m} n={n} k={k} beta={beta_f} rsc={rsc})"
);
assert_eq!(
deep_par, deep_ser,
"{label}: deep-k serial vs parallel must be bit-identical (m={m} n={n} beta={beta_f} rsc={rsc})"
);
}
}
tuning::set_deep_kc_bytes(restore);
}
fn tolerance_case<N: gemmkit::NarrowFloat + gemmkit::GemmScalar + Copy>(
label: &str,
to: fn(u16) -> N,
) {
let _lock = knob_guard();
let restore = tuning::deep_kc_bytes();
let (m, n, k) = (40usize, 50, 4096);
let a: Vec<N> = fill(m * k, 0x2222);
let b: Vec<N> = fill(k * n, 0x3333);
let (alpha, beta) = (N::narrow(0.75), N::narrow(2.5));
for &(rsc, csc) in &[(1isize, m as isize), (2isize, 2 * m as isize)] {
let cbacking = (m as isize * rsc.abs() + n as isize * csc.abs()) as usize + 8;
let c0: Vec<N> = fill(cbacking, 0x4444 ^ rsc as u64);
let single = run(
false,
m,
k,
n,
&a,
&b,
&c0,
rsc,
csc,
alpha,
beta,
Parallelism::Serial,
);
let deep = run(
true,
m,
k,
n,
&a,
&b,
&c0,
rsc,
csc,
alpha,
beta,
Parallelism::Serial,
);
let scale = max_abs(&single, to).max(1e-6);
let mut max_diff = 0.0f32;
for (&s, &d) in single.iter().zip(&deep) {
max_diff = max_diff.max((to(s).widen() - to(d).widen()).abs());
}
assert!(
max_diff <= 0.05 * scale,
"{label}: deep-k must match the single panel within tolerance for general beta (rsc={rsc}, max_diff={max_diff}, scale={scale})"
);
}
tuning::set_deep_kc_bytes(restore);
}
#[test]
fn deep_k_bit_identical_f16() {
bit_identity_case::<f16>("f16");
}
#[test]
fn deep_k_bit_identical_bf16() {
bit_identity_case::<bf16>("bf16");
}
#[test]
fn deep_k_tolerance_f16() {
tolerance_case::<f16>("f16", f16_from_bits);
}
#[test]
fn deep_k_tolerance_bf16() {
tolerance_case::<bf16>("bf16", bf16_from_bits);
}