ferrum_testkit/op_diff/
silu_mul.rs1use super::{random_vec, OpUnderTest, Output};
10
11pub struct SiluMulOp {
12 pub tokens: usize,
13 pub intermediate: usize,
15}
16
17impl SiluMulOp {
18 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 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 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}