1use super::{random_vec, OpUnderTest, Output};
7use ferrum_kernels::backend::Backend;
8
9pub struct GemmOp {
12 pub m: usize,
13 pub n: usize,
14 pub k: usize,
15}
16
17impl GemmOp {
18 pub fn expected_output_len(&self) -> Result<usize, String> {
21 if self.m == 0 || self.n == 0 || self.k == 0 {
22 return Err("GEMM m, n and k must be nonzero".into());
23 }
24 let mut output = 0;
25 for (name, rows, cols) in [
26 ("A", self.m, self.k),
27 ("B", self.n, self.k),
28 ("output", self.m, self.n),
29 ] {
30 let elements = rows
31 .checked_mul(cols)
32 .ok_or_else(|| format!("GEMM {name} element count overflows usize"))?;
33 let bytes = elements
34 .checked_mul(std::mem::size_of::<f32>())
35 .ok_or_else(|| format!("GEMM {name} f32 byte size overflows usize"))?;
36 if bytes > isize::MAX as usize || elements > i32::MAX as usize {
37 return Err(format!(
38 "GEMM {name} exceeds host or signed kernel indexing range"
39 ));
40 }
41 output = elements;
42 }
43 for (name, dimension, tile) in [("m", self.m, 64), ("n", self.n, 32), ("k", self.k, 32)] {
46 let padded = dimension.div_ceil(tile).checked_mul(tile);
47 if padded.is_none_or(|value| value > i32::MAX as usize) {
48 return Err(format!(
49 "GEMM {name} tile exceeds signed kernel indexing range"
50 ));
51 }
52 }
53 Ok(output)
54 }
55
56 fn output_len(&self) -> usize {
57 self.expected_output_len().expect("invalid GEMM fixture")
58 }
59
60 fn build_input(&self, seed: u64) -> (Vec<f32>, Vec<f32>) {
61 self.output_len(); let a = random_vec(self.m * self.k, -1.0, 1.0, seed);
63 let b = random_vec(self.n * self.k, -1.0, 1.0, seed.wrapping_add(1));
64 (a, b)
65 }
66}
67
68impl OpUnderTest for GemmOp {
69 fn name(&self) -> &str {
70 "gemm"
71 }
72
73 fn run_cpu(&self, seed: u64) -> Output {
74 use ferrum_kernels::backend::cpu::CpuBackend;
75 let (a, b) = self.build_input(seed);
76 let mut ctx = CpuBackend::new_context();
77 let a_buf = CpuBackend::from_slice(&a);
78 let b_buf = CpuBackend::from_slice(&b);
79 let mut out = CpuBackend::alloc(self.output_len());
80 CpuBackend::gemm(&mut ctx, &a_buf, &b_buf, &mut out, self.m, self.n, self.k);
81 CpuBackend::sync(&mut ctx);
82 CpuBackend::to_vec(&out, self.output_len())
83 }
84
85 #[cfg(all(target_os = "macos", feature = "metal"))]
86 fn run_metal(&self, seed: u64) -> Output {
87 use ferrum_kernels::backend::metal::MetalBackend;
88 let (a, b) = self.build_input(seed);
89 let mut ctx = MetalBackend::new_context();
90 let a_buf = MetalBackend::from_slice(&a);
91 let b_buf = MetalBackend::from_slice(&b);
92 let mut out = MetalBackend::alloc(self.output_len());
93 MetalBackend::gemm(&mut ctx, &a_buf, &b_buf, &mut out, self.m, self.n, self.k);
94 MetalBackend::sync_checked(&mut ctx)
95 .unwrap_or_else(|error| panic!("GEMM Metal completion failed: {error}"));
96 MetalBackend::to_vec(&out, self.output_len())
97 }
98
99 #[cfg(feature = "cuda")]
100 fn run_cuda(&self, seed: u64) -> Output {
101 use ferrum_kernels::backend::cuda::CudaBackend;
102 let (a, b) = self.build_input(seed);
103 let mut ctx = CudaBackend::new_context();
104 let a_buf = CudaBackend::from_slice(&a);
105 let b_buf = CudaBackend::from_slice(&b);
106 let mut out = CudaBackend::alloc(self.output_len());
107 CudaBackend::gemm(&mut ctx, &a_buf, &b_buf, &mut out, self.m, self.n, self.k);
108 CudaBackend::sync(&mut ctx);
109 CudaBackend::to_vec(&out, self.output_len())
110 }
111}
112
113#[cfg(test)]
114mod tests {
115 use super::*;
116 use crate::op_diff::{required::compare_outputs, NMSE_FP32_TOL};
117
118 #[test]
119 fn validates_gemm_shapes_and_compares_non_square_layout() {
120 for (m, n, k) in [(64, 32, 32), (64, 33, 35), (1, 3, 5)] {
121 let op = GemmOp { m, n, k };
122 assert_eq!(op.expected_output_len(), Ok(m * n));
123 let (a, b) = op.build_input(17);
124 let actual = op.run_cpu(17);
125 let expected: Vec<f32> = a
126 .chunks_exact(k)
127 .flat_map(|row| {
128 b.chunks_exact(k)
129 .map(move |column| row.iter().zip(column).map(|(x, y)| x * y).sum())
130 })
131 .collect();
132 compare_outputs(&expected, &actual, NMSE_FP32_TOL).unwrap();
133 assert!(compare_outputs(&expected, &vec![0.0; m * n], NMSE_FP32_TOL).is_err());
135 }
136 }
137
138 #[test]
139 fn rejects_gemm_empty_product_byte_and_tile_overflow_before_allocation() {
140 for (m, n, k) in [
141 (0, 3, 4),
142 (2, 0, 4),
143 (2, 3, 0),
144 (usize::MAX, 2, 2),
145 (2, usize::MAX, 2),
146 (2, 2, usize::MAX),
147 (1, 1, usize::MAX / std::mem::size_of::<f32>() + 1),
148 (i32::MAX as usize, 1, 1),
149 (1, i32::MAX as usize, 1),
150 (1, 1, i32::MAX as usize),
151 (2, 2, i32::MAX as usize / 2 + 1),
152 ] {
153 assert!(
154 GemmOp { m, n, k }.expected_output_len().is_err(),
155 "{m}/{n}/{k}"
156 );
157 }
158 assert!(std::panic::catch_unwind(|| GemmOp {
159 m: usize::MAX,
160 n: 2,
161 k: 2
162 }
163 .run_cpu(0))
164 .is_err());
165 }
166}