1use super::{gemm::GemmOp, random_vec, silu_mul::SiluMulOp, OpUnderTest, Output};
6pub use ferrum_bench_core::release_regression::numerics::{SubmissionPhase, SUBMISSION_PHASES};
7use ferrum_kernels::backend::{cpu::CpuBackend, Backend};
8
9pub struct MetalContextOp {
10 pub tokens: usize,
11 pub intermediate: usize,
12 pub k: usize,
13}
14
15pub fn compare_submission_segments(
18 op: &MetalContextOp,
19 reference: &[f32],
20 actual: &[f32],
21 tolerance: f64,
22) -> Result<Vec<super::required::NumericalMetrics>, String> {
23 op.expected_output_len()?;
24 ferrum_bench_core::release_regression::numerics::compare_submission_segments(
25 op.segment_len()?,
26 reference,
27 actual,
28 tolerance,
29 )
30}
31
32struct Buffers<B: Backend> {
33 a: B::Buffer,
34 b: B::Buffer,
35 product: B::Buffer,
36 copied: B::Buffer,
37 output: B::Buffer,
38}
39
40impl MetalContextOp {
41 pub fn segment_len(&self) -> Result<usize, String> {
42 let n = self
43 .intermediate
44 .checked_mul(2)
45 .ok_or("context GEMM width overflow")?;
46 GemmOp {
47 m: self.tokens,
48 n,
49 k: self.k,
50 }
51 .expected_output_len()?;
52 SiluMulOp {
53 tokens: self.tokens,
54 intermediate: self.intermediate,
55 }
56 .expected_output_len()
57 }
58
59 pub fn expected_output_len(&self) -> Result<usize, String> {
60 let elements = self
61 .segment_len()?
62 .checked_mul(SUBMISSION_PHASES.len())
63 .ok_or("context output element count overflow")?;
64 let bytes = elements
65 .checked_mul(std::mem::size_of::<f32>())
66 .ok_or("context output byte size overflow")?;
67 if bytes > isize::MAX as usize {
68 return Err("context output exceeds host indexing range".into());
69 }
70 Ok(elements)
71 }
72
73 fn buffers<B: Backend>(&self, seed: u64) -> Buffers<B> {
74 self.expected_output_len().expect("invalid context fixture");
75 let n = 2 * self.intermediate;
76 Buffers {
77 a: B::from_slice(&random_vec(self.tokens * self.k, -0.5, 0.5, seed)),
78 b: B::from_slice(&random_vec(n * self.k, -0.5, 0.5, seed.wrapping_add(1))),
79 product: B::from_slice(&vec![11.0; self.tokens * n]),
81 copied: B::from_slice(&vec![-7.0; self.tokens * n]),
82 output: B::from_slice(&vec![13.0; self.tokens * self.intermediate]),
83 }
84 }
85
86 fn enqueue<B: Backend>(&self, ctx: &mut B::Context, buffers: &mut Buffers<B>) {
87 let n = 2 * self.intermediate;
88 B::gemm(
89 ctx,
90 &buffers.a,
91 &buffers.b,
92 &mut buffers.product,
93 self.tokens,
94 n,
95 self.k,
96 );
97 B::copy_slice(
98 ctx,
99 &buffers.product,
100 0,
101 &mut buffers.copied,
102 0,
103 self.tokens * n,
104 );
105 B::fused_silu_mul_split(
106 ctx,
107 &buffers.copied,
108 &mut buffers.output,
109 self.tokens,
110 self.intermediate,
111 );
112 }
113}
114
115impl OpUnderTest for MetalContextOp {
116 fn name(&self) -> &str {
117 "metal_context"
118 }
119
120 fn run_cpu(&self, seed: u64) -> Output {
121 let len = self.segment_len().expect("valid context shape");
122 let mut output = Vec::with_capacity(self.expected_output_len().unwrap());
123 for (index, _) in SUBMISSION_PHASES.iter().enumerate() {
124 let mut buffers = self.buffers::<CpuBackend>(seed.wrapping_add(index as u64 * 2));
125 let mut ctx = CpuBackend::new_context();
126 self.enqueue::<CpuBackend>(&mut ctx, &mut buffers);
127 CpuBackend::sync(&mut ctx);
128 output.extend(CpuBackend::to_vec(&buffers.output, len));
129 }
130 output
131 }
132
133 #[cfg(all(target_os = "macos", feature = "metal"))]
134 fn run_metal(&self, seed: u64) -> Output {
135 use ferrum_kernels::backend::metal::MetalBackend;
136 let len = self.segment_len().expect("valid context shape");
137 let mut output = Vec::with_capacity(self.expected_output_len().unwrap());
138 let mut first = self.buffers::<MetalBackend>(seed);
139 let mut ctx = MetalBackend::new_context();
140 self.enqueue::<MetalBackend>(&mut ctx, &mut first);
141 MetalBackend::sync_checked(&mut ctx).expect("initial chain Metal completion");
142 output.extend(MetalBackend::to_vec(&first.output, len));
143
144 let mut reused = self.buffers::<MetalBackend>(seed.wrapping_add(2));
145 let mut independent = self.buffers::<MetalBackend>(seed.wrapping_add(4));
146 let mut peer = MetalBackend::new_context();
147 self.enqueue::<MetalBackend>(&mut ctx, &mut reused);
148 self.enqueue::<MetalBackend>(&mut peer, &mut independent);
149 MetalBackend::sync_checked(&mut peer).expect("independent chain Metal completion");
152 MetalBackend::sync_checked(&mut ctx).expect("reused chain Metal completion");
153 output.extend(MetalBackend::to_vec(&reused.output, len));
154 output.extend(MetalBackend::to_vec(&independent.output, len));
155
156 let mut dropped = self.buffers::<MetalBackend>(seed.wrapping_add(6));
157 {
158 let mut pending = MetalBackend::new_context();
159 self.enqueue::<MetalBackend>(&mut pending, &mut dropped);
160 }
163 output.extend(MetalBackend::to_vec(&dropped.output, len));
164 output
165 }
166
167 #[cfg(feature = "cuda")]
168 fn run_cuda(&self, _seed: u64) -> Output {
169 panic!("Metal context lifecycle has no CUDA adapter")
170 }
171}
172
173#[cfg(test)]
174mod tests {
175 use super::*;
176 use crate::op_diff::{required::compare_outputs, NMSE_FP32_TOL};
177
178 #[test]
179 fn cpu_chain_matches_independent_matmul_and_silu_for_each_lifecycle_segment() {
180 let op = MetalContextOp {
181 tokens: 3,
182 intermediate: 33,
183 k: 35,
184 };
185 let actual = op.run_cpu(7);
186 let len = op.segment_len().unwrap();
187 assert_eq!(actual.len(), op.expected_output_len().unwrap());
188 for (index, segment) in actual.chunks_exact(len).enumerate() {
189 let seed = 7 + index as u64 * 2;
190 let a = random_vec(op.tokens * op.k, -0.5, 0.5, seed);
191 let b = random_vec(2 * op.intermediate * op.k, -0.5, 0.5, seed + 1);
192 let expected: Vec<f32> = a
193 .chunks_exact(op.k)
194 .flat_map(|row| {
195 let product: Vec<f32> = b
196 .chunks_exact(op.k)
197 .map(|column| row.iter().zip(column).map(|(a, b)| a * b).sum())
198 .collect();
199 let (gate, up) = product.split_at(op.intermediate);
200 gate.iter()
201 .zip(up)
202 .map(|(g, u)| g / (1.0 + (-g).exp()) * u)
203 .collect::<Vec<_>>()
204 })
205 .collect();
206 compare_outputs(&expected, segment, NMSE_FP32_TOL).unwrap();
207 assert!(compare_outputs(&expected, &vec![13.0; len], NMSE_FP32_TOL).is_err());
208 if index > 0 {
209 assert!(
210 compare_outputs(&actual[..len], segment, NMSE_FP32_TOL).is_err(),
211 "distinct contexts must not share the first result"
212 );
213 }
214 }
215 }
216
217 #[test]
218 fn invalid_context_shapes_fail_before_allocation() {
219 for (tokens, intermediate, k) in [
220 (0, 33, 35),
221 (3, 0, 35),
222 (3, 33, 0),
223 (usize::MAX, 33, 35),
224 (3, usize::MAX, 35),
225 (3, 33, usize::MAX),
226 ] {
227 assert!(MetalContextOp {
228 tokens,
229 intermediate,
230 k
231 }
232 .expected_output_len()
233 .is_err());
234 }
235 }
236}