Skip to main content

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/// Retention order for this controller's bounded process-memory cache.
17#[derive(Debug, Clone, Copy, PartialEq, Eq)]
18pub enum MemoryEviction {
19    /// Evict the oldest inserted key; the existing default.
20    Fifo,
21    /// Promote successful memory lookups and replace the least recently used key.
22    Lru,
23}
24/// Completed-work timing requested from the trial adapter.
25#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26pub enum Timing {
27    /// Use device profiling facilities and wait for the measured work to finish.
28    Device,
29    /// Include candidate-internal allocation, layout conversion, submission and waiting.
30    EndToEnd
31}
32/// Whether the adapter can establish numerical agreement for this trial.
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub enum Validation {
35    /// The reference/candidate comparison completed and met the tolerance.
36    Passed,
37    /// No validator supports this output or its readback size.
38    /// A numerical mismatch is an error, not this outcome.
39    Unsupported
40}
41
42/// Absolute/relative comparison limits and bounded host readback size.
43#[derive(Debug, Clone, Copy, PartialEq)]
44pub struct Tolerance {
45    /// Nonnegative finite absolute error limit; default `1e-4`.
46    pub absolute: f64,
47    /// Nonnegative finite relative error limit; default `1e-3`.
48    pub relative: f64,
49    /// Maximum combined reference/candidate readback bytes, not a GPU allocator quota.
50    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/// Immutable limits shared by operator, graph and application selection.
57/// Defaults are selection rules, not measured performance guarantees.
58#[derive(Debug, Clone)]
59pub struct StackPolicy {
60    /// Explore, cache-only or reference-only operation; defaults to Explore.
61    pub mode: Mode,
62    /// Completed timing scope; defaults to EndToEnd.
63    pub timing: Timing,
64    /// Require numerical validation before promoting a choice; defaults to true.
65    pub require_validation: bool,
66    /// Numerical limits and combined reference/candidate readback quota.
67    pub tolerance: Tolerance,
68    /// Positive warmup count in `1..=100`; default 2.
69    pub warmups: usize,
70    /// Odd sample count. Each non-reference sample is paired with a fresh reference measurement.
71    pub samples: usize,
72    /// Positive candidate trial limit up to 4096; default 32.
73    pub max_candidates: usize,
74    /// A soft deadline checked BETWEEN completed trials; running kernels are never interrupted.
75    pub budget: Duration,
76    /// Required reciprocal median timing ratio for promotion; finite and at least 1.
77    /// Default 1.05 is a threshold, not an observed speedup.
78    pub min_speedup: f64,
79    /// Maximum median absolute deviation of paired ratios divided by their median.
80    pub max_relative_mad: f64,
81    /// Optional candidate scratch-byte ceiling. Unknown estimates do not fit
82    /// an enabled limit, including for the declared reference. Default None.
83    pub workspace_limit: Option<u64>,
84    /// Maximum retained keys in `1..=65536`; default 1024.
85    pub capacity: usize,
86    /// Maximum simultaneous trials in `1..=256`; default 1, also per-device serialized.
87    pub max_parallel_tunes: usize,
88    /// Nonzero cache lifetime; default seven days. Future-dated entries expire.
89    pub ttl: Duration,
90    /// Odd nonzero regression window in `3..=101`; default 7.
91    pub regression_pairs: usize,
92    /// Median selected/reference threshold above 1; default 1.15.
93    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    /// Reject invalid counts, bounds, nonfinite numerical limits or zero durations.
109    /// Returns `FailureKind::InvalidInput` without starting a trial or touching disk.
110    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    /// Mode is intentionally excluded: an offline Explore cache is usable by CacheOnly.
129    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/// Workload level included in cache identity and dependency tracking.
138#[derive(Debug, Clone, Copy, PartialEq, Eq)]
139#[repr(u8)]
140pub enum Scope {
141    /// One complete operator, including its internal setup costs.
142    Operator = 0,
143    /// A fused graph with all externally visible outputs.
144    Graph = 1,
145    /// A complete application plan, such as greedy generation.
146    Pipeline = 2
147}
148
149/// Exact operation/environment signature; shapes are not bucketed or guessed.
150#[derive(Debug, Clone)]
151pub struct Problem {
152    /// Operator, graph or pipeline level.
153    pub scope: Scope,
154    /// Stable library/graph/model operation identity and semantic revision.
155    pub operation: String,
156    /// Backend + physical device + driver/runtime + compiler/build + capabilities.
157    pub environment: String,
158    /// Exact dtype, shape, strides, precision and operation parameters; no lossy bucketing.
159    pub workload: String,
160    /// Caller supplied load/topology/batching/power regime, not inferred from device id.
161    pub execution_context: String,
162    /// False when a reliable driver/build identity cannot be established.
163    pub persistent: bool,
164}
165/// Stable candidate identity and its declared eligibility/scratch requirement.
166#[derive(Debug, Clone)]
167pub struct Candidate {
168    /// Stable algorithm/configuration name, not an index from an old binary.
169    pub name: String,
170    /// Semantic/configuration revision invalidating old choices; default `"1"`.
171    pub revision: String,
172    /// Known temporary workspace bytes, or None when no reliable estimate exists.
173    pub workspace_bytes: Option<u64>,
174    /// Whether the existing operator eligibility rules admit this candidate.
175    pub eligible: bool,
176}
177impl Candidate {
178    /// Create an eligible revision-1 candidate without a workspace estimate.
179    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    /// Check eligibility and the optional workspace ceiling, without executing code.
188    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/// Failure classification used to distinguish rejected trials from device faults.
230#[derive(Debug, Clone, Copy, PartialEq, Eq)]
231pub enum FailureKind {
232    /// Invalid policy, workload or trial metadata.
233    InvalidInput,
234    /// A candidate/reference failed correctness or a safely completed execution.
235    Rejected,
236    /// Validation/timing is unavailable; only the declared reference may be used unchecked.
237    Unavailable,
238    /// Device work could not be confirmed complete; faults the tuning lane.
239    Device,
240    /// This environment's lane was previously faulted and rejects further selection.
241    Quarantined,
242}
243/// Structured selection failure; Display combines the kind and explanatory message.
244#[derive(Debug, Clone)]
245pub struct TuneFailure {
246    /// Failure category controlling adapter error handling.
247    pub kind: FailureKind,
248    /// Human-readable failure reason.
249    pub message: String
250}
251impl TuneFailure {
252    /// Report an invalid caller-supplied policy or signature.
253    pub fn invalid(message: impl Into<String>) -> Self { Self { kind: FailureKind::InvalidInput, message: message.into() } }
254    /// Report incorrect output or a trial rejected after confirmed completion.
255    pub fn rejected(message: impl Into<String>) -> Self { Self { kind: FailureKind::Rejected, message: message.into() } }
256    /// Report a device completion failure, not an ordinary unsupported candidate.
257    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
264/// Canonical, length-delimited encoding prevents concatenation aliases such as (ab,c)/(a,bc).
265pub 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}
285/// Encode the full workload, ordered candidate manifest, reference and numerical policy.
286/// This returns identity text, not a cryptographic hash or proof of correctness.
287pub 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}
301/// Sort positive finite ratios in place and return their median.
302/// Empty, zero, negative or nonfinite input returns None.
303pub 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}
309/// Return `(median_selected_over_reference, relative_mad)` for completed pairs.
310/// Requires at least `policy.samples` nonzero pairs and acceptable relative MAD;
311/// an unavailable/noisy score returns None, not a statistically proven speedup.
312pub 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    // Unlike timing ratios, zero deviations are valid and desirable.
318    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}