qualia_core_db/wgsl_forge/
tune.rs1use serde::{Deserialize, Serialize};
2
3use super::{
4 AdapterConstraints, ComparisonReport, ForgeError, KernelSpec, Schedule, ScheduleSpace,
5 TimingSource, TimingSummary,
6};
7
8#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
9pub struct CandidateEvaluation {
10 pub oracle: ComparisonReport,
11 pub timing_source: TimingSource,
12 pub samples_ns: Vec<u64>,
13}
14
15#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
16pub struct CandidateResult {
17 pub schedule: Schedule,
18 pub oracle: ComparisonReport,
19 pub timing: TimingSummary,
20}
21
22#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
23pub struct CandidateFailure {
24 pub schedule: Schedule,
25 pub reason: String,
26}
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
29pub struct TuningConfig {
30 pub initial_samples: usize,
31 pub finalist_samples: usize,
32 pub finalist_count: usize,
33 pub max_candidates: usize,
34}
35
36impl Default for TuningConfig {
37 fn default() -> Self {
38 Self {
39 initial_samples: 3,
40 finalist_samples: 11,
41 finalist_count: 6,
42 max_candidates: 48,
43 }
44 }
45}
46
47#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
48pub struct TuningResult {
49 pub evaluated_candidates: usize,
50 pub rejected_candidates: usize,
51 pub failures: Vec<CandidateFailure>,
52 pub winner: CandidateResult,
53 pub finalists: Vec<CandidateResult>,
54}
55
56pub fn tune_with<F>(
57 kernel: &KernelSpec,
58 constraints: &AdapterConstraints,
59 space: &ScheduleSpace,
60 config: TuningConfig,
61 mut evaluate: F,
62) -> Result<TuningResult, ForgeError>
63where
64 F: FnMut(Schedule, usize) -> Result<CandidateEvaluation, ForgeError>,
65{
66 if config.initial_samples == 0
67 || config.finalist_samples == 0
68 || config.finalist_count == 0
69 || config.max_candidates == 0
70 {
71 return Err(ForgeError::InvalidSchedule(
72 "tuning sample and candidate budgets must be non-zero".to_string(),
73 ));
74 }
75
76 let candidates = space
77 .candidates(kernel, constraints)
78 .into_iter()
79 .take(config.max_candidates)
80 .collect::<Vec<_>>();
81 if candidates.is_empty() {
82 return Err(ForgeError::InvalidSchedule(
83 "schedule space contains no adapter-compatible candidates".to_string(),
84 ));
85 }
86 let evaluated_candidates = candidates.len();
87
88 let mut accepted = Vec::new();
89 let mut failures = Vec::new();
90 let mut rejected_candidates = 0usize;
91 for schedule in candidates {
92 match evaluate(schedule, config.initial_samples) {
93 Ok(evaluation) if evaluation.oracle.passed() => {
94 if let Some(timing) =
95 TimingSummary::from_samples(evaluation.timing_source, &evaluation.samples_ns)
96 {
97 accepted.push(CandidateResult {
98 schedule,
99 oracle: evaluation.oracle,
100 timing,
101 });
102 } else {
103 rejected_candidates += 1;
104 failures.push(CandidateFailure {
105 schedule,
106 reason: "candidate produced no timing samples".to_string(),
107 });
108 }
109 }
110 Ok(evaluation) => {
111 rejected_candidates += 1;
112 failures.push(CandidateFailure {
113 schedule,
114 reason: format!(
115 "oracle rejected {} value(s), first mismatch {:?}",
116 evaluation.oracle.mismatch_count, evaluation.oracle.first_mismatch
117 ),
118 });
119 }
120 Err(error) => {
121 rejected_candidates += 1;
122 failures.push(CandidateFailure {
123 schedule,
124 reason: error.to_string(),
125 });
126 }
127 }
128 }
129 if accepted.is_empty() {
130 return Err(ForgeError::OracleMismatch(
131 "no schedule passed both oracle and timing gates".to_string(),
132 ));
133 }
134
135 sort_results(&mut accepted);
136 accepted.truncate(config.finalist_count.min(accepted.len()));
137
138 let mut finalists = Vec::with_capacity(accepted.len());
139 for initial in accepted {
140 let final_evaluation = match evaluate(initial.schedule, config.finalist_samples) {
141 Ok(evaluation) => evaluation,
142 Err(error) => {
143 rejected_candidates += 1;
144 failures.push(CandidateFailure {
145 schedule: initial.schedule,
146 reason: format!("finalist evaluation failed: {error}"),
147 });
148 continue;
149 }
150 };
151 if !final_evaluation.oracle.passed() {
152 rejected_candidates += 1;
153 failures.push(CandidateFailure {
154 schedule: initial.schedule,
155 reason: "finalist failed the CPU oracle".to_string(),
156 });
157 continue;
158 }
159 let Some(timing) = TimingSummary::from_samples(
160 final_evaluation.timing_source,
161 &final_evaluation.samples_ns,
162 ) else {
163 rejected_candidates += 1;
164 failures.push(CandidateFailure {
165 schedule: initial.schedule,
166 reason: "finalist produced no timing samples".to_string(),
167 });
168 continue;
169 };
170 finalists.push(CandidateResult {
171 schedule: initial.schedule,
172 oracle: final_evaluation.oracle,
173 timing,
174 });
175 }
176 if finalists.is_empty() {
177 return Err(ForgeError::OracleMismatch(
178 "all finalist schedules failed certification".to_string(),
179 ));
180 }
181 sort_results(&mut finalists);
182 let winner = finalists[0].clone();
183 Ok(TuningResult {
184 evaluated_candidates,
185 rejected_candidates,
186 failures,
187 winner,
188 finalists,
189 })
190}
191
192fn sort_results(results: &mut [CandidateResult]) {
193 results.sort_by(|left, right| {
194 left.timing
195 .median_ns
196 .cmp(&right.timing.median_ns)
197 .then(left.timing.p95_ns.cmp(&right.timing.p95_ns))
198 .then(left.schedule.sort_key().cmp(&right.schedule.sort_key()))
199 });
200}
201
202#[cfg(test)]
203mod tests {
204 use super::*;
205 use crate::wgsl_forge::{BuiltinKernel, OracleTolerance};
206
207 fn passing_report() -> ComparisonReport {
208 ComparisonReport {
209 compared: 16,
210 mismatch_count: 0,
211 first_mismatch: None,
212 max_absolute_error: OracleTolerance::default().absolute / 2.0,
213 max_relative_error: 0.0,
214 }
215 }
216
217 #[test]
218 fn tuner_is_deterministic_and_correctness_gated() {
219 let kernel = BuiltinKernel::AffineF32.spec();
220 let constraints = AdapterConstraints::portable();
221 let space = ScheduleSpace {
222 workgroup_sizes: vec![32, 64],
223 items_per_invocation: vec![1, 2],
224 vector_widths: vec![1],
225 };
226 let run = || {
227 tune_with(
228 &kernel,
229 &constraints,
230 &space,
231 TuningConfig {
232 initial_samples: 2,
233 finalist_samples: 3,
234 finalist_count: 2,
235 max_candidates: 4,
236 },
237 |schedule, samples| {
238 let base = 10_000 / schedule.elements_per_workgroup() as u64;
239 Ok(CandidateEvaluation {
240 oracle: passing_report(),
241 timing_source: TimingSource::Synthetic,
242 samples_ns: (0..samples).map(|index| base + index as u64).collect(),
243 })
244 },
245 )
246 .unwrap()
247 };
248 assert_eq!(run(), run());
249 assert_eq!(run().winner.schedule.workgroup_size, 64);
250 assert_eq!(run().winner.schedule.items_per_invocation, 2);
251 }
252}