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 MemoryEviction {
19 Fifo,
21 Lru,
23}
24#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum Timing {
27 Device,
29 EndToEnd
31}
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub enum Validation {
35 Passed,
37 Unsupported
40}
41
42#[derive(Debug, Clone, Copy, PartialEq)]
44pub struct Tolerance {
45 pub absolute: f64,
47 pub relative: f64,
49 pub max_bytes: u64,
51}
52impl Default for Tolerance {
53 fn default() -> Self { Self { absolute: 1e-4, relative: 1e-3, max_bytes: 64 << 20 } }
54}
55
56#[derive(Debug, Clone)]
59pub struct StackPolicy {
60 pub mode: Mode,
62 pub timing: Timing,
64 pub require_validation: bool,
66 pub tolerance: Tolerance,
68 pub warmups: usize,
70 pub samples: usize,
72 pub max_candidates: usize,
74 pub budget: Duration,
76 pub min_speedup: f64,
79 pub max_relative_mad: f64,
81 pub workspace_limit: Option<u64>,
84 pub capacity: usize,
86 pub max_parallel_tunes: usize,
88 pub ttl: Duration,
90 pub regression_pairs: usize,
92 pub regression_ratio: f64,
94}
95impl Default for StackPolicy {
96 fn default() -> Self {
97 Self {
98 mode: Mode::Explore, timing: Timing::EndToEnd, require_validation: true,
99 tolerance: Tolerance::default(), warmups: 2, samples: 7,
100 max_candidates: 32, budget: Duration::from_secs(30), min_speedup: 1.05,
101 max_relative_mad: 0.15, workspace_limit: None, capacity: 1024,
102 max_parallel_tunes: 1, ttl: Duration::from_secs(7 * 24 * 3600),
103 regression_pairs: 7, regression_ratio: 1.15,
104 }
105 }
106}
107impl StackPolicy {
108 pub fn validate(&self) -> Result<(), TuneFailure> {
111 if self.warmups == 0 || self.warmups > 100 || self.samples < 3
112 || self.samples > 101 || self.samples % 2 == 0
113 || self.max_candidates == 0 || self.max_candidates > 4096
114 || self.budget.is_zero() || self.capacity == 0 || self.capacity > 65536
115 || self.max_parallel_tunes == 0 || self.max_parallel_tunes > 256
116 || self.ttl.is_zero() || self.regression_pairs < 3 || self.regression_pairs > 101
117 || self.regression_pairs % 2 == 0
118 || !self.min_speedup.is_finite() || self.min_speedup < 1.0
119 || !self.max_relative_mad.is_finite() || self.max_relative_mad < 0.0
120 || !self.regression_ratio.is_finite() || self.regression_ratio <= 1.0
121 || !self.tolerance.absolute.is_finite() || self.tolerance.absolute < 0.0
122 || !self.tolerance.relative.is_finite() || self.tolerance.relative < 0.0
123 || self.tolerance.max_bytes == 0 {
124 return Err(TuneFailure::invalid("invalid stack autotune policy"));
125 }
126 Ok(())
127 }
128 pub fn accuracy_key(&self) -> String {
130 std::format!("timing={:?};checked={};atol={:016x};rtol={:016x};workspace={:?};min_speedup={:016x};mad={:016x}",
131 self.timing, self.require_validation, self.tolerance.absolute.to_bits(),
132 self.tolerance.relative.to_bits(), self.workspace_limit,
133 self.min_speedup.to_bits(), self.max_relative_mad.to_bits())
134 }
135}
136
137#[derive(Debug, Clone, Copy, PartialEq, Eq)]
139#[repr(u8)]
140pub enum Scope {
141 Operator = 0,
143 Graph = 1,
145 Pipeline = 2
147}
148
149#[derive(Debug, Clone)]
151pub struct Problem {
152 pub scope: Scope,
154 pub operation: String,
156 pub environment: String,
158 pub workload: String,
160 pub execution_context: String,
162 pub persistent: bool,
164}
165#[derive(Debug, Clone)]
167pub struct Candidate {
168 pub name: String,
170 pub revision: String,
172 pub workspace_bytes: Option<u64>,
174 pub eligible: bool,
176}
177impl Candidate {
178 pub fn new(name: impl Into<String>) -> Self {
180 let name = name.into();
181 let view = CandidateView::new(&name);
182 let revision = view.revision.into();
183 let workspace_bytes = view.workspace_bytes;
184 let eligible = view.eligible;
185 Self { name, revision, workspace_bytes, eligible }
186 }
187 pub fn fits(&self, policy: &StackPolicy) -> bool {
189 self.view().fits(policy)
190 }
191}
192
193#[derive(Clone, Copy)]
194pub(super) struct CandidateView<'a> {
195 pub(super) name: &'a String,
196 pub(super) revision: &'a str,
197 pub(super) workspace_bytes: Option<u64>,
198 pub(super) eligible: bool,
199}
200
201impl<'a> CandidateView<'a> {
202 pub(super) fn new(name: &'a String) -> Self {
203 Self { name, revision: "1", workspace_bytes: None, eligible: true }
204 }
205
206 pub(super) fn fits(&self, policy: &StackPolicy) -> bool {
207 self.eligible && match policy.workspace_limit {
208 Some(limit) => self.workspace_bytes.is_some_and(|n| n <= limit),
209 None => true,
210 }
211 }
212}
213
214pub(super) trait CandidateSource {
215 fn view(&self) -> CandidateView<'_>;
216}
217
218impl CandidateSource for Candidate {
219 fn view(&self) -> CandidateView<'_> {
220 CandidateView { name: &self.name, revision: &self.revision,
221 workspace_bytes: self.workspace_bytes, eligible: self.eligible }
222 }
223}
224
225impl CandidateSource for CandidateView<'_> {
226 fn view(&self) -> CandidateView<'_> { *self }
227}
228
229#[derive(Debug, Clone, Copy, PartialEq, Eq)]
231pub enum FailureKind {
232 InvalidInput,
234 Rejected,
236 Unavailable,
238 Device,
240 Quarantined,
242}
243#[derive(Debug, Clone)]
245pub struct TuneFailure {
246 pub kind: FailureKind,
248 pub message: String
250}
251impl TuneFailure {
252 pub fn invalid(message: impl Into<String>) -> Self { Self { kind: FailureKind::InvalidInput, message: message.into() } }
254 pub fn rejected(message: impl Into<String>) -> Self { Self { kind: FailureKind::Rejected, message: message.into() } }
256 pub fn device(message: impl Into<String>) -> Self { Self { kind: FailureKind::Device, message: message.into() } }
258}
259impl fmt::Display for TuneFailure {
260 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{:?}: {}", self.kind, self.message) }
261}
262impl std::error::Error for TuneFailure {}
263
264pub fn fields(parts: &[&str]) -> String {
266 let mut out = String::new();
267 for part in parts { append_field(&mut out, part); }
268 out
269}
270
271fn append_field(out: &mut String, part: &str) {
272 use std::fmt::Write;
273 let _ = write!(out, "{}:", part.len());
274 out.push_str(part);
275}
276
277fn append_fields(out: &mut String, parts: &[&str]) {
278 use std::fmt::Write;
279 let length = parts.iter().map(|part| {
280 part.len() + 1 + part.len().checked_ilog10().map_or(1, |digits| digits as usize + 1)
281 }).sum::<usize>();
282 let _ = write!(out, "{length}:");
283 for part in parts { append_field(out, part); }
284}
285pub fn cache_key(problem: &Problem, candidates: &[Candidate], reference: usize, policy: &StackPolicy) -> String {
288 cache_key_candidates(problem, candidates, reference, policy)
289}
290
291pub(super) fn cache_key_candidates<C: CandidateSource>(problem: &Problem, candidates: &[C], reference: usize, policy: &StackPolicy) -> String {
292 let mut manifest = String::new();
293 for c in candidates {
294 let c = c.view();
295 append_fields(&mut manifest, &[c.name, c.revision, &std::format!("{:?}/{}", c.workspace_bytes, c.eligible)]);
296 }
297 fields(&["ruda-stack-autotune-v1", &std::format!("{:?}", problem.scope), &problem.operation, &problem.environment, &problem.workload,
298 &problem.execution_context, &reference.to_string(), &manifest,
299 &policy.accuracy_key()])
300}
301pub fn median(values: &mut [f64]) -> Option<f64> {
304 if values.is_empty() || values.iter().any(|v| !v.is_finite() || *v <= 0.0) { return None; }
305 values.sort_by(f64::total_cmp);
306 let m = values.len() / 2;
307 Some(if values.len() % 2 == 0 { values[m - 1] / 2.0 + values[m] / 2.0 } else { values[m] })
308}
309pub fn paired_score(pairs: &[(Duration, Duration)], policy: &StackPolicy) -> Option<(f64, f64)> {
313 if pairs.len() < policy.samples || pairs.iter().any(|(a,b)| a.is_zero() || b.is_zero()) { return None; }
314 let mut ratios: Vec<_> = pairs.iter().map(|(a,b)| b.as_secs_f64()/a.as_secs_f64()).collect();
315 let center = median(&mut ratios)?;
316 let mut deviations: Vec<_> = ratios.iter().map(|r| (r-center).abs()).collect();
317 deviations.sort_by(f64::total_cmp);
319 let mad = deviations[deviations.len()/2] / center;
320 if !mad.is_finite() || mad > policy.max_relative_mad { None } else { Some((center, mad)) }
321}