use std::collections::HashMap;
use cudarc::driver::{CudaStream, sys};
use kime_tensor::Result;
use crate::lt::{self, Ty};
use crate::{WORKSPACE, dev};
const TABLE: &str = include_str!("tuned.txt");
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) struct Key {
pub(crate) dims: (usize, usize, usize),
pub(crate) ab: Ty,
pub(crate) c: Ty,
pub(crate) acc: bool,
}
impl Key {
fn line(&self, gpu: &str, pick: usize) -> String {
let (m, k, n) = self.dims;
let ty = |t: Ty| if t == Ty::F16 { "f16" } else { "f32" };
let acc = if self.acc { "add" } else { "set" };
format!("{gpu} | {m} {k} {n} | {} {} {acc} | {pick}", ty(self.ab), ty(self.c))
}
}
#[derive(Debug, Default)]
pub(crate) struct Picks(HashMap<Key, usize>);
impl Picks {
pub(crate) fn for_gpu(gpu: &str) -> Self {
Self(TABLE.lines().filter_map(|l| parse(l, gpu)).collect())
}
pub(crate) fn get(&self, key: &Key) -> usize {
self.0.get(key).copied().unwrap_or(0)
}
}
fn parse(line: &str, gpu: &str) -> Option<(Key, usize)> {
let line = line.split('#').next().unwrap_or("").trim();
if line.is_empty() {
return None;
}
let f: Vec<&str> = line.split('|').map(str::trim).collect();
let [name, dims, types, pick] = f[..] else { return None };
if name != gpu {
return None;
}
let d: Vec<usize> =
dims.split_whitespace().map(str::parse).collect::<std::result::Result<_, _>>().ok()?;
let t: Vec<&str> = types.split_whitespace().collect();
let ty = |s: &str| match s {
"f16" => Some(Ty::F16),
"f32" => Some(Ty::F32),
_ => None,
};
let (&[m, k, n], &[ab, c, acc]) = (&d[..], &t[..]) else { return None };
let acc = match acc {
"add" => true,
"set" => false,
_ => return None,
};
let key = Key { dims: (m, k, n), ab: ty(ab)?, c: ty(c)?, acc };
Some((key, pick.parse().ok()?))
}
pub(crate) fn enabled() -> bool {
std::env::var_os("KIME_CUDA_TUNE").is_some_and(|v| v != "0")
}
pub(crate) fn tune(
h: <::Handle,
s: &CudaStream,
workspace: u64,
gemms: &mut [lt::Gemm],
keys: &[Key],
gpu: &str,
) -> Result<()> {
let ctx = s.context();
let timed = Some(sys::CUevent_flags::CU_EVENT_DEFAULT);
let marks =
(0..=gemms.len()).map(|_| ctx.new_event(timed)).collect::<std::result::Result<Vec<_>, _>>();
let marks = marks.map_err(dev)?;
let most = gemms.iter().map(lt::Gemm::candidates).max().unwrap_or(0);
let mut fastest = vec![vec![f32::INFINITY; most]; gemms.len()];
for pass in 0..6 {
for rank in 0..most {
marks[0].record(s).map_err(dev)?;
let mut ran = vec![false; gemms.len()];
for (i, g) in gemms.iter_mut().enumerate() {
if rank < g.candidates() {
let keep = g.pick;
g.pick = rank;
let run = unsafe { g.run(h, workspace, WORKSPACE, s.cu_stream().cast()) };
ran[i] = run.is_ok();
g.pick = keep;
}
marks[i + 1].record(s).map_err(dev)?;
}
s.synchronize().map_err(dev)?;
if pass == 0 {
continue;
}
for (i, f) in fastest.iter_mut().enumerate() {
if ran[i] {
f[rank] = f[rank].min(marks[i].elapsed_ms(&marks[i + 1]).map_err(dev)?);
}
}
}
}
let mut total: HashMap<Key, Vec<f32>> = HashMap::new();
for (key, f) in keys.iter().zip(&fastest) {
let t = total.entry(*key).or_insert_with(|| vec![0.0; most]);
for (t, f) in t.iter_mut().zip(f) {
*t += f;
}
}
let mut best: HashMap<Key, usize> = HashMap::new();
for (key, t) in &total {
let (mut pick, mut at) = (0, t[0]);
for (rank, &x) in t.iter().enumerate().skip(1) {
if x < at && x < 0.98 * t[0] {
(pick, at) = (rank, x);
}
}
best.insert(*key, pick);
let calls = keys.iter().filter(|k| *k == key).count() as f32;
let each: Vec<String> = t.iter().map(|x| format!("{:.2}", 1e3 * x / calls)).collect();
println!(
"{} # {} rows, us by rank: {}",
key.line(gpu, pick),
gemms[keys.iter().position(|k| k == key).unwrap_or(0)].dims.0,
each.join(" ")
);
}
for (g, key) in gemms.iter_mut().zip(keys) {
g.pick = best[key];
}
Ok(())
}