1use super::{cache::{DiskCache, Record, now_seconds}, policy::*};
5use std::{cell::Cell, collections::{BTreeMap, BTreeSet, VecDeque}, path::PathBuf,
6 string::{String, ToString}, sync::{Arc, Mutex, MutexGuard}, time::{Duration, Instant}, vec::Vec};
7
8std::thread_local! { static DEPTH: Cell<usize> = const { Cell::new(0) }; }
9pub 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
16pub trait TrialRunner {
19 fn validate(&mut self, reference: usize, candidate: usize, tolerance: Tolerance) -> Result<Validation, TuneFailure>;
23 fn measure(&mut self, candidate: usize) -> Result<Duration, TuneFailure>;
26}
27#[derive(Debug, Clone, Copy, PartialEq, Eq)]
29pub enum DecisionSource {
30 Tuned,
32 MemoryCache,
34 DiskCache,
36 CacheMiss,
38 Disabled,
40 Busy,
42 Nested,
44 ValidationUnavailable
46}
47#[derive(Debug, Clone)]
49pub struct Decision {
50 pub index: usize,
52 pub reference_index: usize,
54 pub name: String,
56 pub source: DecisionSource,
58 pub verified: bool,
60 pub ratio: Option<f64>,
62 pub cache_key: String,
64}
65#[derive(Debug, Clone)]
67pub struct CandidateReport {
68 pub name: String,
70 pub samples: usize,
72 pub ratio: Option<f64>,
74 pub relative_mad: Option<f64>,
76 pub verified: bool,
78 pub note: String,
80}
81#[derive(Debug, Clone)]
83pub struct TuneReport {
84 pub operation: String,
86 pub environment: String,
88 pub workload: String,
90 pub execution_context: String,
92 pub winner: String,
94 pub elapsed: Duration,
96 pub budget_exhausted: bool,
98 pub candidates: Vec<CandidateReport>,
100}
101#[derive(Debug, Clone, Default)]
103pub struct Stats {
104 pub memory_hits: u64,
106 pub disk_hits: u64,
108 pub misses: u64,
110 pub busy_fallbacks: u64,
112 pub tunes: u64,
114 pub cache_warnings: u64,
116 pub invalidations: u64,
118 pub device_failures: u64,
120}
121#[derive(Default)]
122struct State {
123 records: BTreeMap<String, Arc<Record>>, order: MemoryOrder,
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}
129enum MemoryOrder {
130 Fifo(VecDeque<String>),
131 Lru(LruOrder),
132}
133impl Default for MemoryOrder {
134 fn default() -> Self { Self::Fifo(VecDeque::new()) }
135}
136impl MemoryOrder {
137 fn new(eviction: MemoryEviction) -> Self {
138 match eviction { MemoryEviction::Fifo => Self::default(), MemoryEviction::Lru => Self::Lru(LruOrder::default()) }
139 }
140 fn len(&self) -> usize {
141 match self { Self::Fifo(order) => order.len(), Self::Lru(order) => order.indices.len() }
142 }
143 fn contains(&self, key: &String) -> bool {
144 match self { Self::Fifo(order) => order.contains(key), Self::Lru(order) => order.indices.contains_key(key.as_str()) }
145 }
146 fn push_back(&mut self, key: String) {
147 match self { Self::Fifo(order) => order.push_back(key), Self::Lru(order) => order.insert(key) }
148 }
149 fn pop_front(&mut self) -> Option<String> {
150 match self { Self::Fifo(order) => order.pop_front(), Self::Lru(order) => order.pop_front() }
151 }
152 fn promote(&mut self, key: &str) {
153 if let Self::Lru(order) = self { order.promote(key); }
154 }
155}
156#[derive(Default)]
157struct LruOrder {
158 indices: BTreeMap<Arc<str>, usize>,
159 slots: Vec<Option<LruEntry>>,
160 vacant: Vec<usize>,
161 head: Option<usize>,
162 tail: Option<usize>,
163}
164struct LruEntry {
165 key: Arc<str>,
166 previous: Option<usize>,
167 next: Option<usize>,
168}
169impl LruOrder {
170 fn detach(&mut self, index: usize) {
171 let entry = self.slots[index].as_ref().expect("retained LRU entry");
172 let (previous, next) = (entry.previous, entry.next);
173 match previous {
174 Some(previous) => self.slots[previous].as_mut().expect("previous LRU entry").next = next,
175 None => self.head = next,
176 }
177 match next {
178 Some(next) => self.slots[next].as_mut().expect("next LRU entry").previous = previous,
179 None => self.tail = previous,
180 }
181 }
182 fn append(&mut self, index: usize) {
183 let entry = self.slots[index].as_mut().expect("retained LRU entry");
184 entry.previous = self.tail;
185 entry.next = None;
186 match self.tail {
187 Some(tail) => self.slots[tail].as_mut().expect("last LRU entry").next = Some(index),
188 None => self.head = Some(index),
189 }
190 self.tail = Some(index);
191 }
192 fn promote(&mut self, key: &str) {
193 if let Some(&index) = self.indices.get(key) {
194 if self.tail != Some(index) { self.detach(index); self.append(index); }
195 }
196 }
197 fn insert(&mut self, key: String) {
198 if self.indices.contains_key(key.as_str()) { self.promote(&key); return; }
199 let key = Arc::<str>::from(key);
200 let index = match self.vacant.pop() {
201 Some(index) => index,
202 None => { self.slots.push(None); self.slots.len() - 1 },
203 };
204 self.slots[index] = Some(LruEntry { key: key.clone(), previous: None, next: None });
205 self.indices.insert(key, index);
206 self.append(index);
207 }
208 fn pop_front(&mut self) -> Option<String> {
209 let index = self.head?;
210 self.detach(index);
211 let entry = self.slots[index].take().expect("first LRU entry");
212 self.indices.remove(entry.key.as_ref());
213 self.vacant.push(index);
214 Some(entry.key.to_string())
215 }
216}
217pub struct StackTuner { policy: StackPolicy, memory_eviction: MemoryEviction, disk: Option<DiskCache>, state: Mutex<State> }
220struct Permit<'a> { tuner: &'a StackTuner, key: String, device: String }
221impl Drop for Permit<'_> {
222 fn drop(&mut self) { let mut s = self.tuner.lock(); s.pending.remove(&self.key); s.devices.remove(&self.device); }
223}
224impl StackTuner {
225 pub fn new(policy: StackPolicy, cache_directory: Option<PathBuf>) -> Result<Self, TuneFailure> {
229 Self::new_with_memory_eviction(policy, cache_directory, MemoryEviction::Fifo)
230 }
231 pub fn new_with_memory_eviction(policy: StackPolicy, cache_directory: Option<PathBuf>, memory_eviction: MemoryEviction) -> Result<Self, TuneFailure> {
233 policy.validate()?;
234 let disk = cache_directory.map(|dir| DiskCache::new(dir, policy.capacity));
235 Ok(Self { policy, memory_eviction, disk, state: Mutex::new(State { order: MemoryOrder::new(memory_eviction), ..State::default() }) })
236 }
237 fn lock(&self) -> MutexGuard<'_, State> { self.state.lock().unwrap_or_else(|e| e.into_inner()) }
238 pub fn policy(&self) -> &StackPolicy { &self.policy }
240 pub fn memory_eviction(&self) -> MemoryEviction { self.memory_eviction }
242 pub fn stats(&self) -> Stats { self.lock().stats.clone() }
244 pub fn lower_level_fingerprint(&self) -> String {
248 use std::fmt::Write;
249 let s = self.lock();
250 let mut parts = Vec::new();
251 for (key, record) in &s.records {
252 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])); }
253 }
254 for (key, scope) in &s.bypass {
256 if *scope != Scope::Pipeline { parts.push(fields(&[key, "unverified-reference"])); }
257 }
258 parts.sort();
259 let mut encoded = String::new();
260 for part in &parts { let _ = write!(encoded, "{}:{part}", part.len()); }
261 super::cache::digest(encoded.as_bytes())
262 }
263 pub fn reports(&self) -> Vec<TuneReport> { self.lock().reports.iter().cloned().collect() }
266 fn promote(&self, state: &mut State, key: &str) {
267 state.order.promote(key);
268 }
269 fn insert(&self, record: Record) {
270 let mut s = self.lock();
271 if !s.order.contains(&record.key) {
272 while s.order.len() >= self.policy.capacity {
273 if let Some(key) = s.order.pop_front() {
274 s.records.remove(&key); s.banned.remove(&key); s.regressions.remove(&key); s.bypass.remove(&key);
275 }
276 }
277 s.order.push_back(record.key.clone());
278 }
279 self.promote(&mut s, &record.key);
280 s.records.insert(record.key.clone(), Arc::new(record));
281 }
282 fn allowed(&self, key: &str, c: &impl CandidateSource) -> bool {
283 let c = c.view();
284 c.fits(&self.policy) && !self.lock().banned.get(key).is_some_and(|b| b.contains(c.name))
285 }
286 fn from_record<C: CandidateSource>(&self, r: &Record, candidates: &[C], reference: usize, source: DecisionSource) -> Option<Decision> {
287 if !r.is_fresh(now_seconds(), self.policy.ttl.as_secs().max(1)) || (self.policy.require_validation && !r.verified) { return None; }
288 let index = candidates.iter().position(|c| c.view().name == &r.winner && self.allowed(&r.key, c))?;
289 Some(Decision { index, reference_index: reference, name: r.winner.clone(), source, verified: r.verified, ratio: Some(r.ratio), cache_key: r.key.clone() })
290 }
291 fn validation(&self, runner: &mut impl TrialRunner, reference: usize, candidate: usize) -> Result<bool, TuneFailure> {
292 match runner.validate(reference, candidate, self.policy.tolerance)? {
293 Validation::Passed => Ok(true),
294 Validation::Unsupported if !self.policy.require_validation => Ok(false),
295 Validation::Unsupported => Err(TuneFailure { kind: FailureKind::Unavailable, message: "no numerical validator for this output/size".into() }),
296 }
297 }
298 fn fail_device(&self, environment: &str) {
299 let mut s = self.lock(); s.poisoned_devices.insert(environment.to_string()); s.stats.device_failures += 1;
300 }
301 pub fn select(&self, problem: &Problem, candidates: &[Candidate], reference: usize, runner: &mut impl TrialRunner) -> Result<Decision, TuneFailure> {
307 self.select_candidates(problem, candidates, reference, runner)
308 }
309
310 pub(super) fn select_candidates<C: CandidateSource>(&self, problem: &Problem, candidates: &[C], reference: usize, runner: &mut impl TrialRunner) -> Result<Decision, TuneFailure> {
311 if candidates.is_empty() || candidates.len() > 4096 || reference >= candidates.len()
312 || problem.operation.is_empty() || problem.environment.is_empty() || problem.workload.is_empty() {
313 return Err(TuneFailure::invalid("a workload, environment and valid reference candidate are required"));
314 }
315 let mut names = BTreeSet::new();
316 if candidates.iter().any(|c| {
317 let c = c.view();
318 c.name.is_empty() || c.name.len() > 4096 || !names.insert(c.name)
319 }) {
320 return Err(TuneFailure::invalid("candidate names must be nonempty, bounded and unique"));
321 }
322 if !candidates[reference].view().fits(&self.policy) { return Err(TuneFailure::invalid("reference does not satisfy eligibility/workspace policy")); }
323 let key = cache_key_candidates(problem, candidates, reference, &self.policy);
324 if key.len() > 384 * 1024 { return Err(TuneFailure::invalid("autotune key exceeds 384 KiB")); }
325 let fallback = |source| Decision { index: reference, reference_index: reference, name: String::clone(candidates[reference].view().name), source, verified: false, ratio: None, cache_key: key.clone() };
326 if self.lock().poisoned_devices.contains(&problem.environment) {
327 return Err(TuneFailure { kind: FailureKind::Quarantined, message: "device tuning lane is quarantined after an unconfirmed device completion".into() });
328 }
329 if self.policy.mode == Mode::Disabled { return Ok(fallback(DecisionSource::Disabled)); }
330 let memory = { self.lock().records.get(&key).cloned() };
331 if let Some(record) = memory {
332 if let Some(d) = self.from_record(&record, candidates, reference, DecisionSource::MemoryCache) {
333 let mut state = self.lock();
334 state.stats.memory_hits += 1;
335 self.promote(&mut state, &key);
336 return Ok(d);
337 }
338 }
339 let bypass = { self.lock().bypass.contains_key(&key) };
340 if bypass {
341 if self.memory_eviction == MemoryEviction::Lru { self.promote(&mut self.lock(), &key); }
342 return Ok(fallback(DecisionSource::ValidationUnavailable));
343 }
344 if is_tuning() { return Ok(fallback(DecisionSource::Nested)); }
345 let _permit = {
346 let mut s = self.lock();
347 if s.pending.contains(&key) || s.devices.contains(&problem.environment) || s.pending.len() >= self.policy.max_parallel_tunes {
348 s.stats.busy_fallbacks += 1; return Ok(fallback(DecisionSource::Busy));
349 }
350 s.pending.insert(key.clone()); s.devices.insert(problem.environment.clone()); s.stats.misses += 1;
351 Permit { tuner: self, key: key.clone(), device: problem.environment.clone() }
352 };
353 let _depth = DepthGuard::enter();
354 if problem.persistent {
355 if let Some(disk) = &self.disk {
356 match disk.load(&key) {
357 Ok(Some(record)) => {
358 if let Some(mut decision) = self.from_record(&record, candidates, reference, DecisionSource::DiskCache).filter(|_| record.scope == problem.scope as u8) {
359 match self.validation(runner, reference, decision.index) {
362 Ok(checked) => {
363 decision.verified = checked;
364 let mut record = record; record.verified = checked;
365 self.insert(record); self.lock().stats.disk_hits += 1; return Ok(decision);
366 }
367 Err(e) if e.kind == FailureKind::Device => { self.fail_device(&problem.environment); return Err(e); }
368 Err(_) => { if disk.remove(&key).is_err() { self.lock().stats.cache_warnings += 1; } }
369 }
370 }
371 }
372 Err(_) => self.lock().stats.cache_warnings += 1,
373 Ok(None) => {}
374 }
375 }
376 }
377 if self.policy.mode == Mode::CacheOnly { return Ok(fallback(DecisionSource::CacheMiss)); }
378 match self.explore(problem, candidates, reference, &key, runner) {
379 Ok((decision, report, record)) => {
380 if problem.persistent {
381 if let Some(disk) = &self.disk { if disk.save(&record).is_err() { self.lock().stats.cache_warnings += 1; } }
382 }
383 self.insert(record);
384 let mut s = self.lock(); s.stats.tunes += 1;
385 if s.reports.len() >= self.policy.capacity.min(128) { s.reports.pop_front(); }
386 s.reports.push_back(report);
387 Ok(decision)
388 }
389 Err(e) if e.kind == FailureKind::Unavailable => {
390 let mut s = self.lock();
393 if !s.order.contains(&key) { s.order.push_back(key.clone()); }
394 s.bypass.insert(key.clone(), problem.scope);
395 self.promote(&mut s, &key);
396 while s.order.len() > self.policy.capacity {
397 if let Some(old) = s.order.pop_front() { s.records.remove(&old); s.banned.remove(&old); s.regressions.remove(&old); s.bypass.remove(&old); }
398 }
399 if s.reports.len() >= self.policy.capacity.min(128) { s.reports.pop_front(); }
400 s.reports.push_back(TuneReport {
401 operation: problem.operation.clone(), environment: problem.environment.clone(),
402 workload: problem.workload.clone(), execution_context: problem.execution_context.clone(),
403 winner: String::clone(candidates[reference].view().name), elapsed: Duration::ZERO, budget_exhausted: false,
404 candidates: std::vec![CandidateReport { name: String::clone(candidates[reference].view().name), samples: 0,
405 ratio: None, relative_mad: None, verified: false, note: e.message }],
406 });
407 Ok(fallback(DecisionSource::ValidationUnavailable))
408 }
409 Err(e) => { if e.kind == FailureKind::Device { self.fail_device(&problem.environment); } Err(e) }
410 }
411 }
412 fn explore<C: CandidateSource>(&self, problem: &Problem, candidates: &[C], reference: usize, key: &str, runner: &mut impl TrialRunner) -> Result<(Decision, TuneReport, Record), TuneFailure> {
413 let started = Instant::now();
414 for _ in 0..self.policy.warmups { runner.measure(reference)?; }
417 let reference_checked = self.validation(runner, reference, reference)?;
418 let reference_time = runner.measure(reference)?;
419 if reference_time.is_zero() { return Err(TuneFailure { kind: FailureKind::Unavailable, message: "reference has zero elapsed time".into() }); }
420 let mut winner = reference; let mut winner_ratio = 1.0; let mut winner_time = reference_time;
421 let mut winner_reference = reference_time; let mut winner_checked = reference_checked;
422 let mut reports = Vec::new();
423 reports.push(CandidateReport { name: String::clone(candidates[reference].view().name), samples: 1, ratio: Some(1.0), relative_mad: Some(0.0), verified: reference_checked, note: "reference; retained unless a stable measured improvement qualifies".into() });
424 let mut attempted = 1; let mut exhausted = false;
425 for (index, candidate) in candidates.iter().enumerate() {
426 if index == reference { continue; }
427 let mut report = CandidateReport { name: String::clone(candidate.view().name), samples: 0, ratio: None, relative_mad: None, verified: false, note: String::new() };
428 if !self.allowed(key, candidate) { report.note = "ineligible, unknown/over-budget workspace, or regressed candidate".into(); reports.push(report); continue; }
429 if attempted >= self.policy.max_candidates || started.elapsed() >= self.policy.budget {
430 exhausted = true; report.note = "search budget exhausted".into(); reports.push(report); continue;
431 }
432 attempted += 1;
433 let result = (|| -> Result<(bool, Vec<(Duration, Duration)>), TuneFailure> {
434 for _ in 0..self.policy.warmups {
435 if started.elapsed() >= self.policy.budget { return Err(TuneFailure::rejected("budget exhausted during warmup")); }
436 runner.measure(index)?;
437 }
438 let checked = self.validation(runner, reference, index)?;
439 let mut pairs = Vec::with_capacity(self.policy.samples);
440 for sample in 0..self.policy.samples {
441 if started.elapsed() >= self.policy.budget { return Err(TuneFailure::rejected("budget exhausted before enough paired samples")); }
442 let pair = if (sample + index) % 2 == 0 {
444 let a = runner.measure(reference)?; let b = runner.measure(index)?; (a,b)
445 } else {
446 let b = runner.measure(index)?; let a = runner.measure(reference)?; (a,b)
447 };
448 if pair.0.is_zero() || pair.1.is_zero() { return Err(TuneFailure::rejected("zero duration cannot be scored")); }
449 pairs.push(pair);
450 }
451 Ok((checked, pairs))
452 })();
453 match result {
454 Err(e) if e.kind == FailureKind::Device => return Err(e),
455 Err(e) => { report.note = e.message; },
456 Ok((checked, pairs)) => {
457 report.verified = checked; report.samples = pairs.len();
458 match paired_score(&pairs, &self.policy) {
459 None => { report.note = "timings too noisy or incomplete".into(); },
460 Some((ratio, mad)) => {
461 report.ratio = Some(ratio); report.relative_mad = Some(mad);
462 report.note = "paired synchronized measurements".into();
463 if ratio < winner_ratio && ratio * self.policy.min_speedup <= 1.0 {
464 winner = index; winner_ratio = ratio; winner_checked = checked;
465 let mut a: Vec<_> = pairs.iter().map(|p| p.0.as_secs_f64()).collect();
466 let mut b: Vec<_> = pairs.iter().map(|p| p.1.as_secs_f64()).collect();
467 winner_reference = Duration::from_secs_f64(median(&mut a).unwrap());
468 winner_time = Duration::from_secs_f64(median(&mut b).unwrap());
469 }
470 }
471 }
472 }
473 }
474 reports.push(report);
475 }
476 exhausted |= started.elapsed() >= self.policy.budget;
477 let record = Record { scope: problem.scope as u8, key: key.to_string(), winner: String::clone(candidates[winner].view().name), created: now_seconds(),
478 reference_ns: winner_reference.as_nanos().min(u64::MAX as u128) as u64,
479 winner_ns: winner_time.as_nanos().min(u64::MAX as u128) as u64, ratio: winner_ratio, verified: winner_checked };
480 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() };
481 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 };
482 Ok((decision, report, record))
483 }
484 pub fn invalidate(&self, decision: &Decision, ban_candidate: bool) {
488 {
489 let mut s = self.lock(); s.records.remove(&decision.cache_key); s.regressions.remove(&decision.cache_key); s.bypass.remove(&decision.cache_key);
490 if ban_candidate && decision.index != decision.reference_index {
491 s.banned.entry(decision.cache_key.clone()).or_default().insert(decision.name.clone());
492 }
493 if !s.order.contains(&decision.cache_key) { s.order.push_back(decision.cache_key.clone()); }
495 while s.order.len() > self.policy.capacity {
496 if let Some(old) = s.order.pop_front() { s.records.remove(&old); s.banned.remove(&old); s.regressions.remove(&old); s.bypass.remove(&old); }
497 }
498 s.stats.invalidations += 1;
499 }
500 if let Some(disk) = &self.disk { if disk.remove(&decision.cache_key).is_err() { self.lock().stats.cache_warnings += 1; } }
501 }
502 pub fn record_comparison(&self, decision: &Decision, reference: Duration, selected: Duration, correctness_checked: bool) -> Result<bool, TuneFailure> {
508 if !correctness_checked || reference.is_zero() || selected.is_zero() { return Err(TuneFailure::invalid("regression observations require nonzero, correctness-checked paired timings")); }
509 let should_invalidate = {
510 let mut s = self.lock();
511 if !s.records.get(&decision.cache_key).is_some_and(|r| r.winner == decision.name) { return Ok(false); }
512 let history = s.regressions.entry(decision.cache_key.clone()).or_default();
513 if history.len() == self.policy.regression_pairs { history.pop_front(); }
514 history.push_back(selected.as_secs_f64()/reference.as_secs_f64());
515 if history.len() < self.policy.regression_pairs { false } else {
516 let mut ratios: Vec<_> = history.iter().copied().collect(); median(&mut ratios).is_some_and(|r| r > self.policy.regression_ratio)
517 }
518 };
519 if should_invalidate { self.invalidate(decision, true); }
520 Ok(should_invalidate)
521 }
522}