ruda_runtime/runtime/tune/stack/policy.rs
1//! Shared policy and workload vocabulary. This file also builds in the standalone host harness.
2use std::{fmt, string::{String, ToString}, time::Duration, vec::Vec};
3
4/// Selection policy for adapters participating in the shared controller.
5#[derive(Debug, Clone, Copy, PartialEq, Eq)]
6pub enum Mode {
7 /// Reuse valid choices, otherwise validate and time eligible candidates.
8 Explore,
9 /// Reuse choices without fresh timing search; misses use the reference.
10 /// Cold disk hits still require validation in the current process.
11 CacheOnly,
12 /// Use the declared reference without this controller's cache or trials.
13 /// This does not restore the legacy LocalTuner route.
14 Disabled
15}
16/// Completed-work timing requested from the trial adapter.
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum Timing {
19 /// Use device profiling facilities and wait for the measured work to finish.
20 Device,
21 /// Include candidate-internal allocation, layout conversion, submission and waiting.
22 EndToEnd
23}
24/// Whether the adapter can establish numerical agreement for this trial.
25#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum Validation {
27 /// The reference/candidate comparison completed and met the tolerance.
28 Passed,
29 /// No validator supports this output or its readback size.
30 /// A numerical mismatch is an error, not this outcome.
31 Unsupported
32}
33
34/// Absolute/relative comparison limits and bounded host readback size.
35#[derive(Debug, Clone, Copy, PartialEq)]
36pub struct Tolerance {
37 /// Nonnegative finite absolute error limit; default `1e-4`.
38 pub absolute: f64,
39 /// Nonnegative finite relative error limit; default `1e-3`.
40 pub relative: f64,
41 /// Maximum combined reference/candidate readback bytes, not a GPU allocator quota.
42 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/// Immutable limits shared by operator, graph and application selection.
49/// Defaults are selection rules, not measured performance guarantees.
50#[derive(Debug, Clone)]
51pub struct StackPolicy {
52 /// Explore, cache-only or reference-only operation; defaults to Explore.
53 pub mode: Mode,
54 /// Completed timing scope; defaults to EndToEnd.
55 pub timing: Timing,
56 /// Require numerical validation before promoting a choice; defaults to true.
57 pub require_validation: bool,
58 /// Numerical limits and combined reference/candidate readback quota.
59 pub tolerance: Tolerance,
60 /// Positive warmup count in `1..=100`; default 2.
61 pub warmups: usize,
62 /// Odd sample count. Each non-reference sample is paired with a fresh reference measurement.
63 pub samples: usize,
64 /// Positive candidate trial limit up to 4096; default 32.
65 pub max_candidates: usize,
66 /// A soft deadline checked BETWEEN completed trials; running kernels are never interrupted.
67 pub budget: Duration,
68 /// Required reciprocal median timing ratio for promotion; finite and at least 1.
69 /// Default 1.05 is a threshold, not an observed speedup.
70 pub min_speedup: f64,
71 /// Maximum median absolute deviation of paired ratios divided by their median.
72 pub max_relative_mad: f64,
73 /// Optional candidate scratch-byte ceiling. Unknown estimates do not fit
74 /// an enabled limit, including for the declared reference. Default None.
75 pub workspace_limit: Option<u64>,
76 /// Maximum retained keys in `1..=65536`; default 1024.
77 pub capacity: usize,
78 /// Maximum simultaneous trials in `1..=256`; default 1, also per-device serialized.
79 pub max_parallel_tunes: usize,
80 /// Nonzero cache lifetime; default seven days. Future-dated entries expire.
81 pub ttl: Duration,
82 /// Odd nonzero regression window in `3..=101`; default 7.
83 pub regression_pairs: usize,
84 /// Median selected/reference threshold above 1; default 1.15.
85 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 /// Reject invalid counts, bounds, nonfinite numerical limits or zero durations.
101 /// Returns `FailureKind::InvalidInput` without starting a trial or touching disk.
102 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 /// Mode is intentionally excluded: an offline Explore cache is usable by CacheOnly.
121 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/// Workload level included in cache identity and dependency tracking.
130#[derive(Debug, Clone, Copy, PartialEq, Eq)]
131#[repr(u8)]
132pub enum Scope {
133 /// One complete operator, including its internal setup costs.
134 Operator = 0,
135 /// A fused graph with all externally visible outputs.
136 Graph = 1,
137 /// A complete application plan, such as greedy generation.
138 Pipeline = 2
139}
140
141/// Exact operation/environment signature; shapes are not bucketed or guessed.
142#[derive(Debug, Clone)]
143pub struct Problem {
144 /// Operator, graph or pipeline level.
145 pub scope: Scope,
146 /// Stable library/graph/model operation identity and semantic revision.
147 pub operation: String,
148 /// Backend + physical device + driver/runtime + compiler/build + capabilities.
149 pub environment: String,
150 /// Exact dtype, shape, strides, precision and operation parameters; no lossy bucketing.
151 pub workload: String,
152 /// Caller supplied load/topology/batching/power regime, not inferred from device id.
153 pub execution_context: String,
154 /// False when a reliable driver/build identity cannot be established.
155 pub persistent: bool,
156}
157/// Stable candidate identity and its declared eligibility/scratch requirement.
158#[derive(Debug, Clone)]
159pub struct Candidate {
160 /// Stable algorithm/configuration name, not an index from an old binary.
161 pub name: String,
162 /// Semantic/configuration revision invalidating old choices; default `"1"`.
163 pub revision: String,
164 /// Known temporary workspace bytes, or None when no reliable estimate exists.
165 pub workspace_bytes: Option<u64>,
166 /// Whether the existing operator eligibility rules admit this candidate.
167 pub eligible: bool,
168}
169impl Candidate {
170 /// Create an eligible revision-1 candidate without a workspace estimate.
171 pub fn new(name: impl Into<String>) -> Self {
172 Self { name: name.into(), revision: "1".into(), workspace_bytes: None, eligible: true }
173 }
174 /// Check eligibility and the optional workspace ceiling, without executing code.
175 pub fn fits(&self, policy: &StackPolicy) -> bool {
176 self.eligible && match policy.workspace_limit {
177 Some(limit) => self.workspace_bytes.is_some_and(|n| n <= limit),
178 None => true,
179 }
180 }
181}
182
183/// Failure classification used to distinguish rejected trials from device faults.
184#[derive(Debug, Clone, Copy, PartialEq, Eq)]
185pub enum FailureKind {
186 /// Invalid policy, workload or trial metadata.
187 InvalidInput,
188 /// A candidate/reference failed correctness or a safely completed execution.
189 Rejected,
190 /// Validation/timing is unavailable; only the declared reference may be used unchecked.
191 Unavailable,
192 /// Device work could not be confirmed complete; faults the tuning lane.
193 Device,
194 /// This environment's lane was previously faulted and rejects further selection.
195 Quarantined,
196}
197/// Structured selection failure; Display combines the kind and explanatory message.
198#[derive(Debug, Clone)]
199pub struct TuneFailure {
200 /// Failure category controlling adapter error handling.
201 pub kind: FailureKind,
202 /// Human-readable failure reason.
203 pub message: String
204}
205impl TuneFailure {
206 /// Report an invalid caller-supplied policy or signature.
207 pub fn invalid(message: impl Into<String>) -> Self { Self { kind: FailureKind::InvalidInput, message: message.into() } }
208 /// Report incorrect output or a trial rejected after confirmed completion.
209 pub fn rejected(message: impl Into<String>) -> Self { Self { kind: FailureKind::Rejected, message: message.into() } }
210 /// Report a device completion failure, not an ordinary unsupported candidate.
211 pub fn device(message: impl Into<String>) -> Self { Self { kind: FailureKind::Device, message: message.into() } }
212}
213impl fmt::Display for TuneFailure {
214 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{:?}: {}", self.kind, self.message) }
215}
216impl std::error::Error for TuneFailure {}
217
218/// Canonical, length-delimited encoding prevents concatenation aliases such as (ab,c)/(a,bc).
219pub fn fields(parts: &[&str]) -> String {
220 let mut out = String::new();
221 for part in parts { out.push_str(&part.len().to_string()); out.push(':'); out.push_str(part); }
222 out
223}
224/// Encode the full workload, ordered candidate manifest, reference and numerical policy.
225/// This returns identity text, not a cryptographic hash or proof of correctness.
226pub fn cache_key(problem: &Problem, candidates: &[Candidate], reference: usize, policy: &StackPolicy) -> String {
227 let mut manifest = Vec::new();
228 for c in candidates {
229 manifest.push(fields(&[&c.name, &c.revision, &std::format!("{:?}/{}", c.workspace_bytes, c.eligible)]));
230 }
231 fields(&["ruda-stack-autotune-v1", &std::format!("{:?}", problem.scope), &problem.operation, &problem.environment, &problem.workload,
232 &problem.execution_context, &reference.to_string(), &fields(&manifest.iter().map(String::as_str).collect::<Vec<_>>()),
233 &policy.accuracy_key()])
234}
235/// Sort positive finite ratios in place and return their median.
236/// Empty, zero, negative or nonfinite input returns None.
237pub fn median(values: &mut [f64]) -> Option<f64> {
238 if values.is_empty() || values.iter().any(|v| !v.is_finite() || *v <= 0.0) { return None; }
239 values.sort_by(f64::total_cmp);
240 let m = values.len() / 2;
241 Some(if values.len() % 2 == 0 { values[m - 1] / 2.0 + values[m] / 2.0 } else { values[m] })
242}
243/// Return `(median_selected_over_reference, relative_mad)` for completed pairs.
244/// Requires at least `policy.samples` nonzero pairs and acceptable relative MAD;
245/// an unavailable/noisy score returns None, not a statistically proven speedup.
246pub fn paired_score(pairs: &[(Duration, Duration)], policy: &StackPolicy) -> Option<(f64, f64)> {
247 if pairs.len() < policy.samples || pairs.iter().any(|(a,b)| a.is_zero() || b.is_zero()) { return None; }
248 let mut ratios: Vec<_> = pairs.iter().map(|(a,b)| b.as_secs_f64()/a.as_secs_f64()).collect();
249 let center = median(&mut ratios)?;
250 let mut deviations: Vec<_> = ratios.iter().map(|r| (r-center).abs()).collect();
251 // Unlike timing ratios, zero deviations are valid and desirable.
252 deviations.sort_by(f64::total_cmp);
253 let mad = deviations[deviations.len()/2] / center;
254 if !mad.is_finite() || mad > policy.max_relative_mad { None } else { Some((center, mad)) }
255}