1use std::{fmt, string::{String, ToString}, time::Duration, vec::Vec};
3
4#[derive(Debug, Clone, Copy, PartialEq, Eq)]
6pub enum Mode {
7 Explore,
9 CacheOnly,
12 Disabled
15}
16#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum Timing {
19 Device,
21 EndToEnd
23}
24#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum Validation {
27 Passed,
29 Unsupported
32}
33
34#[derive(Debug, Clone, Copy, PartialEq)]
36pub struct Tolerance {
37 pub absolute: f64,
39 pub relative: f64,
41 pub max_bytes: u64,
43}
44impl Default for Tolerance {
45 fn default() -> Self { Self { absolute: 1e-4, relative: 1e-3, max_bytes: 64 << 20 } }
46}
47
48#[derive(Debug, Clone)]
51pub struct StackPolicy {
52 pub mode: Mode,
54 pub timing: Timing,
56 pub require_validation: bool,
58 pub tolerance: Tolerance,
60 pub warmups: usize,
62 pub samples: usize,
64 pub max_candidates: usize,
66 pub budget: Duration,
68 pub min_speedup: f64,
71 pub max_relative_mad: f64,
73 pub workspace_limit: Option<u64>,
76 pub capacity: usize,
78 pub max_parallel_tunes: usize,
80 pub ttl: Duration,
82 pub regression_pairs: usize,
84 pub regression_ratio: f64,
86}
87impl Default for StackPolicy {
88 fn default() -> Self {
89 Self {
90 mode: Mode::Explore, timing: Timing::EndToEnd, require_validation: true,
91 tolerance: Tolerance::default(), warmups: 2, samples: 7,
92 max_candidates: 32, budget: Duration::from_secs(30), min_speedup: 1.05,
93 max_relative_mad: 0.15, workspace_limit: None, capacity: 1024,
94 max_parallel_tunes: 1, ttl: Duration::from_secs(7 * 24 * 3600),
95 regression_pairs: 7, regression_ratio: 1.15,
96 }
97 }
98}
99impl StackPolicy {
100 pub fn validate(&self) -> Result<(), TuneFailure> {
103 if self.warmups == 0 || self.warmups > 100 || self.samples < 3
104 || self.samples > 101 || self.samples % 2 == 0
105 || self.max_candidates == 0 || self.max_candidates > 4096
106 || self.budget.is_zero() || self.capacity == 0 || self.capacity > 65536
107 || self.max_parallel_tunes == 0 || self.max_parallel_tunes > 256
108 || self.ttl.is_zero() || self.regression_pairs < 3 || self.regression_pairs > 101
109 || self.regression_pairs % 2 == 0
110 || !self.min_speedup.is_finite() || self.min_speedup < 1.0
111 || !self.max_relative_mad.is_finite() || self.max_relative_mad < 0.0
112 || !self.regression_ratio.is_finite() || self.regression_ratio <= 1.0
113 || !self.tolerance.absolute.is_finite() || self.tolerance.absolute < 0.0
114 || !self.tolerance.relative.is_finite() || self.tolerance.relative < 0.0
115 || self.tolerance.max_bytes == 0 {
116 return Err(TuneFailure::invalid("invalid stack autotune policy"));
117 }
118 Ok(())
119 }
120 pub fn accuracy_key(&self) -> String {
122 std::format!("timing={:?};checked={};atol={:016x};rtol={:016x};workspace={:?};min_speedup={:016x};mad={:016x}",
123 self.timing, self.require_validation, self.tolerance.absolute.to_bits(),
124 self.tolerance.relative.to_bits(), self.workspace_limit,
125 self.min_speedup.to_bits(), self.max_relative_mad.to_bits())
126 }
127}
128
129#[derive(Debug, Clone, Copy, PartialEq, Eq)]
131#[repr(u8)]
132pub enum Scope {
133 Operator = 0,
135 Graph = 1,
137 Pipeline = 2
139}
140
141#[derive(Debug, Clone)]
143pub struct Problem {
144 pub scope: Scope,
146 pub operation: String,
148 pub environment: String,
150 pub workload: String,
152 pub execution_context: String,
154 pub persistent: bool,
156}
157#[derive(Debug, Clone)]
159pub struct Candidate {
160 pub name: String,
162 pub revision: String,
164 pub workspace_bytes: Option<u64>,
166 pub eligible: bool,
168}
169impl Candidate {
170 pub fn new(name: impl Into<String>) -> Self {
172 let name = name.into();
173 let view = CandidateView::new(&name);
174 let revision = view.revision.into();
175 let workspace_bytes = view.workspace_bytes;
176 let eligible = view.eligible;
177 Self { name, revision, workspace_bytes, eligible }
178 }
179 pub fn fits(&self, policy: &StackPolicy) -> bool {
181 self.view().fits(policy)
182 }
183}
184
185#[derive(Clone, Copy)]
186pub(super) struct CandidateView<'a> {
187 pub(super) name: &'a String,
188 pub(super) revision: &'a str,
189 pub(super) workspace_bytes: Option<u64>,
190 pub(super) eligible: bool,
191}
192
193impl<'a> CandidateView<'a> {
194 pub(super) fn new(name: &'a String) -> Self {
195 Self { name, revision: "1", workspace_bytes: None, eligible: true }
196 }
197
198 pub(super) fn fits(&self, policy: &StackPolicy) -> bool {
199 self.eligible && match policy.workspace_limit {
200 Some(limit) => self.workspace_bytes.is_some_and(|n| n <= limit),
201 None => true,
202 }
203 }
204}
205
206pub(super) trait CandidateSource {
207 fn view(&self) -> CandidateView<'_>;
208}
209
210impl CandidateSource for Candidate {
211 fn view(&self) -> CandidateView<'_> {
212 CandidateView { name: &self.name, revision: &self.revision,
213 workspace_bytes: self.workspace_bytes, eligible: self.eligible }
214 }
215}
216
217impl CandidateSource for CandidateView<'_> {
218 fn view(&self) -> CandidateView<'_> { *self }
219}
220
221#[derive(Debug, Clone, Copy, PartialEq, Eq)]
223pub enum FailureKind {
224 InvalidInput,
226 Rejected,
228 Unavailable,
230 Device,
232 Quarantined,
234}
235#[derive(Debug, Clone)]
237pub struct TuneFailure {
238 pub kind: FailureKind,
240 pub message: String
242}
243impl TuneFailure {
244 pub fn invalid(message: impl Into<String>) -> Self { Self { kind: FailureKind::InvalidInput, message: message.into() } }
246 pub fn rejected(message: impl Into<String>) -> Self { Self { kind: FailureKind::Rejected, message: message.into() } }
248 pub fn device(message: impl Into<String>) -> Self { Self { kind: FailureKind::Device, message: message.into() } }
250}
251impl fmt::Display for TuneFailure {
252 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{:?}: {}", self.kind, self.message) }
253}
254impl std::error::Error for TuneFailure {}
255
256pub fn fields(parts: &[&str]) -> String {
258 let mut out = String::new();
259 for part in parts { append_field(&mut out, part); }
260 out
261}
262
263fn append_field(out: &mut String, part: &str) {
264 use std::fmt::Write;
265 let _ = write!(out, "{}:", part.len());
266 out.push_str(part);
267}
268
269fn append_fields(out: &mut String, parts: &[&str]) {
270 use std::fmt::Write;
271 let length = parts.iter().map(|part| {
272 part.len() + 1 + part.len().checked_ilog10().map_or(1, |digits| digits as usize + 1)
273 }).sum::<usize>();
274 let _ = write!(out, "{length}:");
275 for part in parts { append_field(out, part); }
276}
277pub fn cache_key(problem: &Problem, candidates: &[Candidate], reference: usize, policy: &StackPolicy) -> String {
280 cache_key_candidates(problem, candidates, reference, policy)
281}
282
283pub(super) fn cache_key_candidates<C: CandidateSource>(problem: &Problem, candidates: &[C], reference: usize, policy: &StackPolicy) -> String {
284 let mut manifest = String::new();
285 for c in candidates {
286 let c = c.view();
287 append_fields(&mut manifest, &[c.name, c.revision, &std::format!("{:?}/{}", c.workspace_bytes, c.eligible)]);
288 }
289 fields(&["ruda-stack-autotune-v1", &std::format!("{:?}", problem.scope), &problem.operation, &problem.environment, &problem.workload,
290 &problem.execution_context, &reference.to_string(), &manifest,
291 &policy.accuracy_key()])
292}
293pub fn median(values: &mut [f64]) -> Option<f64> {
296 if values.is_empty() || values.iter().any(|v| !v.is_finite() || *v <= 0.0) { return None; }
297 values.sort_by(f64::total_cmp);
298 let m = values.len() / 2;
299 Some(if values.len() % 2 == 0 { values[m - 1] / 2.0 + values[m] / 2.0 } else { values[m] })
300}
301pub fn paired_score(pairs: &[(Duration, Duration)], policy: &StackPolicy) -> Option<(f64, f64)> {
305 if pairs.len() < policy.samples || pairs.iter().any(|(a,b)| a.is_zero() || b.is_zero()) { return None; }
306 let mut ratios: Vec<_> = pairs.iter().map(|(a,b)| b.as_secs_f64()/a.as_secs_f64()).collect();
307 let center = median(&mut ratios)?;
308 let mut deviations: Vec<_> = ratios.iter().map(|r| (r-center).abs()).collect();
309 deviations.sort_by(f64::total_cmp);
311 let mad = deviations[deviations.len()/2] / center;
312 if !mad.is_finite() || mad > policy.max_relative_mad { None } else { Some((center, mad)) }
313}