use crate::{MatmulProblemSize, TileSize};
pub fn find_instruction_size<IsSupported, FallbackSizes>(
problem_size: MatmulProblemSize,
(tm, tn, tk): (Option<u32>, Option<u32>, Option<u32>),
is_supported: IsSupported,
fallback_sizes: FallbackSizes,
) -> Option<TileSize>
where
IsSupported: Fn(u32, u32, u32) -> bool,
FallbackSizes: Fn() -> Vec<TileSize>,
{
let matches_forced = |m: u32, n: u32, k: u32| {
tm.is_none_or(|v| m == v) && tn.is_none_or(|v| n == v) && tk.is_none_or(|v| k == v)
};
let try_candidate = |m: u32, n: u32, k: u32| {
(is_supported(m, n, k) && matches_forced(m, n, k)).then(|| TileSize::from((m, n, k)))
};
let (m, n) = (problem_size.m, problem_size.n);
if m >= 4 * n
&& let Some(ts) = try_candidate(32, 8, 16)
{
return Some(ts);
}
if n >= 4 * m
&& let Some(ts) = try_candidate(8, 32, 16)
{
return Some(ts);
}
if let Some(ts) = try_candidate(16, 16, 16) {
return Some(ts);
}
if let Some(ts) = try_candidate(8, 8, 8) {
return Some(ts);
}
fallback_sizes()
.into_iter()
.find(|ts| matches_forced(ts.m, ts.n, ts.k))
}