Skip to main content

ruda_runtime/runtime/tune/stack/
engine.rs

1//! Common selection engine for tensor operations, fused graphs and complete inference plans.
2//! GPU code is behind TrialRunner: the built-in runtime adapter uses completed device profiles
3//! or explicit synchronization, while application adapters must honor that same contract.
4use super::{cache::{DiskCache, Record, now_seconds}, policy::*};
5use std::{cell::Cell, collections::{BTreeMap, BTreeSet, VecDeque}, path::PathBuf,
6    string::{String, ToString}, sync::{Mutex, MutexGuard}, time::{Duration, Instant}, vec::Vec};
7
8std::thread_local! { static DEPTH: Cell<usize> = const { Cell::new(0) }; }
9/// Whether the current thread is inside a shared-controller trial.
10/// This is nesting state, not a query of GPU activity or other processes.
11pub fn is_tuning() -> bool { DEPTH.with(|d| d.get() != 0) }
12pub(super) struct DepthGuard;
13impl DepthGuard { pub(super) fn enter() -> Self { DEPTH.with(|d| d.set(d.get()+1)); Self } }
14impl Drop for DepthGuard { fn drop(&mut self) { DEPTH.with(|d| d.set(d.get()-1)); } }
15
16/// Benchmarks MUST use isolated state. A completed measurement includes all device work whose
17/// cost belongs to the plan. A submit-only host timestamp is not a valid measurement.
18pub trait TrialRunner {
19    /// Compare isolated outputs for the candidate indices supplied to select.
20    /// Unsupported comparison returns `Validation::Unsupported`; wrong output
21    /// or unconfirmed completion returns an appropriate failure instead.
22    fn validate(&mut self, reference: usize, candidate: usize, tolerance: Tolerance) -> Result<Validation, TuneFailure>;
23    /// Return nonzero completed-work duration for an isolated candidate trial.
24    /// Stateful requests must not be advanced by measuring their live state.
25    fn measure(&mut self, candidate: usize) -> Result<Duration, TuneFailure>;
26}
27/// Why a reference or selected implementation was returned.
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum DecisionSource {
30    /// A fresh search completed, possibly choosing the reference.
31    Tuned,
32    /// Valid process-memory record, without a fresh synchronization or validation.
33    MemoryCache,
34    /// Persistent record validated once in this process before reuse.
35    DiskCache,
36    /// Cache-only miss; use the declared reference.
37    CacheMiss,
38    /// Reference-only mode.
39    Disabled,
40    /// Another trial owns the device/key or the parallel-trial budget is full.
41    Busy,
42    /// Nested trial without a usable memory record.
43    Nested,
44    /// No supported numerical validator; use the declared reference.
45    ValidationUnavailable
46}
47/// Implementation selection, not a computed output or an instruction to replay a request.
48#[derive(Debug, Clone)]
49pub struct Decision {
50    /// Index in the candidate slice passed to select for this call.
51    pub index: usize,
52    /// Caller-declared reference index in that same slice.
53    pub reference_index: usize,
54    /// Stable selected candidate name, independent of slice indexing in another build.
55    pub name: String,
56    /// Search/cache/reference provenance.
57    pub source: DecisionSource,
58    /// True only after a validator reported Passed, never inferred from successful launch.
59    pub verified: bool,
60    /// Paired selected/reference timing ratio when available; None on unmeasured bypass.
61    pub ratio: Option<f64>,
62    /// Full canonical key used by invalidation and regression observations.
63    pub cache_key: String,
64}
65/// One candidate's local timing/validation metadata.
66#[derive(Debug, Clone)]
67pub struct CandidateReport {
68    /// Candidate identity supplied by the adapter.
69    pub name: String,
70    /// Number of completed timing pairs; the reference row reports one baseline sample.
71    pub samples: usize,
72    /// Median selected/reference ratio, or None when no acceptable score exists.
73    pub ratio: Option<f64>,
74    /// Relative median absolute deviation of paired ratios, when measured.
75    pub relative_mad: Option<f64>,
76    /// Whether numerical validation passed for this trial.
77    pub verified: bool,
78    /// Trial selection/rejection reason for diagnostics.
79    pub note: String,
80}
81/// Bounded in-memory diagnostics for one selection workload.
82#[derive(Debug, Clone)]
83pub struct TuneReport {
84    /// Stable operator/graph/application identity.
85    pub operation: String,
86    /// Backend/device/driver/build signature supplied by the adapter.
87    pub environment: String,
88    /// Exact shapes, strides, precision and operation options.
89    pub workload: String,
90    /// Caller-supplied load/topology regime, not an automatically detected topology.
91    pub execution_context: String,
92    /// Selected candidate name, including a reference selected without usable validation.
93    pub winner: String,
94    /// Time spent in fresh search; unsupported-validation diagnostics may use zero.
95    pub elapsed: Duration,
96    /// Whether the soft budget elapsed; it does not interrupt in-flight kernels.
97    pub budget_exhausted: bool,
98    /// Candidate diagnostic rows considered by this report.
99    pub candidates: Vec<CandidateReport>,
100}
101/// Controller-local counters; these are not GPU utilization or performance estimates.
102#[derive(Debug, Clone, Default)]
103pub struct Stats {
104    /// Successful process-memory lookups.
105    pub memory_hits: u64,
106    /// Persistent records validated and reused in this process.
107    pub disk_hits: u64,
108    /// Calls admitted to the cold selection path, before disk/cache-only handling.
109    pub misses: u64,
110    /// Contending calls returning the declared reference rather than waiting.
111    pub busy_fallbacks: u64,
112    /// Successful fresh searches, including reference wins.
113    pub tunes: u64,
114    /// Disk read/write/removal warnings; normal execution can continue.
115    pub cache_warnings: u64,
116    /// Explicit or regression-triggered invalidation calls.
117    pub invalidations: u64,
118    /// Device completion failures recorded by this controller.
119    pub device_failures: u64,
120}
121#[derive(Default)]
122struct State {
123    records: BTreeMap<String, Record>, order: VecDeque<String>,
124    bypass: BTreeMap<String, Scope>,
125    pending: BTreeSet<String>, devices: BTreeSet<String>, poisoned_devices: BTreeSet<String>,
126    banned: BTreeMap<String, BTreeSet<String>>, regressions: BTreeMap<String, VecDeque<f64>>,
127    reports: VecDeque<TuneReport>, stats: Stats,
128}
129/// No state mutex is held while running a candidate, validator, profile, or disk operation.
130/// Contenders use the declared reference instead of waiting (also avoids nested-tuning deadlocks).
131pub struct StackTuner { policy: StackPolicy, disk: Option<DiskCache>, state: Mutex<State> }
132struct Permit<'a> { tuner: &'a StackTuner, key: String, device: String }
133impl Drop for Permit<'_> {
134    fn drop(&mut self) { let mut s = self.tuner.lock(); s.pending.remove(&self.key); s.devices.remove(&self.device); }
135}
136impl StackTuner {
137    /// Validate policy and create an independent controller.
138    /// `None` disables disk storage. Construction performs no disk I/O and does
139    /// not install this instance as the global runtime controller.
140    pub fn new(policy: StackPolicy, cache_directory: Option<PathBuf>) -> Result<Self, TuneFailure> {
141        policy.validate()?;
142        let disk = cache_directory.map(|dir| DiskCache::new(dir, policy.capacity));
143        Ok(Self { policy, disk, state: Mutex::new(State::default()) })
144    }
145    fn lock(&self) -> MutexGuard<'_, State> { self.state.lock().unwrap_or_else(|e| e.into_inner()) }
146    /// Borrow the immutable policy; runtime reconfiguration is not supported.
147    pub fn policy(&self) -> &StackPolicy { &self.policy }
148    /// Snapshot this instance's counters without launching device work.
149    pub fn stats(&self) -> Stats { self.lock().stats.clone() }
150    /// Conservative dependency snapshot for complete pipelines. It includes all currently
151    /// cached lower-level decisions in this controller, including other workloads/devices.
152    /// That may over-invalidate a pipeline but never hides a changed lower-level choice.
153    pub fn lower_level_fingerprint(&self) -> String {
154        let s = self.lock();
155        let mut parts = Vec::new();
156        for (key, record) in &s.records {
157            if record.scope != Scope::Pipeline as u8 && record.is_fresh(now_seconds(), self.policy.ttl.as_secs().max(1)) { parts.push(fields(&[key, &record.winner])); }
158        }
159        // Bypassed operators retain their declared references, never an unchecked winner.
160        for (key, scope) in &s.bypass {
161            if *scope != Scope::Pipeline { parts.push(fields(&[key, "unverified-reference"])); }
162        }
163        parts.sort();
164        super::cache::digest(fields(&parts.iter().map(String::as_str).collect::<Vec<_>>()).as_bytes())
165    }
166    /// Clone retained search diagnostics, oldest first, bounded by min(capacity, 128).
167    /// Cache hits are counters, not new search reports.
168    pub fn reports(&self) -> Vec<TuneReport> { self.lock().reports.iter().cloned().collect() }
169    fn insert(&self, record: Record) {
170        let mut s = self.lock();
171        if !s.order.contains(&record.key) {
172            while s.order.len() >= self.policy.capacity {
173                if let Some(key) = s.order.pop_front() {
174                    s.records.remove(&key); s.banned.remove(&key); s.regressions.remove(&key); s.bypass.remove(&key);
175                }
176            }
177            s.order.push_back(record.key.clone());
178        }
179        s.records.insert(record.key.clone(), record);
180    }
181    fn allowed(&self, key: &str, c: &Candidate) -> bool {
182        c.fits(&self.policy) && !self.lock().banned.get(key).is_some_and(|b| b.contains(&c.name))
183    }
184    fn from_record(&self, r: &Record, candidates: &[Candidate], reference: usize, source: DecisionSource) -> Option<Decision> {
185        if !r.is_fresh(now_seconds(), self.policy.ttl.as_secs().max(1)) || (self.policy.require_validation && !r.verified) { return None; }
186        let index = candidates.iter().position(|c| c.name == r.winner && self.allowed(&r.key, c))?;
187        Some(Decision { index, reference_index: reference, name: r.winner.clone(), source, verified: r.verified, ratio: Some(r.ratio), cache_key: r.key.clone() })
188    }
189    fn validation(&self, runner: &mut impl TrialRunner, reference: usize, candidate: usize) -> Result<bool, TuneFailure> {
190        match runner.validate(reference, candidate, self.policy.tolerance)? {
191            Validation::Passed => Ok(true),
192            Validation::Unsupported if !self.policy.require_validation => Ok(false),
193            Validation::Unsupported => Err(TuneFailure { kind: FailureKind::Unavailable, message: "no numerical validator for this output/size".into() }),
194        }
195    }
196    fn fail_device(&self, environment: &str) {
197        let mut s = self.lock(); s.poisoned_devices.insert(environment.to_string()); s.stats.device_failures += 1;
198    }
199    /// Reuse or select an eligible implementation for one exact workload.
200    /// Candidate names must be unique/nonempty, count 1..=4096, and the reference
201    /// index must exist and fit policy. Empty operation/environment/workload or an oversized
202    /// key return InvalidInput. Trials may synchronize; memory hits do not.
203    /// A returned decision does not execute the caller's live request.
204    pub fn select(&self, problem: &Problem, candidates: &[Candidate], reference: usize, runner: &mut impl TrialRunner) -> Result<Decision, TuneFailure> {
205        if candidates.is_empty() || candidates.len() > 4096 || reference >= candidates.len()
206            || problem.operation.is_empty() || problem.environment.is_empty() || problem.workload.is_empty() {
207            return Err(TuneFailure::invalid("a workload, environment and valid reference candidate are required"));
208        }
209        let mut names = BTreeSet::new();
210        if candidates.iter().any(|c| c.name.is_empty() || c.name.len() > 4096 || !names.insert(&c.name)) {
211            return Err(TuneFailure::invalid("candidate names must be nonempty, bounded and unique"));
212        }
213        if !candidates[reference].fits(&self.policy) { return Err(TuneFailure::invalid("reference does not satisfy eligibility/workspace policy")); }
214        let key = cache_key(problem, candidates, reference, &self.policy);
215        if key.len() > 384 * 1024 { return Err(TuneFailure::invalid("autotune key exceeds 384 KiB")); }
216        let fallback = |source| Decision { index: reference, reference_index: reference, name: candidates[reference].name.clone(), source, verified: false, ratio: None, cache_key: key.clone() };
217        if self.lock().poisoned_devices.contains(&problem.environment) {
218            return Err(TuneFailure { kind: FailureKind::Quarantined, message: "device tuning lane is quarantined after an unconfirmed device completion".into() });
219        }
220        if self.policy.mode == Mode::Disabled { return Ok(fallback(DecisionSource::Disabled)); }
221        let memory = { self.lock().records.get(&key).cloned() };
222        if let Some(record) = memory {
223            if let Some(d) = self.from_record(&record, candidates, reference, DecisionSource::MemoryCache) {
224                self.lock().stats.memory_hits += 1; return Ok(d);
225            }
226        }
227        if self.lock().bypass.contains_key(&key) { return Ok(fallback(DecisionSource::ValidationUnavailable)); }
228        if is_tuning() { return Ok(fallback(DecisionSource::Nested)); }
229        let _permit = {
230            let mut s = self.lock();
231            if s.pending.contains(&key) || s.devices.contains(&problem.environment) || s.pending.len() >= self.policy.max_parallel_tunes {
232                s.stats.busy_fallbacks += 1; return Ok(fallback(DecisionSource::Busy));
233            }
234            s.pending.insert(key.clone()); s.devices.insert(problem.environment.clone()); s.stats.misses += 1;
235            Permit { tuner: self, key: key.clone(), device: problem.environment.clone() }
236        };
237        let _depth = DepthGuard::enter();
238        if problem.persistent {
239            if let Some(disk) = &self.disk {
240                match disk.load(&key) {
241                    Ok(Some(record)) => {
242                        if let Some(mut decision) = self.from_record(&record, candidates, reference, DecisionSource::DiskCache).filter(|_| record.scope == problem.scope as u8) {
243                            // A disk record is not trusted as evidence for the current inputs until
244                            // the candidate is checked once in this process. CacheOnly still validates.
245                            match self.validation(runner, reference, decision.index) {
246                                Ok(checked) => {
247                                    decision.verified = checked;
248                                    let mut record = record; record.verified = checked;
249                                    self.insert(record); self.lock().stats.disk_hits += 1; return Ok(decision);
250                                }
251                                Err(e) if e.kind == FailureKind::Device => { self.fail_device(&problem.environment); return Err(e); }
252                                Err(_) => { if disk.remove(&key).is_err() { self.lock().stats.cache_warnings += 1; } }
253                            }
254                        }
255                    }
256                    Err(_) => self.lock().stats.cache_warnings += 1,
257                    Ok(None) => {}
258                }
259            }
260        }
261        if self.policy.mode == Mode::CacheOnly { return Ok(fallback(DecisionSource::CacheMiss)); }
262        match self.explore(problem, candidates, reference, &key, runner) {
263            Ok((decision, report, record)) => {
264                if problem.persistent {
265                    if let Some(disk) = &self.disk { if disk.save(&record).is_err() { self.lock().stats.cache_warnings += 1; } }
266                }
267                self.insert(record);
268                let mut s = self.lock(); s.stats.tunes += 1;
269                if s.reports.len() >= self.policy.capacity.min(128) { s.reports.pop_front(); }
270                s.reports.push_back(report);
271                Ok(decision)
272            }
273            Err(e) if e.kind == FailureKind::Unavailable => {
274                // Reference benchmarking may exceed the validation/workspace budget. Preserve
275                // normal execution on the declared reference; never promote an unchecked winner.
276                let mut s = self.lock();
277                if !s.order.contains(&key) { s.order.push_back(key.clone()); }
278                s.bypass.insert(key.clone(), problem.scope);
279                while s.order.len() > self.policy.capacity {
280                    if let Some(old) = s.order.pop_front() { s.records.remove(&old); s.banned.remove(&old); s.regressions.remove(&old); s.bypass.remove(&old); }
281                }
282                if s.reports.len() >= self.policy.capacity.min(128) { s.reports.pop_front(); }
283                s.reports.push_back(TuneReport {
284                    operation: problem.operation.clone(), environment: problem.environment.clone(),
285                    workload: problem.workload.clone(), execution_context: problem.execution_context.clone(),
286                    winner: candidates[reference].name.clone(), elapsed: Duration::ZERO, budget_exhausted: false,
287                    candidates: std::vec![CandidateReport { name: candidates[reference].name.clone(), samples: 0,
288                        ratio: None, relative_mad: None, verified: false, note: e.message }],
289                });
290                Ok(fallback(DecisionSource::ValidationUnavailable))
291            }
292            Err(e) => { if e.kind == FailureKind::Device { self.fail_device(&problem.environment); } Err(e) }
293        }
294    }
295    fn explore(&self, problem: &Problem, candidates: &[Candidate], reference: usize, key: &str, runner: &mut impl TrialRunner) -> Result<(Decision, TuneReport, Record), TuneFailure> {
296        let started = Instant::now();
297        // The mandatory reference is validated even if its compilation takes the whole budget.
298        // Never treat a failed reference as permission to choose an unverified fast candidate.
299        for _ in 0..self.policy.warmups { runner.measure(reference)?; }
300        let reference_checked = self.validation(runner, reference, reference)?;
301        let reference_time = runner.measure(reference)?;
302        if reference_time.is_zero() { return Err(TuneFailure { kind: FailureKind::Unavailable, message: "reference has zero elapsed time".into() }); }
303        let mut winner = reference; let mut winner_ratio = 1.0; let mut winner_time = reference_time;
304        let mut winner_reference = reference_time; let mut winner_checked = reference_checked;
305        let mut reports = Vec::new();
306        reports.push(CandidateReport { name: candidates[reference].name.clone(), samples: 1, ratio: Some(1.0), relative_mad: Some(0.0), verified: reference_checked, note: "reference; retained unless a stable measured improvement qualifies".into() });
307        let mut attempted = 1; let mut exhausted = false;
308        for (index, candidate) in candidates.iter().enumerate() {
309            if index == reference { continue; }
310            let mut report = CandidateReport { name: candidate.name.clone(), samples: 0, ratio: None, relative_mad: None, verified: false, note: String::new() };
311            if !self.allowed(key, candidate) { report.note = "ineligible, unknown/over-budget workspace, or regressed candidate".into(); reports.push(report); continue; }
312            if attempted >= self.policy.max_candidates || started.elapsed() >= self.policy.budget {
313                exhausted = true; report.note = "search budget exhausted".into(); reports.push(report); continue;
314            }
315            attempted += 1;
316            let result = (|| -> Result<(bool, Vec<(Duration, Duration)>), TuneFailure> {
317                for _ in 0..self.policy.warmups {
318                    if started.elapsed() >= self.policy.budget { return Err(TuneFailure::rejected("budget exhausted during warmup")); }
319                    runner.measure(index)?;
320                }
321                let checked = self.validation(runner, reference, index)?;
322                let mut pairs = Vec::with_capacity(self.policy.samples);
323                for sample in 0..self.policy.samples {
324                    if started.elapsed() >= self.policy.budget { return Err(TuneFailure::rejected("budget exhausted before enough paired samples")); }
325                    // Alternate order to reduce monotonic warmup/clock/thermal bias.
326                    let pair = if (sample + index) % 2 == 0 {
327                        let a = runner.measure(reference)?; let b = runner.measure(index)?; (a,b)
328                    } else {
329                        let b = runner.measure(index)?; let a = runner.measure(reference)?; (a,b)
330                    };
331                    if pair.0.is_zero() || pair.1.is_zero() { return Err(TuneFailure::rejected("zero duration cannot be scored")); }
332                    pairs.push(pair);
333                }
334                Ok((checked, pairs))
335            })();
336            match result {
337                Err(e) if e.kind == FailureKind::Device => return Err(e),
338                Err(e) => { report.note = e.message; },
339                Ok((checked, pairs)) => {
340                    report.verified = checked; report.samples = pairs.len();
341                    match paired_score(&pairs, &self.policy) {
342                        None => { report.note = "timings too noisy or incomplete".into(); },
343                        Some((ratio, mad)) => {
344                            report.ratio = Some(ratio); report.relative_mad = Some(mad);
345                            report.note = "paired synchronized measurements".into();
346                            if ratio < winner_ratio && ratio * self.policy.min_speedup <= 1.0 {
347                                winner = index; winner_ratio = ratio; winner_checked = checked;
348                                let mut a: Vec<_> = pairs.iter().map(|p| p.0.as_secs_f64()).collect();
349                                let mut b: Vec<_> = pairs.iter().map(|p| p.1.as_secs_f64()).collect();
350                                winner_reference = Duration::from_secs_f64(median(&mut a).unwrap());
351                                winner_time = Duration::from_secs_f64(median(&mut b).unwrap());
352                            }
353                        }
354                    }
355                }
356            }
357            reports.push(report);
358        }
359        exhausted |= started.elapsed() >= self.policy.budget;
360        let record = Record { scope: problem.scope as u8, key: key.to_string(), winner: candidates[winner].name.clone(), created: now_seconds(),
361            reference_ns: winner_reference.as_nanos().min(u64::MAX as u128) as u64,
362            winner_ns: winner_time.as_nanos().min(u64::MAX as u128) as u64, ratio: winner_ratio, verified: winner_checked };
363        let decision = Decision { index: winner, reference_index: reference, name: record.winner.clone(), source: DecisionSource::Tuned, verified: winner_checked, ratio: Some(winner_ratio), cache_key: key.to_string() };
364        let report = TuneReport { operation: problem.operation.clone(), environment: problem.environment.clone(), workload: problem.workload.clone(), execution_context: problem.execution_context.clone(), winner: record.winner.clone(), elapsed: started.elapsed(), budget_exhausted: exhausted, candidates: reports };
365        Ok((decision, report, record))
366    }
367    /// Invalidate future selections; NEVER re-run the current stateful request after a failure.
368    /// `ban_candidate=true` additionally bans a non-reference name for this key.
369    /// Removes the persistent record when possible; disk failures increment cache warnings.
370    pub fn invalidate(&self, decision: &Decision, ban_candidate: bool) {
371        {
372            let mut s = self.lock(); s.records.remove(&decision.cache_key); s.regressions.remove(&decision.cache_key); s.bypass.remove(&decision.cache_key);
373            if ban_candidate && decision.index != decision.reference_index {
374                s.banned.entry(decision.cache_key.clone()).or_default().insert(decision.name.clone());
375            }
376            // Retain one bounded FIFO slot while the key is banned.
377            if !s.order.contains(&decision.cache_key) { s.order.push_back(decision.cache_key.clone()); }
378            while s.order.len() > self.policy.capacity {
379                if let Some(old) = s.order.pop_front() { s.records.remove(&old); s.banned.remove(&old); s.regressions.remove(&old); s.bypass.remove(&old); }
380            }
381            s.stats.invalidations += 1;
382        }
383        if let Some(disk) = &self.disk { if disk.remove(&decision.cache_key).is_err() { self.lock().stats.cache_warnings += 1; } }
384    }
385    /// Feed paired observations from a controlled replay of the SAME workload and load regime.
386    /// Ordinary production request latency is NOT a comparable baseline measurement.
387    /// Returns true only when the accumulated ratio window invalidates this decision.
388    /// Missing records or changed winner names return false; zero or unchecked
389    /// timings return InvalidInput. This observation does not revalidate TTL.
390    pub fn record_comparison(&self, decision: &Decision, reference: Duration, selected: Duration, correctness_checked: bool) -> Result<bool, TuneFailure> {
391        if !correctness_checked || reference.is_zero() || selected.is_zero() { return Err(TuneFailure::invalid("regression observations require nonzero, correctness-checked paired timings")); }
392        let should_invalidate = {
393            let mut s = self.lock();
394            if !s.records.get(&decision.cache_key).is_some_and(|r| r.winner == decision.name) { return Ok(false); }
395            let history = s.regressions.entry(decision.cache_key.clone()).or_default();
396            if history.len() == self.policy.regression_pairs { history.pop_front(); }
397            history.push_back(selected.as_secs_f64()/reference.as_secs_f64());
398            if history.len() < self.policy.regression_pairs { false } else {
399                let mut ratios: Vec<_> = history.iter().copied().collect(); median(&mut ratios).is_some_and(|r| r > self.policy.regression_ratio)
400            }
401        };
402        if should_invalidate { self.invalidate(decision, true); }
403        Ok(should_invalidate)
404    }
405}