use crate::components::{
MatmulIdent,
global::{MaxGlobalReaderPlanes, specialization::roles::PlaneRoles},
};
#[derive(Default, Copy, Clone, Debug, Hash, PartialEq, Eq)]
pub struct LoadSpecializationConfig {
pub lhs: SpecializationTensorConfig,
pub rhs: SpecializationTensorConfig,
}
#[derive(Default, Copy, Clone, Debug, Hash, PartialEq, Eq)]
pub enum SpecializationTensorConfig {
#[default]
MainFlowOnly,
LoadFlowOnly,
}
impl LoadSpecializationConfig {
pub fn has_specialization(&self) -> bool {
self.lhs.has_specialization() || self.rhs.has_specialization()
}
}
impl SpecializationTensorConfig {
pub fn has_specialization(&self) -> bool {
match self {
SpecializationTensorConfig::MainFlowOnly => false,
SpecializationTensorConfig::LoadFlowOnly => true,
}
}
}
impl LoadSpecializationConfig {
pub fn to_plane_roles(
&self,
main_flow: u32,
reader_tasks: MaxGlobalReaderPlanes,
) -> PlaneRoles {
use SpecializationTensorConfig::*;
let ideal_load_only = match (self.lhs, self.rhs) {
(MainFlowOnly, MainFlowOnly) => 0,
(MainFlowOnly, LoadFlowOnly) => reader_tasks.rhs,
(LoadFlowOnly, MainFlowOnly) => reader_tasks.lhs,
(LoadFlowOnly, LoadFlowOnly) => gcd(reader_tasks.lhs, reader_tasks.rhs),
};
let load_only = best_divisor_close_to_reference(ideal_load_only, main_flow);
PlaneRoles {
main_flow,
load_only,
}
}
}
fn best_divisor_close_to_reference(dividible_value: u32, reference: u32) -> u32 {
let mut best = 1;
let mut best_dist = reference.abs_diff(1);
for d in 1..=dividible_value {
if dividible_value.is_multiple_of(d) {
let dist = reference.abs_diff(d);
if dist < best_dist || (dist == best_dist && d > best) {
best = d;
best_dist = dist;
}
}
}
best
}
#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
pub enum LoadingSides {
Both,
Lhs,
Rhs,
None,
}
impl LoadingSides {
pub fn includes_lhs(&self) -> bool {
self.includes(MatmulIdent::Lhs)
}
pub fn includes_rhs(&self) -> bool {
self.includes(MatmulIdent::Rhs)
}
pub fn includes(&self, ident: MatmulIdent) -> bool {
matches!(
(self, ident),
(LoadingSides::Both, _)
| (LoadingSides::Lhs, MatmulIdent::Lhs)
| (LoadingSides::Rhs, MatmulIdent::Rhs)
)
}
}
#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
pub struct SpecializedLoadingSides {
pub main_flow: LoadingSides,
pub load_only: LoadingSides,
}
impl SpecializedLoadingSides {
pub fn num_loading_planes(
&self,
specialized: bool,
ident: MatmulIdent,
plane_roles: PlaneRoles,
) -> u32 {
if specialized {
let mut num_loading_planes = 0;
if self.main_flow.includes(ident) {
num_loading_planes += plane_roles.main_flow;
}
if self.load_only.includes(ident) {
num_loading_planes += plane_roles.load_only;
}
num_loading_planes
} else {
plane_roles.main_flow
}
}
}
impl From<LoadSpecializationConfig> for SpecializedLoadingSides {
fn from(lsc: LoadSpecializationConfig) -> Self {
use SpecializationTensorConfig::*;
match (lsc.lhs, lsc.rhs) {
(MainFlowOnly, MainFlowOnly) => SpecializedLoadingSides {
main_flow: LoadingSides::Both,
load_only: LoadingSides::None,
},
(MainFlowOnly, LoadFlowOnly) => SpecializedLoadingSides {
main_flow: LoadingSides::Lhs,
load_only: LoadingSides::Rhs,
},
(LoadFlowOnly, MainFlowOnly) => SpecializedLoadingSides {
main_flow: LoadingSides::Rhs,
load_only: LoadingSides::Lhs,
},
(LoadFlowOnly, LoadFlowOnly) => SpecializedLoadingSides {
main_flow: LoadingSides::None,
load_only: LoadingSides::Both,
},
}
}
}
pub(crate) fn gcd(mut a: u32, mut b: u32) -> u32 {
while b != 0 {
let r = a % b;
a = b;
b = r;
}
a
}