Skip to main content

qualia_core_db/wgsl_forge/
tune.rs

1use 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}