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#[derive(Debug, Clone, Copy, PartialEq, Eq)]
5pub enum Mode { Explore, CacheOnly, Disabled }
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
7pub enum Timing { Device, EndToEnd }
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum Validation { Passed, Unsupported }
10
11#[derive(Debug, Clone, Copy, PartialEq)]
12pub struct Tolerance {
13    pub absolute: f64,
14    pub relative: f64,
15    /// Maximum combined reference/candidate readback bytes, not a GPU allocator quota.
16    pub max_bytes: u64,
17}
18impl Default for Tolerance {
19    fn default() -> Self { Self { absolute: 1e-4, relative: 1e-3, max_bytes: 64 << 20 } }
20}
21
22#[derive(Debug, Clone)]
23pub struct StackPolicy {
24    pub mode: Mode,
25    pub timing: Timing,
26    pub require_validation: bool,
27    pub tolerance: Tolerance,
28    pub warmups: usize,
29    /// Odd sample count. Each non-reference sample is paired with a fresh reference measurement.
30    pub samples: usize,
31    pub max_candidates: usize,
32    /// A soft deadline checked BETWEEN completed trials; running kernels are never interrupted.
33    pub budget: Duration,
34    pub min_speedup: f64,
35    /// Maximum median absolute deviation of paired ratios divided by their median.
36    pub max_relative_mad: f64,
37    pub workspace_limit: Option<u64>,
38    pub capacity: usize,
39    pub max_parallel_tunes: usize,
40    pub ttl: Duration,
41    pub regression_pairs: usize,
42    pub regression_ratio: f64,
43}
44impl Default for StackPolicy {
45    fn default() -> Self {
46        Self {
47            mode: Mode::Explore, timing: Timing::EndToEnd, require_validation: true,
48            tolerance: Tolerance::default(), warmups: 2, samples: 7,
49            max_candidates: 32, budget: Duration::from_secs(30), min_speedup: 1.05,
50            max_relative_mad: 0.15, workspace_limit: None, capacity: 1024,
51            max_parallel_tunes: 1, ttl: Duration::from_secs(7 * 24 * 3600),
52            regression_pairs: 7, regression_ratio: 1.15,
53        }
54    }
55}
56impl StackPolicy {
57    pub fn validate(&self) -> Result<(), TuneFailure> {
58        if self.warmups == 0 || self.warmups > 100 || self.samples < 3
59            || self.samples > 101 || self.samples % 2 == 0
60            || self.max_candidates == 0 || self.max_candidates > 4096
61            || self.budget.is_zero() || self.capacity == 0 || self.capacity > 65536
62            || self.max_parallel_tunes == 0 || self.max_parallel_tunes > 256
63            || self.ttl.is_zero() || self.regression_pairs < 3 || self.regression_pairs > 101
64            || self.regression_pairs % 2 == 0
65            || !self.min_speedup.is_finite() || self.min_speedup < 1.0
66            || !self.max_relative_mad.is_finite() || self.max_relative_mad < 0.0
67            || !self.regression_ratio.is_finite() || self.regression_ratio <= 1.0
68            || !self.tolerance.absolute.is_finite() || self.tolerance.absolute < 0.0
69            || !self.tolerance.relative.is_finite() || self.tolerance.relative < 0.0
70            || self.tolerance.max_bytes == 0 {
71            return Err(TuneFailure::invalid("invalid stack autotune policy"));
72        }
73        Ok(())
74    }
75    /// Mode is intentionally excluded: an offline Explore cache is usable by CacheOnly.
76    pub fn accuracy_key(&self) -> String {
77        std::format!("timing={:?};checked={};atol={:016x};rtol={:016x};workspace={:?};min_speedup={:016x};mad={:016x}",
78            self.timing, self.require_validation, self.tolerance.absolute.to_bits(),
79            self.tolerance.relative.to_bits(), self.workspace_limit,
80            self.min_speedup.to_bits(), self.max_relative_mad.to_bits())
81    }
82}
83
84#[derive(Debug, Clone, Copy, PartialEq, Eq)]
85#[repr(u8)]
86pub enum Scope { Operator = 0, Graph = 1, Pipeline = 2 }
87
88#[derive(Debug, Clone)]
89pub struct Problem {
90    pub scope: Scope,
91    /// Stable library/graph/model operation identity and semantic revision.
92    pub operation: String,
93    /// Backend + physical device + driver/runtime + compiler/build + capabilities.
94    pub environment: String,
95    /// Exact dtype, shape, strides, precision and operation parameters; no lossy bucketing.
96    pub workload: String,
97    /// Caller supplied load/topology/batching/power regime, not inferred from device id.
98    pub execution_context: String,
99    /// False when a reliable driver/build identity cannot be established.
100    pub persistent: bool,
101}
102#[derive(Debug, Clone)]
103pub struct Candidate {
104    /// Stable algorithm/configuration name, not an index from an old binary.
105    pub name: String,
106    pub revision: String,
107    pub workspace_bytes: Option<u64>,
108    pub eligible: bool,
109}
110impl Candidate {
111    pub fn new(name: impl Into<String>) -> Self {
112        Self { name: name.into(), revision: "1".into(), workspace_bytes: None, eligible: true }
113    }
114    pub fn fits(&self, policy: &StackPolicy) -> bool {
115        self.eligible && match policy.workspace_limit {
116            Some(limit) => self.workspace_bytes.is_some_and(|n| n <= limit),
117            None => true,
118        }
119    }
120}
121
122#[derive(Debug, Clone, Copy, PartialEq, Eq)]
123pub enum FailureKind {
124    InvalidInput,
125    /// A candidate/reference failed correctness or a safely completed execution.
126    Rejected,
127    /// Validation/timing is unavailable; only the declared reference may be used unchecked.
128    Unavailable,
129    Device,
130    Quarantined,
131}
132#[derive(Debug, Clone)]
133pub struct TuneFailure { pub kind: FailureKind, pub message: String }
134impl TuneFailure {
135    pub fn invalid(message: impl Into<String>) -> Self { Self { kind: FailureKind::InvalidInput, message: message.into() } }
136    pub fn rejected(message: impl Into<String>) -> Self { Self { kind: FailureKind::Rejected, message: message.into() } }
137    pub fn device(message: impl Into<String>) -> Self { Self { kind: FailureKind::Device, message: message.into() } }
138}
139impl fmt::Display for TuneFailure {
140    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { write!(f, "{:?}: {}", self.kind, self.message) }
141}
142impl std::error::Error for TuneFailure {}
143
144/// Canonical, length-delimited encoding prevents concatenation aliases such as (ab,c)/(a,bc).
145pub fn fields(parts: &[&str]) -> String {
146    let mut out = String::new();
147    for part in parts { out.push_str(&part.len().to_string()); out.push(':'); out.push_str(part); }
148    out
149}
150pub fn cache_key(problem: &Problem, candidates: &[Candidate], reference: usize, policy: &StackPolicy) -> String {
151    let mut manifest = Vec::new();
152    for c in candidates {
153        manifest.push(fields(&[&c.name, &c.revision, &std::format!("{:?}/{}", c.workspace_bytes, c.eligible)]));
154    }
155    fields(&["ruda-stack-autotune-v1", &std::format!("{:?}", problem.scope), &problem.operation, &problem.environment, &problem.workload,
156        &problem.execution_context, &reference.to_string(), &fields(&manifest.iter().map(String::as_str).collect::<Vec<_>>()),
157        &policy.accuracy_key()])
158}
159pub fn median(values: &mut [f64]) -> Option<f64> {
160    if values.is_empty() || values.iter().any(|v| !v.is_finite() || *v <= 0.0) { return None; }
161    values.sort_by(f64::total_cmp);
162    let m = values.len() / 2;
163    Some(if values.len() % 2 == 0 { values[m - 1] / 2.0 + values[m] / 2.0 } else { values[m] })
164}
165pub fn paired_score(pairs: &[(Duration, Duration)], policy: &StackPolicy) -> Option<(f64, f64)> {
166    if pairs.len() < policy.samples || pairs.iter().any(|(a,b)| a.is_zero() || b.is_zero()) { return None; }
167    let mut ratios: Vec<_> = pairs.iter().map(|(a,b)| b.as_secs_f64()/a.as_secs_f64()).collect();
168    let center = median(&mut ratios)?;
169    let mut deviations: Vec<_> = ratios.iter().map(|r| (r-center).abs()).collect();
170    // Unlike timing ratios, zero deviations are valid and desirable.
171    deviations.sort_by(f64::total_cmp);
172    let mad = deviations[deviations.len()/2] / center;
173    if !mad.is_finite() || mad > policy.max_relative_mad { None } else { Some((center, mad)) }
174}