use super::*;
fn round_to_f16(data: &[f32]) -> Vec<f32> {
data.iter()
.map(|&x| half::f16::from_f32(x).to_f32())
.collect()
}
fn f16_tol(expected: f32) -> f32 {
1e-2 * (1.0 + expected.abs())
}
#[allow(clippy::too_many_arguments)]
fn run_gemm_f16_multi_tile_case(
trans_a: BackendTranspose,
trans_b: BackendTranspose,
m: usize,
n: usize,
k: usize,
) {
let Some(b) = try_init() else { return };
if !b.supports_f16() {
return; }
let ta = trans_a != BackendTranspose::NoTrans;
let tb = trans_b != BackendTranspose::NoTrans;
let a_data: Vec<f32> = (0..m * k).map(|x| (x as f32) * 0.02 - 1.0).collect();
let b_data: Vec<f32> = (0..k * n).map(|x| (x as f32) * 0.015 + 0.3).collect();
let c_init: Vec<f32> = (0..m * n).map(|x| (x as f32) * 0.01).collect();
let alpha = 1.25f32;
let beta = 0.75f32;
let a_r = round_to_f16(&a_data);
let b_r = round_to_f16(&b_data);
let c_r = round_to_f16(&c_init);
let expected = cpu_gemm(ta, tb, m, n, k, alpha, &a_r, &b_r, beta, &c_r);
let a_h = upload_f16(&b, &a_r);
let b_h = upload_f16(&b, &b_r);
let c_h = upload_f16(&b, &c_r);
b.gemm_f16(
trans_a,
trans_b,
m,
n,
k,
alpha as f64,
a_h,
if ta { m } else { k },
b_h,
if tb { k } else { n },
beta as f64,
c_h,
n,
)
.expect("gemm_f16 multi-tile transpose case");
let result = download_f16(&b, c_h, m * n);
for (idx, (r, e)) in result.iter().zip(expected.iter()).enumerate() {
assert!(
(r - e).abs() < f16_tol(*e),
"trans_a={trans_a}, trans_b={trans_b}, m={m}, n={n}, k={k}, slot={idx}: \
got {r}, expected {e}"
);
}
b.free(a_h).expect("free");
b.free(b_h).expect("free");
b.free(c_h).expect("free");
}
#[test]
fn gemm_f16_multi_tile_matches_reference_across_transpose_combinations() {
for &ta in &[BackendTranspose::NoTrans, BackendTranspose::Trans] {
for &tb in &[BackendTranspose::NoTrans, BackendTranspose::Trans] {
run_gemm_f16_multi_tile_case(ta, tb, 20, 18, 40);
}
}
}
#[test]
fn gemm_f16_transpose_with_padded_lda_multi_tile_matches_reference() {
let Some(b) = try_init() else { return };
if !b.supports_f16() {
return;
}
let m = 20usize;
let k = 40usize;
let n = 18usize;
let lda_a = m + 3; let ldb = n;
let ldc = n;
let a_logical = |row: usize, i: usize| -> f32 { ((row * k + i) as f32) * 0.02 - 1.0 };
let mut a_packed_t = vec![0.0f32; k * m];
for row in 0..m {
for i in 0..k {
a_packed_t[i * m + row] = a_logical(row, i);
}
}
let mut a_padded_t = vec![-9999.0f32; k * lda_a];
for row in 0..m {
for i in 0..k {
a_padded_t[i * lda_a + row] = a_logical(row, i);
}
}
let b_data: Vec<f32> = (0..k * n).map(|x| (x as f32) * 0.015 + 0.3).collect();
let c_init: Vec<f32> = (0..m * n).map(|x| (x as f32) * 0.01).collect();
let alpha = 1.5f32;
let beta = 0.5f32;
let a_packed_t_r = round_to_f16(&a_packed_t);
let a_padded_t_r = round_to_f16(&a_padded_t);
let b_r = round_to_f16(&b_data);
let c_r = round_to_f16(&c_init);
let expected = cpu_gemm(true, false, m, n, k, alpha, &a_packed_t_r, &b_r, beta, &c_r);
let a_h = upload_f16(&b, &a_padded_t_r);
let b_h = upload_f16(&b, &b_r);
let c_h = upload_f16(&b, &c_r);
b.gemm_f16(
BackendTranspose::Trans,
BackendTranspose::NoTrans,
m,
n,
k,
alpha as f64,
a_h,
lda_a,
b_h,
ldb,
beta as f64,
c_h,
ldc,
)
.expect("gemm_f16 TN, padded lda, multi-tile k");
let result = download_f16(&b, c_h, m * n);
for (idx, (r, e)) in result.iter().zip(expected.iter()).enumerate() {
assert!(
(r - e).abs() < f16_tol(*e),
"slot={idx}: got {r}, expected {e}"
);
}
b.free(a_h).expect("free");
b.free(b_h).expect("free");
b.free(c_h).expect("free");
}