Skip to main content

ruda_runtime/runtime/tune/stack/
runtime_adapter.rs

1use super::*;
2use super::policy::CandidateView;
3use crate::runtime::{backend::Runtime, client::ComputeClient, tune::{AutotuneKey, AutotuneOutput, TunableSet, TuneInputs, AutotuneError}};
4use std::{any::type_name, path::PathBuf, string::{String, ToString}, sync::{Arc, OnceLock}, time::{Duration, Instant, SystemTime, UNIX_EPOCH}, vec::Vec};
5
6static GLOBAL: OnceLock<StackTuner> = OnceLock::new();
7static SESSION: OnceLock<String> = OnceLock::new();
8/// Install once, before worker threads/model loading. No implicit environment mutation or I/O.
9pub fn enable_stack_autotune(policy: StackPolicy, cache_directory: Option<PathBuf>) -> Result<&'static StackTuner, TuneFailure> {
10    enable_stack_autotune_with_memory_eviction(policy, cache_directory, MemoryEviction::Fifo)
11}
12/// Install a controller with explicit memory eviction; disk behavior and trials stay unchanged.
13pub fn enable_stack_autotune_with_memory_eviction(policy: StackPolicy, cache_directory: Option<PathBuf>, memory_eviction: MemoryEviction) -> Result<&'static StackTuner, TuneFailure> {
14    let tuner = StackTuner::new_with_memory_eviction(policy, cache_directory, memory_eviction)?;
15    GLOBAL.set(tuner).map_err(|_| TuneFailure::invalid("stack autotuning was already configured"))?;
16    Ok(GLOBAL.get().expect("controller was initialized"))
17}
18/// Return the installed controller, or None to preserve the legacy tuner route.
19/// Merely querying it does not create a controller, inspect files or run trials.
20pub fn stack_autotuner() -> Option<&'static StackTuner> { GLOBAL.get() }
21
22/// Supply extra tags for a deployment's compiler flags and load/topology/power regime.
23/// Set these BEFORE constructing devices/starting inference; changing them live is unsupported.
24#[derive(Debug, Clone)]
25pub struct RuntimeEnvironment {
26    /// Canonical backend/device/driver/build/runtime-options signature.
27    pub fingerprint: String,
28    /// Whether a driver identity and compiled build identity permit disk reuse.
29    pub persistent: bool,
30    /// Deployment-supplied regime, defaulting to isolated-single-device.
31    /// The default does not establish actual isolation or topology.
32    pub execution_context: String
33}
34/// Probe and memoize immutable environment identity for this runtime/client.
35/// Deployment tags must be set before first use; live mutation is not supported.
36/// Missing identity yields session-only caching rather than a fabricated driver version.
37pub fn runtime_environment<R: Runtime>(client: &ComputeClient<R>) -> RuntimeEnvironment {
38    use std::{any::TypeId, collections::BTreeMap, sync::RwLock};
39    type Key = (TypeId, u16, u16, u64);
40    static ENVIRONMENTS: OnceLock<RwLock<BTreeMap<Key, RuntimeEnvironment>>> = OnceLock::new();
41    let id = client.device_id();
42    let key = (TypeId::of::<R>(), id.type_id, id.index_id, client.properties_fingerprint());
43    let cache = ENVIRONMENTS.get_or_init(|| RwLock::new(BTreeMap::new()));
44    if let Some(env) = cache.read().unwrap_or_else(|p| p.into_inner()).get(&key).cloned() { return env; }
45    let env = probe_environment(client);
46    let mut cache = cache.write().unwrap_or_else(|p| p.into_inner());
47    // Device/runtime metadata is immutable after initialization. Keep this memo bounded.
48    if cache.len() < 256 { cache.insert(key, env.clone()); }
49    env
50}
51fn probe_environment<R: Runtime>(client: &ComputeClient<R>) -> RuntimeEnvironment {
52    let driver_override = std::env::var("RUDA_AUTOTUNE_DRIVER_TAG").ok().filter(|s| !s.is_empty());
53    let driver = match (R::autotune_driver_fingerprint(client), driver_override) {
54        (Some(probed), Some(extra)) => Some(fields(&[&probed, &extra])),
55        (Some(probed), None) => Some(probed),
56        (None, Some(asserted)) => Some(asserted),
57        (None, None) => None,
58    };
59    let persistent = driver.is_some() && env!("RUDA_STACK_BUILD_ID") != "unavailable";
60    let driver = driver.unwrap_or_else(|| SESSION.get_or_init(|| std::format!("session-only:{}:{}", std::process::id(), SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_nanos())).clone());
61    let build_extra = std::env::var("RUDA_AUTOTUNE_BUILD_TAG").unwrap_or_default();
62    let execution_context = std::env::var("RUDA_AUTOTUNE_CONTEXT_TAG").unwrap_or_else(|_| "isolated-single-device".into());
63    use crate::runtime::config::{RudaRuntimeConfig, RuntimeConfig};
64    let config = RudaRuntimeConfig::get();
65    let runtime_options = std::format!("compilation={:?};streaming={:?};memory={:?}", config.compilation, config.streaming, config.memory);
66    let fingerprint = fields(&[R::name(client), &std::format!("{:?}", client.device_id()), &std::format!("{:016x}", client.properties_fingerprint()),
67        &driver, type_name::<R::Compiler>(), env!("RUDA_STACK_BUILD_ID"), &build_extra,
68        &std::format!("{:?}", client.info()), &runtime_options]);
69    RuntimeEnvironment { fingerprint, persistent, execution_context }
70}
71
72fn complete<R: Runtime>(client: &ComputeClient<R>) -> Result<(), TuneFailure> {
73    let submitted = client.flush();
74    // Completion errors are never demoted to a failed candidate followed by another launch.
75    ruda_core::future::block_on(client.sync()).map_err(|e| TuneFailure::device(std::format!("autotune synchronization failed: {e:?}")))?;
76    submitted.map_err(|e| TuneFailure::rejected(std::format!("candidate submission failed after a successful completion fence: {e:?}")))
77}
78fn run_output<R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput>(
79    client: &ComputeClient<R>, set: &TunableSet<K,I,O>, index: usize, input: I::At<'_>,
80) -> Result<O, TuneFailure> {
81    let result = set.fastest(index).execute(input);
82    complete(client)?;
83    result.map_err(|e| TuneFailure::rejected(std::format!("candidate rejected: {e:?}")))
84}
85struct RuntimeTrials<'set, 'inp, R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput> {
86    client: ComputeClient<R>, set: &'set TunableSet<K,I,O>, key: &'set K,
87    input: &'set I::At<'inp>, indices: &'set [usize], timing: Timing,
88}
89impl<R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput> TrialRunner for RuntimeTrials<'_, '_, R,K,I,O> {
90    fn validate(&mut self, reference: usize, candidate: usize, tolerance: Tolerance) -> Result<Validation, TuneFailure> {
91        let a = self.set.generate_inputs(self.key, self.input);
92        let b = self.set.generate_inputs(self.key, self.input);
93        complete(&self.client)?;
94        let client = self.client.clone(); let set = self.set;
95        let ri = self.indices[reference]; let ci = self.indices[candidate];
96        self.client.exclusive(move || {
97            let _depth = super::engine::DepthGuard::enter();
98            let expected = run_output(&client, set, ri, a)?;
99            let actual = run_output(&client, set, ci, b)?;
100            let result = expected.validate_for_tuning(&actual, tolerance.absolute, tolerance.relative, tolerance.max_bytes);
101            // Readback/contiguous conversion can also launch kernels. Fence before accepting a
102            // validation mismatch or freeing trial state.
103            complete(&client)?;
104            match result {
105                Ok(true) => Ok(Validation::Passed), Ok(false) => Ok(Validation::Unsupported),
106                Err(e) => Err(TuneFailure::rejected(e)),
107            }
108        }).map_err(|e| TuneFailure::device(std::format!("exclusive validation failed: {e:?}")))?
109    }
110    fn measure(&mut self, candidate: usize) -> Result<Duration, TuneFailure> {
111        let input = self.set.generate_inputs(self.key, self.input);
112        complete(&self.client)?;
113        let client = self.client.clone(); let set = self.set;
114        let index = self.indices[candidate]; let timing = self.timing;
115        self.client.exclusive(move || {
116            let _depth = super::engine::DepthGuard::enter();
117            // Preparation/sandbox copies are outside the measured region. All allocations and
118            // relayouts made INSIDE the actual candidate operation are part of EndToEnd timing.
119            complete(&client)?;
120            match timing {
121                Timing::EndToEnd => {
122                    let start = Instant::now();
123                    let out = run_output(&client, set, index, input)?;
124                    let elapsed = start.elapsed();
125                    std::hint::black_box(&out); drop(out);
126                    Ok(elapsed)
127                }
128                Timing::Device => {
129                    let profiled = client.profile(move || set.fastest(index).execute(input), "stack-autotune");
130                    complete(&client)?;
131                    let (output, profile) = profiled.map_err(|e| TuneFailure::rejected(std::format!("device profile failed: {e:?}")))?;
132                    let output = output.map_err(|e| TuneFailure::rejected(std::format!("candidate rejected: {e:?}")))?;
133                    let elapsed = ruda_core::future::block_on(profile.resolve()).duration();
134                    std::hint::black_box(&output); drop(output);
135                    Ok(elapsed)
136                }
137            }
138        }).map_err(|e| TuneFailure::device(std::format!("exclusive measurement failed: {e:?}")))?
139    }
140}
141
142/// Fallible full-stack path. Actual request execution is performed ONCE after selection; an
143/// execution error invalidates future choices but is never silently replayed on another kernel.
144/// Requires an installed controller, an in-range explicit reference and an exact workload
145/// signature on the set. Missing metadata returns AutotuneError before live execution.
146pub fn try_execute_stack<'a, R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput>(
147    name: &str, device_id: &str, client: &ComputeClient<R>, set: Arc<TunableSet<K,I,O>>, input: I::At<'a>,
148) -> Result<O, AutotuneError> {
149    let tuner = stack_autotuner().ok_or_else(|| autotune_error(name, "stack autotuning is not enabled"))?;
150    let reference = set.stack_reference().ok_or_else(|| autotune_error(name, "no explicit reference/workload signature was registered"))?;
151    if reference >= set.len() { return Err(autotune_error(name, "reference index out of range")); }
152    let key = set.generate_key(&input);
153    let env = runtime_environment(client);
154    let problem = Problem { scope: if name.contains("fusion") { Scope::Graph } else { Scope::Operator }, operation: fields(&[name, device_id, set.stack_checksum_ref()]), environment: env.fingerprint,
155        workload: set.stack_workload(&input).ok_or_else(|| autotune_error(name, "no exact workload signature"))?,
156        execution_context: env.execution_context, persistent: env.persistent };
157    // Enumerate EVERY eligible priority group. Group order remains a useful bounded-search
158    // heuristic but a first valid group no longer terminates the search.
159    let mut indices = std::vec![reference]; let mut plan = set.plan(&key);
160    let mut included = std::vec![false; set.len()];
161    included[reference] = true;
162    loop { let batch = plan.next(None); if batch.is_empty() { break; }
163        for index in batch { if !included[index] { indices.push(index); included[index] = true; } }
164    }
165    let candidates: Vec<_> = indices.iter().map(|&i| CandidateView::new(&set.fastest(i).name)).collect();
166    // Cache hits must retain normal asynchronous execution. TrialRunner fences before every
167    // validation/measurement; do NOT synchronize ordinary requests just to read this cache.
168    let mut trials = RuntimeTrials { client: client.clone(), set: &set, key: &key, input: &input, indices: &indices, timing: tuner.policy().timing };
169    let decision = tuner.select_candidates(&problem, &candidates, 0, &mut trials).map_err(|e| autotune_error(name, &e.to_string()))?;
170    let output = set.fastest(indices[decision.index]).execute(input);
171    if output.is_err() { tuner.invalidate(&decision, true); }
172    output
173}
174fn autotune_error(name: &str, message: &str) -> AutotuneError {
175    AutotuneError::Unknown { name: name.to_string(), err: message.to_string() }
176}