Skip to main content

ferrum_testkit/op_diff/
gemm.rs

1//! Non-quantized `Backend::gemm` comparison against CPU F32.
2//! Metal uses F32 buffers (GEMV for m=1, tiled GEMM otherwise); CUDA
3//! uses F16 input/output buffers and cuBLAS F32 accumulation. This fixture
4//! does not exercise quantized Marlin or the production plan runtime.
5
6use super::{random_vec, OpUnderTest, Output};
7use ferrum_kernels::backend::Backend;
8
9/// `C[m, n] = A[m, k] ยท B[n, k]^T` (row-major, B already transposed
10/// to head-major). Matches the Backend::gemm signature used by Linear.
11pub struct GemmOp {
12    pub m: usize,
13    pub n: usize,
14    pub k: usize,
15}
16
17impl GemmOp {
18    /// Validate buffer sizes and the signed indices used by Metal/cuBLAS.
19    /// These are arithmetic limits, not a guarantee of available device memory.
20    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        // The Metal kernel adds tile offsets before checking edge coordinates,
44        // and increments K by 32. Include that final increment, also for GEMV.
45        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(); // Reject invalid direct fixture use before allocation.
62        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            // A missing dispatch/zero-filled output must fail this same oracle.
134            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}