use crate::common::*;
use gemmkit::{Activation, Bias, MatMut, MatRef, Parallelism, Workspace, gemm, gemm_fused};
fn fused_matrix<T: Flt>(par: Parallelism) {
let mut rng = Rng::new(0xE91109E1);
let shapes = [
(17usize, 17usize, 17usize), (33, 40, 24), (64, 64, 64),
(48, 96, 129), ];
let acts: [Option<Activation<T>>; 3] = [
None,
Some(Activation::Relu),
Some(Activation::LeakyRelu(T::of(0.1))),
];
let fast = fast_test();
let full_lattice = 0usize;
for (si, &(m, k, n)) in shapes.iter().enumerate() {
for &beta in &[T::ZERO, T::ONE, T::of(0.7)] {
for &alpha in &[T::ONE, T::of(0.9)] {
for layout in [Layout::Col, Layout::Row, Layout::ColPadded] {
for bias_kind in 0u8..=2 {
for act in &acts {
if fast
&& si != full_lattice
&& !(beta == T::of(0.7)
&& alpha == T::of(0.9)
&& matches!(layout, Layout::ColPadded)
&& bias_kind == 2
&& matches!(act, Some(Activation::LeakyRelu(_))))
{
continue;
}
check_fused::<T>(
&mut rng,
m,
k,
n,
alpha,
beta,
layout,
bias_kind,
act.clone_like(),
par,
"matrix",
);
}
}
}
}
}
}
}
#[test]
fn fused_eq_gemm_then_map_serial() {
fused_matrix::<f32>(Parallelism::Serial);
fused_matrix::<f64>(Parallelism::Serial);
}
#[test]
fn fused_eq_gemm_then_map_parallel() {
fused_matrix::<f32>(Parallelism::Rayon(8));
fused_matrix::<f64>(Parallelism::Rayon(8));
}
#[test]
fn identity_delegates_to_gemm() {
let mut rng = Rng::new(42);
for &(m, k, n) in &[(17usize, 20usize, 19usize), (64, 33, 48)] {
for layout in [Layout::Col, Layout::Row, Layout::ColPadded] {
for par in [Parallelism::Serial, Parallelism::Rayon(8)] {
let a = make::<f32>(&mut rng, m, k);
let b = make::<f32>(&mut rng, k, n);
let (rsc, csc, clen) = c_strides(layout, m, n);
let c0 = make::<f32>(&mut rng, clen, 1);
let mut c_fused = c0.clone();
let mut c_ref = c0.clone();
{
let ar = MatRef::new(&a, m, k, 1, m as isize);
let br = MatRef::new(&b, k, n, 1, k as isize);
let cm = MatMut::new(&mut c_fused, m, n, rsc, csc);
gemm_fused(1.0f32, ar, br, 0.5, cm, None, None, par);
}
{
let ar = MatRef::new(&a, m, k, 1, m as isize);
let br = MatRef::new(&b, k, n, 1, k as isize);
let cm = MatMut::new(&mut c_ref, m, n, rsc, csc);
gemm(1.0f32, ar, br, 0.5, cm, par);
}
for (x, y) in c_fused.iter().zip(c_ref.iter()) {
assert_eq!(x.to_bits(), y.to_bits(), "identity-fused != gemm");
}
}
}
}
}
#[test]
fn run_epilogue_identity_matches_run() {
use gemmkit::driver;
use gemmkit::kernel::{FloatGemm, Identity};
use gemmkit::simd::ScalarTok;
let mut rng = Rng::new(7);
for &(m, k, n) in &[(20usize, 24usize, 18usize), (40, 32, 40)] {
for par in [Parallelism::Serial, Parallelism::Rayon(4)] {
let a = make::<f32>(&mut rng, m, k);
let b = make::<f32>(&mut rng, k, n);
let c0 = make::<f32>(&mut rng, m * n, 1);
let mut c_run = c0.clone();
let mut c_epi = c0.clone();
let mut ws = Workspace::new();
unsafe {
driver::run::<FloatGemm<f32>, ScalarTok, 4, 4>(
ScalarTok,
m,
k,
n,
1.0,
a.as_ptr(),
1,
m as isize,
b.as_ptr(),
1,
k as isize,
0.7,
c_run.as_mut_ptr(),
1,
m as isize,
par,
&mut ws,
);
driver::run_epilogue::<FloatGemm<f32>, ScalarTok, Identity, 4, 4>(
ScalarTok,
m,
k,
n,
1.0,
a.as_ptr(),
1,
m as isize,
b.as_ptr(),
1,
k as isize,
0.7,
c_epi.as_mut_ptr(),
1,
m as isize,
&Identity,
par,
&mut ws,
);
}
for (x, y) in c_run.iter().zip(c_epi.iter()) {
assert_eq!(x.to_bits(), y.to_bits(), "run != run_epilogue::<Identity>");
}
}
}
}
#[test]
fn fire_once_multi_panel() {
let mut rng = Rng::new(0xF11E);
check_fused::<f32>(
&mut rng,
40,
4096,
40,
1.0,
0.7,
Layout::Col,
1, Some(Activation::Relu),
Parallelism::Serial,
"fire-once/serial",
);
check_fused::<f32>(
&mut rng,
40,
4096,
40,
0.9,
0.7,
Layout::Row,
2, Some(Activation::LeakyRelu(0.1)),
Parallelism::Rayon(8),
"fire-once/parallel",
);
}
#[test]
fn bias_orientation() {
let mut rng = Rng::new(0xB1A5);
for bias_kind in [1u8, 2u8] {
for layout in [Layout::Col, Layout::Row] {
check_fused::<f32>(
&mut rng,
33,
40,
21,
1.0,
0.0,
layout,
bias_kind,
None,
Parallelism::Serial,
"orient",
);
check_fused::<f64>(
&mut rng,
33,
40,
21,
1.0,
0.0,
layout,
bias_kind,
None,
Parallelism::Serial,
"orient",
);
}
}
}
#[test]
fn nan_and_neg_zero() {
nan_and_neg_zero_for::<f32>();
nan_and_neg_zero_for::<f64>();
}
fn nan_and_neg_zero_for<T: Flt>() {
let m = 64usize;
let k = 2usize;
let n = 64usize;
let ctx = T::name();
let mut a = vec![T::of(0.0); m * k];
let mut b = vec![T::of(0.0); k * n];
for i in 0..m {
a[i] = T::of(f64::INFINITY); a[m + i] = T::of(f64::INFINITY); }
for j in 0..n {
b[k * j] = T::of(1.0); b[k * j + 1] = T::of(-1.0); }
for &act in &[0u8, 1u8] {
let activation = if act == 1 {
Some(Activation::Relu)
} else {
Some(Activation::LeakyRelu(T::of(0.25)))
};
let mut c = vec![T::of(0.0); m * n];
{
let ar = MatRef::new(&a, m, k, 1, m as isize);
let br = MatRef::new(&b, k, n, 1, k as isize);
let cm = MatMut::new(&mut c, m, n, 1, m as isize);
gemm_fused(
T::of(1.0),
ar,
br,
T::of(0.0),
cm,
None,
activation.clone_like(),
Parallelism::Serial,
);
}
for &v in &c {
assert_eq!(
v.bits(),
T::of(0.0).bits(),
"{ctx}: ReLU/Leaky(NaN) must be +0.0"
);
}
}
let a2 = vec![T::of(0.0); m * k];
let b2 = vec![T::of(1.0); k * n];
let mut c2 = vec![T::of(0.0); m * n];
{
let ar = MatRef::new(&a2, m, k, 1, m as isize);
let br = MatRef::new(&b2, k, n, 1, k as isize);
let cm = MatMut::new(&mut c2, m, n, 1, m as isize);
gemm_fused(
T::of(1.0),
ar,
br,
T::of(0.0),
cm,
None,
Some(Activation::LeakyRelu(T::of(-0.5))),
Parallelism::Serial,
);
}
for &v in &c2 {
assert_eq!(
v.bits(),
T::of(0.0).bits(),
"{ctx}: LeakyReLU(0) must be +0.0"
);
}
}
#[test]
fn fused_degenerate() {
let mut rng = Rng::new(0xDE6E);
for &(m, n) in &[(20usize, 24usize)] {
let bias: Vec<f32> = (0..m).map(|_| rng.unit() as f32).collect();
let c0 = make::<f32>(&mut rng, m * n, 1);
for &(k, alpha) in &[(0usize, 1.0f32), (24usize, 0.0f32)] {
let a = make::<f32>(&mut rng, m, k.max(1));
let b = make::<f32>(&mut rng, k.max(1), n);
let mut c = c0.clone();
{
let ar = MatRef::new(&a, m, k, 1, m as isize);
let br = MatRef::new(&b, k, n, 1, k as isize);
let cm = MatMut::new(&mut c, m, n, 1, m as isize);
gemm_fused(
alpha,
ar,
br,
0.5,
cm,
Some(Bias::PerRow(&bias)),
Some(Activation::Relu),
Parallelism::Serial,
);
}
for j in 0..n {
for i in 0..m {
let idx = i + j * m;
let want = ref_apply(0.5f32 * c0[idx], Some(bias[i]), &Some(Activation::Relu));
assert_eq!(c[idx].to_bits(), want.to_bits(), "degenerate fused");
}
}
}
}
}
mod validation {
use super::*;
fn base() -> (Vec<f32>, Vec<f32>, Vec<f32>) {
(
vec![1.0f32; 4 * 4],
vec![1.0f32; 4 * 4],
vec![0.0f32; 4 * 4],
)
}
#[test]
#[should_panic(expected = "bias length")]
fn bias_wrong_length() {
let (a, b, mut c) = base();
let bias = vec![0.0f32; 3]; gemm_fused(
1.0,
MatRef::from_col_major(&a, 4, 4),
MatRef::from_col_major(&b, 4, 4),
0.0,
MatMut::from_col_major(&mut c, 4, 4),
Some(Bias::PerRow(&bias)),
None,
Parallelism::Serial,
);
}
#[test]
#[should_panic(expected = "LeakyRelu slope must be finite")]
fn leaky_slope_not_finite() {
let (a, b, mut c) = base();
gemm_fused(
1.0,
MatRef::from_col_major(&a, 4, 4),
MatRef::from_col_major(&b, 4, 4),
0.0,
MatMut::from_col_major(&mut c, 4, 4),
None,
Some(Activation::LeakyRelu(f32::INFINITY)),
Parallelism::Serial,
);
}
#[test]
#[should_panic(expected = "bias slice overlaps C")]
fn bias_overlaps_c() {
let a = vec![1.0f32; 16];
let b = vec![1.0f32; 16];
let mut buf = vec![0.0f32; 16];
let bias: &[f32] = unsafe { core::slice::from_raw_parts(buf.as_ptr(), 4) };
gemm_fused(
1.0,
MatRef::from_col_major(&a, 4, 4),
MatRef::from_col_major(&b, 4, 4),
0.0,
MatMut::from_col_major(&mut buf, 4, 4),
Some(Bias::PerRow(bias)),
None,
Parallelism::Serial,
);
}
}
#[test]
fn fused_unchecked_matches_checked() {
use gemmkit::{BiasDim, gemm_fused, gemm_fused_unchecked};
let mut rng = Rng::new(0x0F05_ED12);
let (m, k, n) = (33usize, 24usize, 40usize);
let a = make::<f32>(&mut rng, m, k); let b = make::<f32>(&mut rng, k, n); let c0 = make::<f32>(&mut rng, m * n, 1); let bias_row: Vec<f32> = (0..m).map(|_| (rng.unit() * 3.0) as f32).collect();
let bias_col: Vec<f32> = (0..n).map(|_| (rng.unit() * 3.0) as f32).collect();
let (alpha, beta) = (0.9f32, 0.7f32);
let par = Parallelism::Serial;
let mk_act = |kind: u8| match kind {
1 => Some(Activation::Relu),
2 => Some(Activation::LeakyRelu(0.1f32)),
_ => None,
};
for bias_kind in 0u8..=2 {
for act_kind in 0u8..=2 {
let bias_checked = match bias_kind {
1 => Some(Bias::PerRow(&bias_row)),
2 => Some(Bias::PerCol(&bias_col)),
_ => None,
};
let (bptr, bdim, has_bias) = match bias_kind {
1 => (bias_row.as_ptr(), BiasDim::PerRow, true),
2 => (bias_col.as_ptr(), BiasDim::PerCol, true),
_ => (core::ptr::null(), BiasDim::PerRow, false),
};
let mut c_checked = c0.clone();
{
let ar = MatRef::new(&a, m, k, 1, m as isize);
let br = MatRef::new(&b, k, n, 1, k as isize);
let cm = MatMut::new(&mut c_checked, m, n, 1, m as isize);
gemm_fused(alpha, ar, br, beta, cm, bias_checked, mk_act(act_kind), par);
}
let mut c_unchecked = c0.clone();
unsafe {
gemm_fused_unchecked(
m,
k,
n,
alpha,
a.as_ptr(),
1,
m as isize,
b.as_ptr(),
1,
k as isize,
beta,
c_unchecked.as_mut_ptr(),
1,
m as isize,
bptr,
bdim,
has_bias,
mk_act(act_kind),
par,
);
}
for idx in 0..m * n {
assert_eq!(
c_checked[idx].to_bits(),
c_unchecked[idx].to_bits(),
"fused unchecked != checked at {idx} [bias_kind={bias_kind} act_kind={act_kind}]",
);
}
}
}
}
#[test]
fn fused_unchecked_with_matches_checked() {
use gemmkit::{BiasDim, gemm_fused_unchecked_with, gemm_fused_with};
let mut rng = Rng::new(0x0F05_ED13);
let (m, k, n) = (40usize, 33usize, 24usize);
let a = make::<f64>(&mut rng, m, k); let b = make::<f64>(&mut rng, k, n); let c0 = make::<f64>(&mut rng, m * n, 1); let bias_row: Vec<f64> = (0..m).map(|_| rng.unit() * 3.0).collect();
let (alpha, beta) = (0.9f64, 0.7f64);
let par = Parallelism::Serial;
let mut c_checked = c0.clone();
{
let mut ws = Workspace::new();
let ar = MatRef::new(&a, m, k, 1, m as isize);
let br = MatRef::new(&b, k, n, 1, k as isize);
let cm = MatMut::new(&mut c_checked, m, n, 1, m as isize);
gemm_fused_with(
&mut ws,
alpha,
ar,
br,
beta,
cm,
Some(Bias::PerRow(&bias_row)),
Some(Activation::LeakyRelu(0.1)),
par,
);
}
let mut c_unchecked = c0.clone();
let mut ws = Workspace::new();
unsafe {
gemm_fused_unchecked_with(
&mut ws,
m,
k,
n,
alpha,
a.as_ptr(),
1,
m as isize,
b.as_ptr(),
1,
k as isize,
beta,
c_unchecked.as_mut_ptr(),
1,
m as isize,
bias_row.as_ptr(),
BiasDim::PerRow,
true,
Some(Activation::LeakyRelu(0.1)),
par,
);
}
for idx in 0..m * n {
assert_eq!(
c_checked[idx].to_bits(),
c_unchecked[idx].to_bits(),
"fused unchecked_with != checked at {idx}",
);
}
}