use std::ffi::c_char;
use std::ffi::c_int;
unsafe extern "C" {
fn sgemm_(
transa: *const c_char,
transb: *const c_char,
m: *const c_int,
n: *const c_int,
k: *const c_int,
alpha: *const f32,
a: *const f32,
lda: *const c_int,
b: *const f32,
ldb: *const c_int,
beta: *const f32,
c: *mut f32,
ldc: *const c_int,
);
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn hgemm_(
transa: *const c_char,
transb: *const c_char,
m: *const c_int,
n: *const c_int,
k: *const c_int,
alpha: *const half::f16,
a: *const half::f16,
lda: *const c_int,
b: *const half::f16,
ldb: *const c_int,
beta: *const half::f16,
c: *mut half::f16,
ldc: *const c_int,
) {
let (m, n, k) = (*m, *n, *k);
let (lda, ldb, ldc) = (*lda, *ldb, *ldc);
if m <= 0 || n <= 0 || k <= 0 || lda <= 0 || ldb <= 0 || ldc <= 0 {
return;
}
let (mu, nu, ku) = (m as usize, n as usize, k as usize);
let ta = *transa as u8;
let tb = *transb as u8;
let is_trans = |flag: u8| flag == b'T' || flag == b't' || flag == b'C' || flag == b'c';
let a_rows = if is_trans(ta) { ku } else { mu };
let a_cols = if is_trans(ta) { mu } else { ku };
let b_rows = if is_trans(tb) { nu } else { ku };
let b_cols = if is_trans(tb) { ku } else { nu };
let a_f32: Vec<f32> = (0..a_cols)
.flat_map(|col| {
let base = col * lda as usize;
(0..a_rows).map(move |row| *a.add(base + row))
})
.map(|v| v.to_f32())
.collect();
let b_f32: Vec<f32> = (0..b_cols)
.flat_map(|col| {
let base = col * ldb as usize;
(0..b_rows).map(move |row| *b.add(base + row))
})
.map(|v| v.to_f32())
.collect();
let mut out = vec![0.0_f32; mu * nu];
let alpha_f32 = (*alpha).to_f32();
let beta_f32 = (*beta).to_f32();
sgemm_(
transa,
transb,
&m,
&n,
&k,
&alpha_f32,
a_f32.as_ptr(),
&lda,
b_f32.as_ptr(),
&ldb,
&beta_f32,
out.as_mut_ptr(),
&ldc,
);
for col in 0..nu {
let base = col * ldc as usize;
for row in 0..mu {
*c.add(base + row) = half::f16::from_f32(out[col * mu + row]);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hgemm_matches_f32_reference() {
let (m, n, k) = (2usize, 3usize, 4usize);
let a: Vec<half::f16> = (0..m * k)
.map(|i| half::f16::from_f32((i % 7) as f32 - 3.0))
.collect();
let b: Vec<half::f16> = (0..k * n)
.map(|i| half::f16::from_f32((i % 5) as f32 - 2.0))
.collect();
let mut c = vec![half::f16::ZERO; m * n]; let one = half::f16::ONE;
let zero = half::f16::ZERO;
unsafe {
hgemm_(
&(b'N' as c_char),
&(b'N' as c_char),
&(m as c_int),
&(n as c_int),
&(k as c_int),
&one,
a.as_ptr(),
&(m as c_int),
b.as_ptr(),
&(k as c_int),
&zero,
c.as_mut_ptr(),
&(m as c_int),
);
}
let a2d: Vec<Vec<f32>> = (0..m)
.map(|i| (0..k).map(|j| a[j * m + i].to_f32()).collect())
.collect();
let b2d: Vec<Vec<f32>> = (0..k)
.map(|i| (0..n).map(|j| b[j * k + i].to_f32()).collect())
.collect();
for row in 0..m {
for col in 0..n {
let expect: f32 = (0..k).map(|t| a2d[row][t] * b2d[t][col]).sum();
let got = c[col * m + row].to_f32();
assert!(
(expect - got).abs() < 1e-2,
"C[{row}][{col}] 期望 {expect} 实得 {got}"
);
}
}
}
}