use std::collections::{HashMap, HashSet};
use std::path::{Path, PathBuf};
use std::sync::{Mutex, OnceLock};
use crate::prelude::{ComputeClient, Runtime};
pub struct Variant<'a, T> {
pub name: &'static str,
f: Box<dyn Fn(usize) -> (T, f64) + 'a>,
}
impl<'a, T> Variant<'a, T> {
pub fn new(name: &'static str, f: impl Fn(usize) -> (T, f64) + 'a) -> Self {
Variant {
name,
f: Box::new(f),
}
}
}
#[derive(Debug)]
pub struct Pick<T> {
pub output: T,
pub winner: String,
pub from_cache: bool,
pub benched: usize,
pub timings: Vec<(String, f64)>,
}
const TUNE_ITERS: usize = 25;
#[derive(Default)]
struct Inner {
winners: HashMap<(String, String, String), String>,
loaded: HashSet<String>,
}
pub struct Tuner {
cache_root: PathBuf,
inner: Mutex<Inner>,
}
impl Tuner {
pub fn new(cache_root: impl Into<PathBuf>) -> Self {
Tuner {
cache_root: cache_root.into(),
inner: Mutex::new(Inner::default()),
}
}
fn device_file(&self, device: &str) -> PathBuf {
self.cache_root
.join("hanzo-kernel")
.join("autotune")
.join(format!("{}.tsv", sanitize(device)))
}
pub fn select<T>(
&self,
device: &str,
op: &str,
key: &str,
variants: Vec<Variant<'_, T>>,
) -> Pick<T> {
assert!(
!variants.is_empty(),
"tune::select needs at least one variant"
);
self.ensure_loaded(device);
let ck = (device.to_string(), op.to_string(), key.to_string());
let cached = self
.inner
.lock()
.expect("tuner cache poisoned")
.winners
.get(&ck)
.cloned();
if let Some(winner) = cached {
if let Some(v) = variants.iter().find(|v| v.name == winner) {
let (output, _ms) = (v.f)(1);
return Pick {
output,
winner,
from_cache: true,
benched: 0,
timings: Vec::new(),
};
}
}
let mut timings: Vec<(String, f64)> = Vec::with_capacity(variants.len());
let mut best: Option<(usize, f64)> = None;
let mut outputs: Vec<Option<T>> = Vec::with_capacity(variants.len());
for (i, v) in variants.iter().enumerate() {
let (out, ms) = (v.f)(TUNE_ITERS);
timings.push((v.name.to_string(), ms));
outputs.push(Some(out));
if best.map(|(_, bms)| ms < bms).unwrap_or(true) {
best = Some((i, ms));
}
}
let (bi, _) = best.expect("at least one variant timed");
let winner = variants[bi].name.to_string();
let output = outputs[bi].take().expect("winner output present");
self.record(device, op, key, &winner);
timings.sort_by(|a, b| a.1.total_cmp(&b.1));
Pick {
output,
winner,
from_cache: false,
benched: variants.len(),
timings,
}
}
fn ensure_loaded(&self, device: &str) {
{
let inner = self.inner.lock().expect("tuner cache poisoned");
if inner.loaded.contains(device) {
return;
}
}
let entries = read_device_file(&self.device_file(device));
let mut inner = self.inner.lock().expect("tuner cache poisoned");
for (op, key, winner) in entries {
inner.winners.insert((device.to_string(), op, key), winner);
}
inner.loaded.insert(device.to_string());
}
fn record(&self, device: &str, op: &str, key: &str, winner: &str) {
{
let mut inner = self.inner.lock().expect("tuner cache poisoned");
inner.winners.insert(
(device.to_string(), op.to_string(), key.to_string()),
winner.to_string(),
);
}
append_device_file(&self.device_file(device), op, key, winner);
}
pub fn cached_winner(&self, device: &str, op: &str, key: &str) -> Option<String> {
self.ensure_loaded(device);
self.inner
.lock()
.expect("tuner cache poisoned")
.winners
.get(&(device.to_string(), op.to_string(), key.to_string()))
.cloned()
}
}
pub fn global() -> &'static Tuner {
static G: OnceLock<Tuner> = OnceLock::new();
G.get_or_init(|| Tuner::new(xdg_cache_dir()))
}
pub fn xdg_cache_dir() -> PathBuf {
if let Some(x) = std::env::var_os("XDG_CACHE_HOME") {
if !x.is_empty() {
return PathBuf::from(x);
}
}
if let Some(h) = std::env::var_os("HOME") {
if !h.is_empty() {
return PathBuf::from(h).join(".cache");
}
}
PathBuf::from(".")
}
pub fn device_id<R: Runtime>(_client: &ComputeClient<R>) -> String {
let name = std::any::type_name::<R>();
name.rsplit("::").next().unwrap_or(name).to_string()
}
fn sanitize(s: &str) -> String {
s.chars()
.map(|c| {
if c.is_ascii_alphanumeric() || c == '-' || c == '_' {
c
} else {
'_'
}
})
.collect()
}
fn cell(s: &str) -> String {
s.chars()
.map(|c| {
if c == '\t' || c == '\n' || c == '\r' {
' '
} else {
c
}
})
.collect()
}
fn read_device_file(path: &Path) -> Vec<(String, String, String)> {
let Ok(text) = std::fs::read_to_string(path) else {
return Vec::new();
};
let mut map: HashMap<(String, String), String> = HashMap::new();
for line in text.lines() {
let mut it = line.splitn(3, '\t');
if let (Some(op), Some(key), Some(winner)) = (it.next(), it.next(), it.next()) {
map.insert((op.to_string(), key.to_string()), winner.to_string());
}
}
map.into_iter()
.map(|((op, key), winner)| (op, key, winner))
.collect()
}
fn append_device_file(path: &Path, op: &str, key: &str, winner: &str) {
use std::io::Write;
if let Some(dir) = path.parent() {
let _ = std::fs::create_dir_all(dir);
}
if let Ok(mut f) = std::fs::OpenOptions::new()
.create(true)
.append(true)
.open(path)
{
let _ = writeln!(f, "{}\t{}\t{}", cell(op), cell(key), cell(winner));
}
}
pub struct Tuned<'a, T> {
op: &'static str,
key: String,
variants: Vec<Variant<'a, T>>,
}
impl<'a, T> Tuned<'a, T> {
pub fn new(op: &'static str, key: impl Into<String>) -> Self {
Tuned {
op,
key: key.into(),
variants: Vec::new(),
}
}
pub fn variant(mut self, name: &'static str, f: impl Fn(usize) -> (T, f64) + 'a) -> Self {
self.variants.push(Variant::new(name, f));
self
}
pub fn pick_with(self, tuner: &Tuner, device: &str) -> Pick<T> {
tuner.select(device, self.op, &self.key, self.variants)
}
pub fn pick<R: Runtime>(self, client: &ComputeClient<R>) -> Pick<T> {
let device = device_id::<R>(client);
self.pick_with(global(), &device)
}
pub fn run<R: Runtime>(self, client: &ComputeClient<R>) -> T {
self.pick(client).output
}
}
use std::hash::Hash;
pub struct Rng(u64);
impl Rng {
pub fn new(seed: u64) -> Self {
Rng(seed ^ 0x9E3779B97F4A7C15 | 1)
}
fn next_u64(&mut self) -> u64 {
let mut x = self.0;
x ^= x >> 12;
x ^= x << 25;
x ^= x >> 27;
self.0 = x;
x.wrapping_mul(0x2545F4914F6CDD1D)
}
pub fn below(&mut self, n: usize) -> usize {
(self.next_u64() % n as u64) as usize
}
pub fn chance(&mut self, p: f64) -> bool {
((self.next_u64() >> 11) as f64) / ((1u64 << 53) as f64) < p
}
}
pub struct Param {
pub name: &'static str,
pub values: Vec<i64>,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub struct Config {
vals: Vec<i64>,
}
impl Config {
pub fn get(&self, space: &Space, name: &str) -> i64 {
let i = space
.index(name)
.unwrap_or_else(|| panic!("no param `{name}` in space"));
self.vals[i]
}
pub fn name(&self, space: &Space) -> String {
space
.params
.iter()
.zip(&self.vals)
.map(|(p, v)| format!("{}={}", p.name, v))
.collect::<Vec<_>>()
.join(",")
}
}
type Constraint = Box<dyn Fn(&Config, &Space) -> bool>;
pub struct Space {
params: Vec<Param>,
constraints: Vec<Constraint>,
denies: Vec<Vec<(&'static str, i64)>>,
}
impl Default for Space {
fn default() -> Self {
Space {
params: Vec::new(),
constraints: Vec::new(),
denies: Vec::new(),
}
}
}
impl Space {
pub fn new() -> Self {
Space::default()
}
pub fn param(mut self, name: &'static str, values: impl IntoIterator<Item = i64>) -> Self {
self.params.push(Param {
name,
values: values.into_iter().collect(),
});
self
}
pub fn constraint(mut self, f: impl Fn(&Config, &Space) -> bool + 'static) -> Self {
self.constraints.push(Box::new(f));
self
}
pub fn deny(mut self, subspace: impl IntoIterator<Item = (&'static str, i64)>) -> Self {
self.denies.push(subspace.into_iter().collect());
self
}
pub fn arity(&self) -> usize {
self.params.len()
}
fn index(&self, name: &str) -> Option<usize> {
self.params.iter().position(|p| p.name == name)
}
pub fn config(&self, assignments: &[(&str, i64)]) -> Config {
let mut vals = vec![i64::MIN; self.params.len()];
let mut set = vec![false; self.params.len()];
for (name, v) in assignments {
let i = self
.index(name)
.unwrap_or_else(|| panic!("no param `{name}`"));
vals[i] = *v;
set[i] = true;
}
assert!(
set.iter().all(|&b| b),
"config() must assign every parameter"
);
Config { vals }
}
pub fn parse(&self, name: &str) -> Option<Config> {
let mut vals = vec![i64::MIN; self.params.len()];
let mut set = vec![false; self.params.len()];
for tok in name.split(',') {
let (k, v) = tok.split_once('=')?;
let i = self.index(k)?;
let v: i64 = v.parse().ok()?;
if !self.params[i].values.contains(&v) {
return None;
}
vals[i] = v;
set[i] = true;
}
if !set.iter().all(|&b| b) {
return None;
}
Some(Config { vals })
}
pub fn denied(&self, c: &Config) -> bool {
self.denies.iter().any(|entry| {
entry
.iter()
.all(|(name, val)| self.index(name).is_some_and(|i| c.vals[i] == *val))
})
}
pub fn feasible(&self, c: &Config) -> bool {
if c.vals.len() != self.params.len() {
return false;
}
if !self
.params
.iter()
.zip(&c.vals)
.all(|(p, v)| p.values.contains(v))
{
return false;
}
if self.denied(c) {
return false;
}
self.constraints.iter().all(|f| f(c, self))
}
pub fn enumerate(&self) -> Vec<Config> {
let mut out = vec![Config {
vals: Vec::with_capacity(self.params.len()),
}];
for p in &self.params {
let mut next = Vec::with_capacity(out.len() * p.values.len());
for base in &out {
for &v in &p.values {
let mut vals = base.vals.clone();
vals.push(v);
next.push(Config { vals });
}
}
out = next;
}
out.retain(|c| self.feasible(c));
out
}
pub fn random(&self, rng: &mut Rng, tries: usize) -> Option<Config> {
for _ in 0..tries {
let vals = self
.params
.iter()
.map(|p| p.values[rng.below(p.values.len())])
.collect();
let c = Config { vals };
if self.feasible(&c) {
return Some(c);
}
}
None
}
}
#[derive(Clone, Debug)]
pub enum Verdict {
Pass,
Reject(String),
}
pub trait Evaluator {
fn static_check(&self, cfg: &Config) -> Verdict;
fn measure(&self, cfg: &Config, iters: usize) -> f64;
}
#[derive(Clone)]
struct Cell {
fitness: f64,
ms: Option<f64>,
reason: Option<String>,
}
#[derive(Debug)]
pub struct EvoReport {
pub best: Config,
pub best_name: String,
pub best_ms: f64,
pub measured: Vec<(String, f64)>,
pub rejected: Vec<(String, String)>,
pub evaluated: usize,
pub generations: usize,
}
pub struct Evolution {
pop: usize,
generations: usize,
tournament: usize,
mutation: f64,
elitism: usize,
measure_iters: usize,
seeds: Vec<Config>,
}
impl Default for Evolution {
fn default() -> Self {
Evolution {
pop: 16,
generations: 8,
tournament: 3,
mutation: 0.25,
elitism: 2,
measure_iters: TUNE_ITERS,
seeds: Vec::new(),
}
}
}
impl Evolution {
pub fn new() -> Self {
Evolution::default()
}
pub fn population(mut self, n: usize) -> Self {
self.pop = n.max(1);
self
}
pub fn generations(mut self, n: usize) -> Self {
self.generations = n;
self
}
pub fn tournament(mut self, k: usize) -> Self {
self.tournament = k.max(1);
self
}
pub fn mutation(mut self, p: f64) -> Self {
self.mutation = p.clamp(0.0, 1.0);
self
}
pub fn elitism(mut self, k: usize) -> Self {
self.elitism = k;
self
}
pub fn measure_iters(mut self, n: usize) -> Self {
self.measure_iters = n;
self
}
pub fn seed_config(mut self, c: Config) -> Self {
self.seeds.push(c);
self
}
fn cmp_fit(
space: &Space,
memo: &HashMap<Config, Cell>,
a: &Config,
b: &Config,
) -> std::cmp::Ordering {
let fa = memo[a].fitness;
let fb = memo[b].fitness;
fa.total_cmp(&fb)
.then_with(|| a.name(space).cmp(&b.name(space)))
}
fn evaluate<E: Evaluator>(
&self,
eval: &E,
c: &Config,
memo: &mut HashMap<Config, Cell>,
order: &mut Vec<Config>,
) -> f64 {
if let Some(cell) = memo.get(c) {
return cell.fitness;
}
let cell = match eval.static_check(c) {
Verdict::Reject(reason) => Cell {
fitness: f64::INFINITY,
ms: None,
reason: Some(reason),
},
Verdict::Pass => {
let ms = eval.measure(c, self.measure_iters);
Cell {
fitness: ms,
ms: Some(ms),
reason: None,
}
}
};
let f = cell.fitness;
memo.insert(c.clone(), cell);
order.push(c.clone());
f
}
fn select<'p>(
&self,
space: &Space,
memo: &HashMap<Config, Cell>,
pop: &'p [Config],
rng: &mut Rng,
) -> &'p Config {
let mut best = &pop[rng.below(pop.len())];
for _ in 1..self.tournament {
let c = &pop[rng.below(pop.len())];
if Self::cmp_fit(space, memo, c, best).is_lt() {
best = c;
}
}
best
}
fn crossover(&self, a: &Config, b: &Config, rng: &mut Rng) -> Config {
let vals = a
.vals
.iter()
.zip(&b.vals)
.map(|(&va, &vb)| if rng.chance(0.5) { va } else { vb })
.collect();
Config { vals }
}
fn mutate(&self, space: &Space, c: &mut Config, rng: &mut Rng) {
for (i, p) in space.params.iter().enumerate() {
if p.values.len() > 1 && rng.chance(self.mutation) {
let cur = c.vals[i];
let mut j = rng.below(p.values.len() - 1);
if p.values[j] == cur {
j = p.values.len() - 1;
}
c.vals[i] = p.values[j];
}
}
}
fn repair(&self, space: &Space, mut c: Config, rng: &mut Rng) -> Option<Config> {
for _ in 0..64 {
if space.feasible(&c) {
return Some(c);
}
self.mutate(space, &mut c, rng);
}
space.random(rng, 1024)
}
pub fn hunt<E: Evaluator>(&self, space: &Space, eval: &E, seed: u64) -> EvoReport {
let mut rng = Rng::new(seed);
let mut memo: HashMap<Config, Cell> = HashMap::new();
let mut order: Vec<Config> = Vec::new();
let mut population: Vec<Config> = Vec::new();
let mut seen: HashSet<Config> = HashSet::new();
for s in &self.seeds {
if space.feasible(s) && seen.insert(s.clone()) {
population.push(s.clone());
}
}
let mut fill_tries = 0usize;
while population.len() < self.pop && fill_tries < self.pop * 128 {
fill_tries += 1;
match space.random(&mut rng, 256) {
Some(c) if seen.insert(c.clone()) => population.push(c),
Some(_) => {} None => break, }
}
if population.is_empty() {
return EvoReport {
best: Config { vals: vec![] },
best_name: "<none>".to_string(),
best_ms: f64::INFINITY,
measured: Vec::new(),
rejected: Vec::new(),
evaluated: 0,
generations: 0,
};
}
for c in &population.clone() {
self.evaluate(eval, c, &mut memo, &mut order);
}
let mut gens_run = 0usize;
for _ in 0..self.generations {
gens_run += 1;
let mut ranked = population.clone();
ranked.sort_by(|a, b| Self::cmp_fit(space, &memo, a, b));
let mut next: Vec<Config> = Vec::new();
let mut nseen: HashSet<Config> = HashSet::new();
for e in ranked.iter().take(self.elitism) {
if nseen.insert(e.clone()) {
next.push(e.clone());
}
}
let mut tries = 0usize;
while next.len() < self.pop && tries < self.pop * 128 {
tries += 1;
let a = self.select(space, &memo, &population, &mut rng).clone();
let b = self.select(space, &memo, &population, &mut rng).clone();
let mut child = self.crossover(&a, &b, &mut rng);
self.mutate(space, &mut child, &mut rng);
let child = match self.repair(space, child, &mut rng) {
Some(c) => c,
None => continue,
};
if nseen.insert(child.clone()) {
next.push(child);
}
}
if next.is_empty() {
break;
}
for c in &next.clone() {
self.evaluate(eval, c, &mut memo, &mut order);
}
population = next;
}
let mut measured: Vec<(String, f64)> = order
.iter()
.filter_map(|c| memo[c].ms.map(|ms| (c.name(space), ms)))
.collect();
measured.sort_by(|x, y| x.1.total_cmp(&y.1).then_with(|| x.0.cmp(&y.0)));
let rejected: Vec<(String, String)> = order
.iter()
.filter_map(|c| memo[c].reason.clone().map(|r| (c.name(space), r)))
.collect();
let best = order
.iter()
.filter(|c| memo[*c].ms.is_some())
.min_by(|a, b| Self::cmp_fit(space, &memo, a, b))
.cloned()
.unwrap_or_else(|| order[0].clone());
let best_ms = memo[&best].ms.unwrap_or(f64::INFINITY);
let best_name = best.name(space);
EvoReport {
best,
best_name,
best_ms,
measured,
rejected,
evaluated: order.len(),
generations: gens_run,
}
}
}
pub struct Evolved {
pub winner: String,
pub from_cache: bool,
pub report: Option<EvoReport>,
}
impl Tuner {
pub fn evolve<E: Evaluator>(
&self,
device: &str,
op: &str,
key: &str,
space: &Space,
eval: &E,
evo: &Evolution,
seed: u64,
) -> Evolved {
if let Some(w) = self.cached_winner(device, op, key) {
if space.parse(&w).is_some() {
return Evolved {
winner: w,
from_cache: true,
report: None,
};
}
}
let report = evo.hunt(space, eval, seed);
self.record(device, op, key, &report.best_name);
Evolved {
winner: report.best_name.clone(),
from_cache: false,
report: Some(report),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn select_caches_and_reloads_from_disk() {
let dir = tmp_dir("select");
let variants = || {
vec![
Variant::new("slow", |_it| (10u32, 9.0)),
Variant::new("fast", |_it| (20u32, 1.0)), Variant::new("mid", |_it| (30u32, 5.0)),
]
};
let t1 = Tuner::new(&dir);
let p = t1.select("cpu", "op", "s1", variants());
assert!(!p.from_cache);
assert_eq!(p.benched, 3);
assert_eq!(p.winner, "fast");
assert_eq!(p.output, 20);
assert_eq!(p.timings.first().unwrap().0, "fast");
let p = t1.select("cpu", "op", "s1", variants());
assert!(p.from_cache);
assert_eq!(p.benched, 0);
assert_eq!(p.winner, "fast");
assert_eq!(p.output, 20);
let t2 = Tuner::new(&dir);
assert_eq!(t2.cached_winner("cpu", "op", "s1").as_deref(), Some("fast"));
let p = t2.select("cpu", "op", "s1", variants());
assert!(p.from_cache);
assert_eq!(p.benched, 0);
assert_eq!(p.winner, "fast");
let p = t2.select("cpu", "op", "s2", variants());
assert!(!p.from_cache);
assert_eq!(p.benched, 3);
let shrunk = vec![
Variant::new("slow", |_it| (10u32, 9.0)),
Variant::new("mid", |_it| (30u32, 2.0)),
];
let p = t2.select("cpu", "op", "s1", shrunk);
assert!(!p.from_cache);
assert_eq!(p.winner, "mid");
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn disk_format_round_trips() {
let dir = tmp_dir("disk");
let t = Tuner::new(&dir);
let file = t.device_file("CpuRuntime");
append_device_file(&file, "rms_norm", "rows=8,n=4", "b64_r1");
append_device_file(&file, "matvec", "rows=8,k=64", "b128_v2");
append_device_file(&file, "rms_norm", "rows=8,n=4", "b256_r1");
let mut got = read_device_file(&file);
got.sort();
assert_eq!(
got,
vec![
(
"matvec".to_string(),
"rows=8,k=64".to_string(),
"b128_v2".to_string()
),
(
"rms_norm".to_string(),
"rows=8,n=4".to_string(),
"b256_r1".to_string()
),
]
);
std::fs::remove_dir_all(&dir).ok();
}
#[test]
fn device_file_path_is_xdg_shaped() {
let t = Tuner::new("/x/y");
assert_eq!(
t.device_file("Cuda::Runtime").to_str().unwrap(),
"/x/y/hanzo-kernel/autotune/Cuda__Runtime.tsv"
);
}
struct Deceptive<'a> {
space: &'a Space,
}
impl<'a> Evaluator for Deceptive<'a> {
fn static_check(&self, _c: &Config) -> Verdict {
Verdict::Pass
}
fn measure(&self, c: &Config, _iters: usize) -> f64 {
let a = c.get(self.space, "a") as f64;
let b = c.get(self.space, "b") as f64;
let cc = c.get(self.space, "c") as f64;
let g = ((a - 7.0).powi(2) + (b - 2.0).powi(2) + (cc - 5.0).powi(2)) / 3.5;
let l = ((a - 2.0).powi(2) + (b - 7.0).powi(2) + (cc - 4.0).powi(2)) / 18.0;
10.0 - 10.0 * (-g).exp() - 6.0 * (-l).exp()
}
}
fn grid3() -> Space {
Space::new()
.param("a", 0..=9)
.param("b", 0..=9)
.param("c", 0..=9)
}
fn brute_argmin(
space: &Space,
eval: &Deceptive,
keep: impl Fn(&Config) -> bool,
) -> (Config, f64) {
let mut best: Option<(Config, f64)> = None;
for a in 0..=9 {
for b in 0..=9 {
for c in 0..=9 {
let cfg = space.config(&[("a", a), ("b", b), ("c", c)]);
if !keep(&cfg) {
continue;
}
let ms = eval.measure(&cfg, 1);
let better = match &best {
None => true,
Some((bc, bms)) => {
ms < *bms || (ms == *bms && cfg.name(space) < bc.name(space))
}
};
if better {
best = Some((cfg, ms));
}
}
}
}
best.unwrap()
}
#[test]
fn evo_converges_on_deceptive_landscape_and_is_deterministic() {
let space = grid3();
let eval = Deceptive { space: &space };
let (opt, opt_ms) = brute_argmin(&space, &eval, |_| true);
assert_eq!(opt.name(&space), "a=7,b=2,c=5"); let local_ms = eval.measure(&space.config(&[("a", 2), ("b", 7), ("c", 4)]), 1);
let evo = Evolution::new()
.population(28)
.generations(22)
.tournament(3)
.mutation(0.35)
.elitism(2);
let r = evo.hunt(&space, &eval, 0xC0FFEE);
assert_eq!(
r.best_name, "a=7,b=2,c=5",
"GA trapped in the local well; best={}",
r.best_name
);
assert!((r.best_ms - opt_ms).abs() < 1e-9);
assert!(r.best_ms < local_ms - 1.0, "did not escape the local trap");
assert!(
r.evaluated < 700,
"evaluated {} of 1000 -- not searching",
r.evaluated
);
assert_eq!(r.rejected.len(), 0);
let r2 = evo.hunt(&space, &eval, 0xC0FFEE);
assert_eq!(r.best_name, r2.best_name);
assert_eq!(r.evaluated, r2.evaluated);
assert_eq!(r.measured, r2.measured);
let seeds: [u64; 10] = [1, 42, 7777, 0xABCDEF, 0xC0FFEE, 2, 3, 99, 12345, 0xDEADBEEF];
let hits = seeds
.iter()
.filter(|&&s| evo.hunt(&space, &eval, s).best_name == "a=7,b=2,c=5")
.count();
assert!(
hits >= 8,
"only {hits}/10 seeds reached the global optimum -- escape not systematic"
);
}
#[test]
fn evo_respects_the_deny_list() {
let space = grid3().deny([("a", 7)]);
let eval = Deceptive { space: &space };
assert!(!space.feasible(&space.config(&[("a", 7), ("b", 2), ("c", 5)])));
assert!(space.denied(&space.config(&[("a", 7), ("b", 0), ("c", 0)])));
assert!(space.feasible(&space.config(&[("a", 6), ("b", 2), ("c", 5)])));
let mut rng = Rng::new(5);
for _ in 0..1000 {
let c = space.random(&mut rng, 256).unwrap();
assert_ne!(c.get(&space, "a"), 7);
}
let (_opt, opt_ms) = brute_argmin(&space, &eval, |c| c.get(&space, "a") != 7);
let local_ms = eval.measure(&space.config(&[("a", 2), ("b", 7), ("c", 4)]), 1);
let evo = Evolution::new()
.population(28)
.generations(22)
.tournament(3)
.mutation(0.35)
.elitism(2);
let r = evo.hunt(&space, &eval, 99);
assert!(
space.feasible(&r.best) && r.best.get(&space, "a") != 7,
"best is denied: {}",
r.best_name
);
assert!(
r.best_ms < local_ms - 0.1,
"did not escape the deceptive well: best_ms={}",
r.best_ms
);
assert!(
r.best_ms <= opt_ms + 2.0,
"best far from the admissible optimum: {}",
r.best_ms
);
for (name, _) in r.measured.iter() {
let c = space.parse(name).unwrap();
assert_ne!(c.get(&space, "a"), 7, "measured a denied config: {name}");
}
}
struct Gated<'a> {
space: &'a Space,
measured: std::cell::RefCell<usize>,
}
impl<'a> Evaluator for Gated<'a> {
fn static_check(&self, c: &Config) -> Verdict {
let sum = c.get(self.space, "a") + c.get(self.space, "b") + c.get(self.space, "c");
if sum > 20 {
Verdict::Reject(format!("synthetic spill: sum={sum} > 20"))
} else {
Verdict::Pass
}
}
fn measure(&self, c: &Config, _iters: usize) -> f64 {
let sum = c.get(self.space, "a") + c.get(self.space, "b") + c.get(self.space, "c");
assert!(
sum <= 20,
"measure() ran on a statically-rejected config: sum={sum}"
);
*self.measured.borrow_mut() += 1;
-(sum as f64)
}
}
#[test]
fn evo_multi_fidelity_gates_the_gpu_tier() {
let space = grid3();
let eval = Gated {
space: &space,
measured: std::cell::RefCell::new(0),
};
let evo = Evolution::new()
.population(24)
.generations(15)
.mutation(0.35);
let r = evo.hunt(&space, &eval, 2024);
assert!(!r.rejected.is_empty(), "nothing was statically rejected");
assert!(r.rejected.iter().all(|(_, why)| why.contains("spill")));
assert!(r.best_ms.is_finite());
assert_eq!(
r.best_ms, -20.0,
"winner should sit on the sum==20 boundary"
);
assert_eq!(*eval.measured.borrow(), r.measured.len());
assert_eq!(r.measured.len() + r.rejected.len(), r.evaluated);
}
#[test]
fn space_name_parse_round_trips() {
let space = Space::new()
.param("NWARP", [2, 4, 8])
.param("RM", [1, 2, 4]);
let c = space.config(&[("NWARP", 4), ("RM", 2)]);
assert_eq!(c.name(&space), "NWARP=4,RM=2");
assert_eq!(space.parse("NWARP=4,RM=2").as_ref(), Some(&c));
assert!(space.parse("NWARP=4").is_none()); assert!(space.parse("NWARP=3,RM=2").is_none()); assert!(space.parse("BOGUS=1,RM=2").is_none()); }
#[test]
fn space_enumerate_lists_only_feasible() {
let space = Space::new()
.param("a", [1, 2, 3])
.param("b", [1, 2, 3])
.constraint(|c, s| c.get(s, "a") + c.get(s, "b") <= 4)
.deny([("a", 2)]);
let all = space.enumerate();
assert!(all.iter().all(|c| space.feasible(c)));
assert!(all.iter().all(|c| c.get(&space, "a") != 2)); assert_eq!(all.len(), 4);
}
struct Bowl<'a> {
space: &'a Space,
calls: std::cell::RefCell<usize>,
}
impl<'a> Evaluator for Bowl<'a> {
fn static_check(&self, _c: &Config) -> Verdict {
Verdict::Pass
}
fn measure(&self, c: &Config, _iters: usize) -> f64 {
*self.calls.borrow_mut() += 1;
(c.get(self.space, "a") + c.get(self.space, "b")) as f64
}
}
#[test]
fn tuner_evolve_caches_and_skips_the_hunt() {
let dir = tmp_dir("evolve");
let space = Space::new().param("a", 0..=4).param("b", 0..=4);
let eval = Bowl {
space: &space,
calls: std::cell::RefCell::new(0),
};
let evo = Evolution::new().population(8).generations(5);
let t1 = Tuner::new(&dir);
let r1 = t1.evolve("cpu", "coopmat", "k1", &space, &eval, &evo, 7);
assert!(!r1.from_cache);
assert!(r1.report.is_some());
assert_eq!(r1.winner, "a=0,b=0"); let after_first = *eval.calls.borrow();
assert!(after_first > 0);
let r2 = t1.evolve("cpu", "coopmat", "k1", &space, &eval, &evo, 7);
assert!(r2.from_cache);
assert!(r2.report.is_none());
assert_eq!(r2.winner, r1.winner);
assert_eq!(
*eval.calls.borrow(),
after_first,
"cache hit still ran the hunt"
);
let t2 = Tuner::new(&dir);
let r3 = t2.evolve("cpu", "coopmat", "k1", &space, &eval, &evo, 7);
assert!(r3.from_cache);
assert_eq!(r3.winner, r1.winner);
assert_eq!(*eval.calls.borrow(), after_first);
std::fs::remove_dir_all(&dir).ok();
}
fn tmp_dir(tag: &str) -> PathBuf {
let p = std::env::temp_dir().join(format!(
"hanzo-kernel-tune-{tag}-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&p).unwrap();
p
}
}