1#[path = "backend_numerics/config.rs"]
3mod config;
4use config::{parse_args, Args, Config, Precision, USAGE};
5use ferrum_testkit::op_diff::metal_context::{
6 compare_submission_segments, MetalContextOp, SubmissionPhase, SUBMISSION_PHASES,
7};
8use ferrum_testkit::op_diff::required::{NumericalMetrics, RequiredReport, RequiredStatus};
9use serde::Serialize;
10use std::fs::{File, OpenOptions};
11use std::io::{Seek, Write};
12use std::process::ExitCode;
13use std::time::Instant;
14
15#[derive(Serialize)]
16struct Document {
17 schema_version: u32,
18 config: Config,
19 output_shape: [usize; 2],
20 expected_output_elements: usize,
21 configured_precision: Precision,
22 executed_precision: Option<Precision>,
25 execution_path: &'static str,
26 coverage: &'static str,
27 #[serde(skip_serializing_if = "Option::is_none")]
28 submission_phases: Option<&'static [SubmissionPhase]>,
29 #[serde(skip_serializing_if = "Option::is_none")]
30 submission_metrics: Option<Vec<NumericalMetrics>>,
31 started_at: String,
32 completed: bool,
33 finished_at: Option<String>,
34 execution_elapsed_ms: Option<f64>,
37 result: Option<RequiredReport>,
38}
39
40fn bind_output_shape(report: &mut RequiredReport, expected_elements: usize) {
41 let mut errors = Vec::new();
42 for (label, output) in [("reference", &report.reference), ("actual", &report.actual)] {
43 match output {
44 Some(output) if output.f32_bits.len() != expected_elements => errors.push(format!(
45 "{label} output has {} elements; configured shape requires {expected_elements}",
46 output.f32_bits.len()
47 )),
48 None if report.status == RequiredStatus::Passed => {
49 errors.push(format!("passed result is missing {label} output"))
50 }
51 _ => {}
52 }
53 }
54 if !errors.is_empty() {
55 if let Some(previous) = report.reason.take() {
56 errors.insert(0, previous);
57 }
58 report.reason = Some(errors.join("; "));
59 report.status = RequiredStatus::Failed;
60 }
61}
62
63fn bind_submission_segments(
64 report: &mut RequiredReport,
65 config: &Config,
66) -> Option<Vec<NumericalMetrics>> {
67 let config::Operation::MetalContext {
68 tokens,
69 intermediate,
70 k,
71 } = config.op
72 else {
73 return None;
74 };
75 let (Some(reference), Some(actual)) = (&report.reference, &report.actual) else {
76 return None;
77 };
78 match compare_submission_segments(
79 &MetalContextOp {
80 tokens,
81 intermediate,
82 k,
83 },
84 &reference.to_f32(),
85 &actual.to_f32(),
86 config.max_nmse,
87 ) {
88 Ok(metrics) => Some(metrics),
89 Err(error) => {
90 report.status = RequiredStatus::Failed;
91 report.reason = Some(match report.reason.take() {
92 Some(previous) => format!("{previous}; {error}"),
93 None => error,
94 });
95 None
96 }
97 }
98}
99
100fn write_document(file: &mut File, document: &Document) -> Result<(), String> {
101 let mut bytes = serde_json::to_vec_pretty(document)
104 .map_err(|error| format!("serialize report: {error}"))?;
105 bytes.push(b'\n');
106 file.rewind()
107 .map_err(|error| format!("seek report: {error}"))?;
108 file.set_len(0)
109 .map_err(|error| format!("truncate report: {error}"))?;
110 file.write_all(&bytes)
111 .map_err(|error| format!("write report: {error}"))?;
112 file.sync_all()
113 .map_err(|error| format!("flush report: {error}"))
114}
115
116fn run(args: Args) -> Result<(), String> {
117 let expected_elements = args.config.validate()?;
118 let mut file = OpenOptions::new()
119 .write(true)
120 .create_new(true)
121 .open(&args.report)
122 .map_err(|error| format!("create report {}: {error}", args.report.display()))?;
123 let is_context = matches!(args.config.op, config::Operation::MetalContext { .. });
124 let mut document = Document {
125 schema_version: 1,
126 output_shape: args.config.op.output_shape(),
127 expected_output_elements: expected_elements,
128 configured_precision: args.config.precision(),
129 executed_precision: None,
130 execution_path: args.config.op.execution_path(),
131 coverage: if is_context {
132 "Legacy MetalContext F32 compute/blit/compute, checked initial/reused/independent submissions and post-Drop data. Drop exposes no driver status. Excludes quantized kernels, production-plan dispatch and model performance."
133 } else {
134 "Only the selected Backend trait operation, shape and fixed adapter precision; excludes production-plan dispatch, quantized Marlin, paged attention, full models and performance claims"
135 },
136 submission_phases: is_context.then_some(&SUBMISSION_PHASES),
137 submission_metrics: None,
138 started_at: chrono::Utc::now().to_rfc3339(),
139 completed: false,
140 finished_at: None,
141 execution_elapsed_ms: None,
142 result: None,
143 config: args.config,
144 };
145 write_document(&mut file, &document)?;
146 let started = Instant::now();
147 let mut result = document.config.execute();
148 bind_output_shape(&mut result, expected_elements);
149 document.submission_metrics = bind_submission_segments(&mut result, &document.config);
150 document.execution_elapsed_ms = Some(started.elapsed().as_secs_f64() * 1000.0);
151 if result.actual.is_some() {
152 document.executed_precision = Some(document.config.precision());
153 }
154 document.completed = true;
155 document.finished_at = Some(chrono::Utc::now().to_rfc3339());
156 let passed = result.is_passed();
157 let failure = result
158 .reason
159 .clone()
160 .unwrap_or_else(|| format!("required backend result: {:?}", result.status));
161 document.result = Some(result);
162 write_document(&mut file, &document)?;
163 if passed {
164 println!(
165 "Numerical check completed; report: {}",
166 args.report.display()
167 );
168 Ok(())
169 } else {
170 Err(format!("{failure}; report: {}", args.report.display()))
171 }
172}
173
174fn entry(arguments: impl IntoIterator<Item = String>) -> ExitCode {
175 let result = match parse_args(arguments) {
176 Ok(Some(args)) => run(args),
177 Ok(None) => {
178 println!("{USAGE}");
179 return ExitCode::SUCCESS;
180 }
181 Err(error) => Err(error),
182 };
183 match result {
184 Ok(()) => ExitCode::SUCCESS,
185 Err(error) => {
186 eprintln!("backend numerics: {error}");
187 ExitCode::FAILURE
188 }
189 }
190}
191
192fn main() -> ExitCode {
193 entry(std::env::args().skip(1))
194}
195
196#[cfg(test)]
197#[path = "backend_numerics/tests.rs"]
198mod tests;