use super::MatMatMul;
fn order_f(a: f32, b: f32) -> std::cmp::Ordering {
if a < b { std::cmp::Ordering::Less } else { std::cmp::Ordering::Greater }
}
#[derive(Debug)]
pub struct LinearCostModel<'a> {
pub default_kernel: &'a str,
pub kernels: &'a [&'a str],
pub coeffs: &'a [[f32; 3]],
pub restream: f32,
}
impl LinearCostModel<'_> {
fn predicted(&self, ix: usize, m: usize, k: usize, n: usize, mr: usize, nr: usize) -> f32 {
let coeffs = &self.coeffs[ix];
let padded_work = (m.div_ceil(mr) * mr * n.div_ceil(nr) * nr * k) as f32;
let n_tiles = (m.div_ceil(mr) * n.div_ceil(nr)) as f32;
let a_restream = (m.div_ceil(mr) * mr * n.div_ceil(nr) * k) as f32;
coeffs[0] * padded_work + coeffs[1] * n_tiles + coeffs[2] + self.restream * a_restream
}
pub fn pick(
&self,
impls: &[Box<dyn MatMatMul>],
m: Option<usize>,
k: Option<usize>,
n: Option<usize>,
) -> Box<dyn MatMatMul> {
if let (Some(m), Some(k), Some(n)) = (m, k, n) {
let best = impls
.iter()
.filter_map(|imp| {
if imp.nr() == 1 && n != 1 {
return None;
}
let ix = self.kernels.iter().position(|name| *name == imp.name())?;
let t = self.predicted(ix, m, k, n, imp.mr(), imp.nr());
Some((t, imp))
})
.min_by(|a, b| order_f(a.0, b.0))
.map(|(_, imp)| imp.clone());
if let Some(best) = best {
return best;
}
}
impls.iter().find(|k| k.name() == self.default_kernel).unwrap().clone()
}
}