use crate::isa::Arch;
#[cfg(test)]
use crate::isa::Isa;
use crate::isa::IsaSet;
use crate::mmm::{Query, Suitable};
use tract_data::prelude::DatumType;
pub struct MmmTier {
pub arch: Option<Arch>,
pub precedence: u8,
pub name: &'static str,
pub applies: fn(&IsaSet) -> bool,
pub preferred: fn(&IsaSet, DatumType, &Query, &[Suitable]) -> Option<&'static str>,
}
inventory::collect!(MmmTier);
pub fn declared() -> impl Iterator<Item = &'static MmmTier> {
inventory::iter::<MmmTier>()
}
pub fn for_isa(isa: &IsaSet) -> Vec<&'static MmmTier> {
let arch = isa.arch();
let mut tiers: Vec<&'static MmmTier> = declared()
.filter(|t| t.arch.is_none() || t.arch == arch)
.filter(|t| (t.applies)(isa))
.collect();
tiers.sort_by_key(|t| std::cmp::Reverse(t.precedence));
log::debug!(
"mmm tiers for {isa:?}: {}",
tiers.iter().map(|t| t.name).collect::<Vec<_>>().join(" > ")
);
tiers
}
pub fn preferred(
isa: &IsaSet,
tiers: &[&'static MmmTier],
accumulator: DatumType,
query: &Query,
suitable: &[Suitable],
) -> Option<usize> {
tiers.iter().find_map(|t| {
let name = (t.preferred)(isa, accumulator, query, suitable)?;
crate::mmm::suitable_named(suitable, name)
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn precedence_is_unique_per_arch() {
let tiers: Vec<&MmmTier> = declared().collect();
for (ix, a) in tiers.iter().enumerate() {
for b in &tiers[ix + 1..] {
assert!(
a.arch != b.arch || a.precedence != b.precedence,
"tiers {} and {} both claim {:?} precedence {}",
a.name,
b.name,
a.arch,
a.precedence
);
}
}
}
#[test]
fn a_tier_only_needs_its_own_architecture() {
for tier in declared() {
let Some(arch) = tier.arch else { continue };
for isa in Isa::ALL.into_iter().filter(|i| !i.is_arch()) {
if (tier.applies)(&IsaSet::of_arch(arch).with(isa))
!= (tier.applies)(&IsaSet::of_arch(arch))
{
assert_eq!(
isa.arch(),
arch,
"tier {} is {arch:?} but reacts to {isa}, which is {:?}",
tier.name,
isa.arch()
);
}
}
}
}
#[test]
fn a_tier_that_answers_names_a_reachable_kernel() {
let mut answered = std::collections::HashSet::new();
let mut reached = std::collections::HashSet::new();
let native = crate::isa::native();
for isa in IsaSet::every_ladder().filter(|l| l.iter().all(|i| native.has(i))) {
let dispatch = crate::MmmDispatch::for_isa(isa);
for acc in [DatumType::F32, DatumType::F16, DatumType::I32] {
let query = Query::plain(acc, None, None, None);
let suitable = dispatch.suitable(&query);
for tier in dispatch.tiers() {
let Some(name) = (tier.preferred)(&isa, acc, &query, &suitable) else {
continue;
};
answered.insert(tier.name);
if crate::mmm::suitable_named(&suitable, name).is_some() {
reached.insert(tier.name);
}
}
}
}
let unheard: Vec<&str> =
answered.iter().copied().filter(|name| !reached.contains(name)).collect();
assert!(unheard.is_empty(), "these tiers name a kernel no machine can reach: {unheard:?}");
}
}