use super::Suitable;
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<'a> LinearCostModel<'a> {
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 preferred(
&self,
suitable: &[Suitable],
m: Option<usize>,
k: Option<usize>,
n: Option<usize>,
) -> Option<&'a str> {
if let (Some(m), Some(k), Some(n)) = (m, k, n) {
let best = suitable
.iter()
.filter_map(|(mmm, _, _)| {
if mmm.nr() == 1 && n != 1 {
return None;
}
let ix = self.kernels.iter().position(|name| *name == mmm.name())?;
let t = self.predicted(ix, m, k, n, mmm.mr(), mmm.nr());
Some((t, self.kernels[ix]))
})
.min_by(|a, b| order_f(a.0, b.0))
.map(|(_, name)| name);
if best.is_some() {
return best;
}
}
Some(self.default_kernel)
}
}