1use super::*;
2use crate::runtime::{backend::Runtime, client::ComputeClient, tune::{AutotuneKey, AutotuneOutput, TunableSet, TuneInputs, AutotuneError}};
3use std::{any::type_name, path::PathBuf, string::{String, ToString}, sync::{Arc, OnceLock}, time::{Duration, Instant, SystemTime, UNIX_EPOCH}, vec::Vec};
4
5static GLOBAL: OnceLock<StackTuner> = OnceLock::new();
6static SESSION: OnceLock<String> = OnceLock::new();
7pub fn enable_stack_autotune(policy: StackPolicy, cache_directory: Option<PathBuf>) -> Result<&'static StackTuner, TuneFailure> {
9 let tuner = StackTuner::new(policy, cache_directory)?;
10 GLOBAL.set(tuner).map_err(|_| TuneFailure::invalid("stack autotuning was already configured"))?;
11 Ok(GLOBAL.get().expect("controller was initialized"))
12}
13pub fn stack_autotuner() -> Option<&'static StackTuner> { GLOBAL.get() }
14
15#[derive(Debug, Clone)]
18pub struct RuntimeEnvironment { pub fingerprint: String, pub persistent: bool, pub execution_context: String }
19pub fn runtime_environment<R: Runtime>(client: &ComputeClient<R>) -> RuntimeEnvironment {
20 use std::{any::TypeId, collections::BTreeMap, sync::Mutex};
21 type Key = (TypeId, u16, u16, u64);
22 static ENVIRONMENTS: OnceLock<Mutex<BTreeMap<Key, RuntimeEnvironment>>> = OnceLock::new();
23 let id = client.device_id();
24 let key = (TypeId::of::<R>(), id.type_id, id.index_id, client.properties_fingerprint());
25 let cache = ENVIRONMENTS.get_or_init(|| Mutex::new(BTreeMap::new()));
26 if let Some(env) = cache.lock().unwrap_or_else(|p| p.into_inner()).get(&key).cloned() { return env; }
27 let env = probe_environment(client);
28 let mut cache = cache.lock().unwrap_or_else(|p| p.into_inner());
29 if cache.len() < 256 { cache.insert(key, env.clone()); }
31 env
32}
33fn probe_environment<R: Runtime>(client: &ComputeClient<R>) -> RuntimeEnvironment {
34 let driver_override = std::env::var("RUDA_AUTOTUNE_DRIVER_TAG").ok().filter(|s| !s.is_empty());
35 let driver = match (R::autotune_driver_fingerprint(client), driver_override) {
36 (Some(probed), Some(extra)) => Some(fields(&[&probed, &extra])),
37 (Some(probed), None) => Some(probed),
38 (None, Some(asserted)) => Some(asserted),
39 (None, None) => None,
40 };
41 let persistent = driver.is_some() && env!("RUDA_STACK_BUILD_ID") != "unavailable";
42 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());
43 let build_extra = std::env::var("RUDA_AUTOTUNE_BUILD_TAG").unwrap_or_default();
44 let execution_context = std::env::var("RUDA_AUTOTUNE_CONTEXT_TAG").unwrap_or_else(|_| "isolated-single-device".into());
45 use crate::runtime::config::{RudaRuntimeConfig, RuntimeConfig};
46 let config = RudaRuntimeConfig::get();
47 let runtime_options = std::format!("compilation={:?};streaming={:?};memory={:?}", config.compilation, config.streaming, config.memory);
48 let fingerprint = fields(&[R::name(client), &std::format!("{:?}", client.device_id()), &std::format!("{:016x}", client.properties_fingerprint()),
49 &driver, type_name::<R::Compiler>(), env!("RUDA_STACK_BUILD_ID"), &build_extra,
50 &std::format!("{:?}", client.info()), &runtime_options]);
51 RuntimeEnvironment { fingerprint, persistent, execution_context }
52}
53
54fn complete<R: Runtime>(client: &ComputeClient<R>) -> Result<(), TuneFailure> {
55 let submitted = client.flush();
56 ruda_core::future::block_on(client.sync()).map_err(|e| TuneFailure::device(std::format!("autotune synchronization failed: {e:?}")))?;
58 submitted.map_err(|e| TuneFailure::rejected(std::format!("candidate submission failed after a successful completion fence: {e:?}")))
59}
60fn run_output<R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput>(
61 client: &ComputeClient<R>, set: &TunableSet<K,I,O>, index: usize, input: I::At<'_>,
62) -> Result<O, TuneFailure> {
63 let result = set.fastest(index).execute(input);
64 complete(client)?;
65 result.map_err(|e| TuneFailure::rejected(std::format!("candidate rejected: {e:?}")))
66}
67struct RuntimeTrials<'set, 'inp, R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput> {
68 client: ComputeClient<R>, set: &'set TunableSet<K,I,O>, key: &'set K,
69 input: &'set I::At<'inp>, indices: &'set [usize], timing: Timing,
70}
71impl<R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput> TrialRunner for RuntimeTrials<'_, '_, R,K,I,O> {
72 fn validate(&mut self, reference: usize, candidate: usize, tolerance: Tolerance) -> Result<Validation, TuneFailure> {
73 let a = self.set.generate_inputs(self.key, self.input);
74 let b = self.set.generate_inputs(self.key, self.input);
75 complete(&self.client)?;
76 let client = self.client.clone(); let set = self.set;
77 let ri = self.indices[reference]; let ci = self.indices[candidate];
78 self.client.exclusive(move || {
79 let _depth = super::engine::DepthGuard::enter();
80 let expected = run_output(&client, set, ri, a)?;
81 let actual = run_output(&client, set, ci, b)?;
82 let result = expected.validate_for_tuning(&actual, tolerance.absolute, tolerance.relative, tolerance.max_bytes);
83 complete(&client)?;
86 match result {
87 Ok(true) => Ok(Validation::Passed), Ok(false) => Ok(Validation::Unsupported),
88 Err(e) => Err(TuneFailure::rejected(e)),
89 }
90 }).map_err(|e| TuneFailure::device(std::format!("exclusive validation failed: {e:?}")))?
91 }
92 fn measure(&mut self, candidate: usize) -> Result<Duration, TuneFailure> {
93 let input = self.set.generate_inputs(self.key, self.input);
94 complete(&self.client)?;
95 let client = self.client.clone(); let set = self.set;
96 let index = self.indices[candidate]; let timing = self.timing;
97 self.client.exclusive(move || {
98 let _depth = super::engine::DepthGuard::enter();
99 complete(&client)?;
102 match timing {
103 Timing::EndToEnd => {
104 let start = Instant::now();
105 let out = run_output(&client, set, index, input)?;
106 let elapsed = start.elapsed();
107 std::hint::black_box(&out); drop(out);
108 Ok(elapsed)
109 }
110 Timing::Device => {
111 let profiled = client.profile(move || set.fastest(index).execute(input), "stack-autotune");
112 complete(&client)?;
113 let (output, profile) = profiled.map_err(|e| TuneFailure::rejected(std::format!("device profile failed: {e:?}")))?;
114 let output = output.map_err(|e| TuneFailure::rejected(std::format!("candidate rejected: {e:?}")))?;
115 let elapsed = ruda_core::future::block_on(profile.resolve()).duration();
116 std::hint::black_box(&output); drop(output);
117 Ok(elapsed)
118 }
119 }
120 }).map_err(|e| TuneFailure::device(std::format!("exclusive measurement failed: {e:?}")))?
121 }
122}
123
124pub fn try_execute_stack<'a, R: Runtime, K: AutotuneKey, I: TuneInputs, O: AutotuneOutput>(
127 name: &str, device_id: &str, client: &ComputeClient<R>, set: Arc<TunableSet<K,I,O>>, input: I::At<'a>,
128) -> Result<O, AutotuneError> {
129 let tuner = stack_autotuner().ok_or_else(|| autotune_error(name, "stack autotuning is not enabled"))?;
130 let reference = set.stack_reference().ok_or_else(|| autotune_error(name, "no explicit reference/workload signature was registered"))?;
131 if reference >= set.len() { return Err(autotune_error(name, "reference index out of range")); }
132 let key = set.generate_key(&input);
133 let env = runtime_environment(client);
134 let problem = Problem { scope: if name.contains("fusion") { Scope::Graph } else { Scope::Operator }, operation: fields(&[name, device_id, &set.stack_checksum()]), environment: env.fingerprint,
135 workload: set.stack_workload(&input).ok_or_else(|| autotune_error(name, "no exact workload signature"))?,
136 execution_context: env.execution_context, persistent: env.persistent };
137 let mut indices = std::vec![reference]; let mut plan = set.plan(&key);
140 loop { let batch = plan.next(None); if batch.is_empty() { break; }
141 for index in batch { if !indices.contains(&index) { indices.push(index); } }
142 }
143 let candidates: Vec<_> = indices.iter().map(|&i| Candidate::new(set.fastest(i).name.clone())).collect();
144 let mut trials = RuntimeTrials { client: client.clone(), set: &set, key: &key, input: &input, indices: &indices, timing: tuner.policy().timing };
147 let decision = tuner.select(&problem, &candidates, 0, &mut trials).map_err(|e| autotune_error(name, &e.to_string()))?;
148 let output = set.fastest(indices[decision.index]).execute(input);
149 if output.is_err() { tuner.invalidate(&decision, true); }
150 output
151}
152fn autotune_error(name: &str, message: &str) -> AutotuneError {
153 AutotuneError::Unknown { name: name.to_string(), err: message.to_string() }
154}