Skip to main content

ferrum_testkit/op_diff/
metal_context.rs

1//! Legacy Metal command-context lifecycle, using F32 GEMM -> blit -> SiLU.
2//! This proves dependent submission/readback behavior, not quantized kernels,
3//! production-plan dispatch or model scheduling performance.
4
5use 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
15/// Replay every lifecycle segment independently; a high-energy successful
16/// segment must not hide a failed lower-energy one in the aggregate NMSE.
17pub 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            // Distinct finite sentinels make a missing compute/blit visible.
80            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        // Submit the independently recorded context first: the reused context
150        // must retain its own command buffer, encoder and bound resources.
151        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            // Production Drop flushes and waits. Its API returns no driver
161            // status; this segment asserts the actual post-Drop data only.
162        }
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}