Skip to main content

ferrum_testkit/op_diff/
silu_mul.rs

1//! `fused_silu_mul_split` op-diff harness.
2//!
3//! Input layout (matches the kernel API):
4//!   - `gate_up`: tokens × (2 * intermediate)
5//!   - For each token row: `[gate ‖ up]` concatenated
6//! Output:
7//!   - `out`: tokens × intermediate, where `out[i,j] = silu(gate[i,j]) * up[i,j]`
8
9use super::{random_vec, OpUnderTest, Output};
10
11pub struct SiluMulOp {
12    pub tokens: usize,
13    /// One side; the gate_up buffer is `tokens × (2*intermediate)`.
14    pub intermediate: usize,
15}
16
17impl SiluMulOp {
18    /// Checked row-split input/output sizes, before allocation or submission.
19    /// Both backends use signed 32-bit offsets into the full gate/up buffer.
20    pub fn expected_output_len(&self) -> Result<usize, String> {
21        if self.tokens == 0 || self.intermediate == 0 {
22            return Err("SiLU Mul tokens and intermediate must be nonzero".into());
23        }
24        let output = self
25            .tokens
26            .checked_mul(self.intermediate)
27            .ok_or("SiLU Mul output element count overflows usize")?;
28        let input = output
29            .checked_mul(2)
30            .ok_or("SiLU Mul gate/up element count overflows usize")?;
31        let bytes = input
32            .checked_mul(std::mem::size_of::<f32>())
33            .ok_or("SiLU Mul f32 input byte size overflows usize")?;
34        if bytes > isize::MAX as usize || input > i32::MAX as usize {
35            return Err("SiLU Mul input exceeds host or signed kernel indexing range".into());
36        }
37        // input <= i32::MAX also keeps the 256-wide padded output launch
38        // representable, including Metal's uint thread ID -> int conversion.
39        Ok(output)
40    }
41
42    fn output_len(&self) -> usize {
43        self.expected_output_len()
44            .expect("invalid SiLU Mul fixture")
45    }
46
47    fn build_input(&self, seed: u64) -> Vec<f32> {
48        random_vec(self.output_len() * 2, -3.0, 3.0, seed)
49    }
50}
51
52impl OpUnderTest for SiluMulOp {
53    fn name(&self) -> &str {
54        "fused_silu_mul"
55    }
56
57    fn run_cpu(&self, seed: u64) -> Output {
58        use ferrum_kernels::backend::cpu::CpuBackend;
59        use ferrum_kernels::backend::Backend;
60
61        let gate_up = self.build_input(seed);
62        let mut ctx = CpuBackend::new_context();
63        let gu_buf = CpuBackend::from_slice(&gate_up);
64        let mut out = CpuBackend::alloc(self.output_len());
65        CpuBackend::fused_silu_mul_split(
66            &mut ctx,
67            &gu_buf,
68            &mut out,
69            self.tokens,
70            self.intermediate,
71        );
72        CpuBackend::sync(&mut ctx);
73        CpuBackend::to_vec(&out, self.output_len())
74    }
75
76    #[cfg(all(target_os = "macos", feature = "metal"))]
77    fn run_metal(&self, seed: u64) -> Output {
78        use ferrum_kernels::backend::metal::MetalBackend;
79        use ferrum_kernels::backend::Backend;
80
81        let gate_up = self.build_input(seed);
82        let mut ctx = MetalBackend::new_context();
83        let gu_buf = MetalBackend::from_slice(&gate_up);
84        let mut out = MetalBackend::alloc(self.output_len());
85        MetalBackend::fused_silu_mul_split(
86            &mut ctx,
87            &gu_buf,
88            &mut out,
89            self.tokens,
90            self.intermediate,
91        );
92        MetalBackend::sync_checked(&mut ctx)
93            .unwrap_or_else(|error| panic!("SiLU Mul Metal completion failed: {error}"));
94        MetalBackend::to_vec(&out, self.output_len())
95    }
96
97    #[cfg(feature = "cuda")]
98    fn run_cuda(&self, seed: u64) -> Output {
99        use ferrum_kernels::backend::cuda::CudaBackend;
100        use ferrum_kernels::backend::Backend;
101
102        let gate_up = self.build_input(seed);
103        let mut ctx = CudaBackend::new_context();
104        let gu_buf = CudaBackend::from_slice(&gate_up);
105        let mut out = CudaBackend::alloc(self.output_len());
106        CudaBackend::fused_silu_mul_split(
107            &mut ctx,
108            &gu_buf,
109            &mut out,
110            self.tokens,
111            self.intermediate,
112        );
113        CudaBackend::sync(&mut ctx);
114        CudaBackend::to_vec(&out, self.output_len())
115    }
116}
117
118#[cfg(test)]
119mod tests {
120    use super::*;
121    use crate::op_diff::{required::compare_outputs, NMSE_FP32_TOL};
122
123    #[test]
124    fn validates_silu_mul_shapes_and_checks_each_row_split() {
125        for (tokens, intermediate) in [(4, 256), (3, 257)] {
126            let op = SiluMulOp {
127                tokens,
128                intermediate,
129            };
130            assert_eq!(op.expected_output_len(), Ok(tokens * intermediate));
131            let input = op.build_input(7);
132            let actual = op.run_cpu(7);
133            let expected: Vec<f32> = input
134                .chunks_exact(2 * intermediate)
135                .flat_map(|row| {
136                    let (gate, up) = row.split_at(intermediate);
137                    gate.iter().zip(up).map(|(g, u)| g / (1.0 + (-g).exp()) * u)
138                })
139                .collect();
140            compare_outputs(&expected, &actual, NMSE_FP32_TOL).unwrap();
141            // Accidentally returning only gate*up is finite and the right
142            // length, but omits SiLU and must fail the numerical comparison.
143            let without_activation: Vec<f32> = input
144                .chunks_exact(2 * intermediate)
145                .flat_map(|row| {
146                    row[..intermediate]
147                        .iter()
148                        .zip(&row[intermediate..])
149                        .map(|(g, u)| g * u)
150                })
151                .collect();
152            assert!(compare_outputs(&expected, &without_activation, NMSE_FP32_TOL).is_err());
153        }
154    }
155
156    #[test]
157    fn rejects_silu_mul_empty_product_byte_and_input_index_overflow() {
158        for (tokens, intermediate) in [
159            (0, 256),
160            (1, 0),
161            (usize::MAX, 2),
162            (1, usize::MAX / 2 + 1),
163            (1, usize::MAX / 8 + 1),
164            (1, i32::MAX as usize / 2 + 1),
165            (2, i32::MAX as usize / 4 + 1),
166        ] {
167            assert!(
168                SiluMulOp {
169                    tokens,
170                    intermediate
171                }
172                .expected_output_len()
173                .is_err(),
174                "{tokens}/{intermediate}"
175            );
176        }
177        assert!(std::panic::catch_unwind(|| SiluMulOp {
178            tokens: usize::MAX,
179            intermediate: 2
180        }
181        .run_cpu(0))
182        .is_err());
183    }
184}