use crate::dsl::prelude::*;
use crate::tiling::ruda_count::{RudaCountPlan, RudaCountPlanKind, GlobalOrder, swizzle};
#[derive(RudaType, RudaLaunch)]
pub struct RudaMapping {
strategy: RudaMappingStrategy,
#[ruda(comptime)]
pub can_yield_extra_rudas: bool,
#[ruda(comptime)]
global_order: GlobalOrder,
}
#[derive(RudaType, RudaLaunch)]
#[allow(unused)] pub(crate) enum RudaMappingStrategy {
FromProblem,
SmFirst {
x_rudas: u32,
y_rudas: u32,
z_rudas: u32,
},
RudaFirst {
x_rudas: u32,
y_rudas: u32,
z_rudas: u32,
},
Flattened {
x_rudas: u32,
y_rudas: u32,
},
Spread {
x_rudas: u32,
y_rudas: u32,
z_rudas: u32,
},
}
#[ruda]
impl RudaMapping {
pub fn num_valid_rudas(&self) -> usize {
match &self.strategy {
RudaMappingStrategy::FromProblem | RudaMappingStrategy::Flattened { .. } => {
panic!("Shouldn't need to be called because the ruda count should always be exact")
}
RudaMappingStrategy::SmFirst {
x_rudas,
y_rudas,
z_rudas,
}
| RudaMappingStrategy::RudaFirst {
x_rudas,
y_rudas,
z_rudas,
}
| RudaMappingStrategy::Spread {
x_rudas,
y_rudas,
z_rudas,
} => *x_rudas as usize * *y_rudas as usize * *z_rudas as usize,
}
}
pub fn ruda_pos_to_xyz(&self) -> (u32, u32, u32) {
match &self.strategy {
RudaMappingStrategy::FromProblem => (RUDA_POS_X, RUDA_POS_Y, RUDA_POS_Z),
RudaMappingStrategy::SmFirst {
x_rudas, y_rudas, ..
} => {
self.strategy
.absolute_index_to_xyz(RUDA_POS, *x_rudas, *y_rudas, self.global_order)
}
RudaMappingStrategy::RudaFirst {
x_rudas, y_rudas, ..
} => self.strategy.absolute_index_to_xyz(
RUDA_POS_Y as usize * RUDA_COUNT_X as usize + RUDA_POS_X as usize,
*x_rudas,
*y_rudas,
self.global_order,
),
RudaMappingStrategy::Flattened { x_rudas, y_rudas } => self
.strategy
.absolute_index_to_xyz(RUDA_POS_X as usize, *x_rudas, *y_rudas, self.global_order),
RudaMappingStrategy::Spread {
x_rudas, y_rudas, ..
} => {
self.strategy
.absolute_index_to_xyz(RUDA_POS, *x_rudas, *y_rudas, self.global_order)
}
}
}
}
#[ruda]
impl RudaMappingStrategy {
fn absolute_index_to_xyz(
&self,
absolute_index: usize,
x_rudas: u32,
y_rudas: u32,
#[comptime] global_order: GlobalOrder,
) -> (u32, u32, u32) {
let z_stride = (x_rudas * y_rudas) as usize;
let z_pos = absolute_index / z_stride;
let xy_pos = absolute_index % z_stride;
let (x_pos, y_pos) = match comptime!(global_order) {
GlobalOrder::RowMajor => ((xy_pos / y_rudas as usize) as u32, xy_pos as u32 % y_rudas),
GlobalOrder::ColMajor => (xy_pos as u32 % x_rudas, (xy_pos / x_rudas as usize) as u32),
GlobalOrder::SwizzleRow(w) => {
let (x, y) = swizzle(xy_pos, y_rudas as usize, w);
(y, x)
}
GlobalOrder::SwizzleCol(w) => swizzle(xy_pos, x_rudas as usize, w),
};
(x_pos, y_pos, z_pos as u32)
}
}
pub fn ruda_mapping_launch<R: Runtime>(ruda_count_plan: &RudaCountPlan) -> RudaMappingLaunch<R> {
RudaMappingLaunch::new(
mapping_strategy(&ruda_count_plan.kind),
ruda_count_plan.kind.can_yield_extra_rudas(),
ruda_count_plan.global_order,
)
}
fn mapping_strategy<R: Runtime>(
ruda_count_plan_kind: &RudaCountPlanKind,
) -> RudaMappingStrategyArgs<R> {
match ruda_count_plan_kind {
RudaCountPlanKind::FromProblem { .. } => RudaMappingStrategyArgs::FromProblem,
RudaCountPlanKind::Sm {
rudas_first,
problem_count,
..
} => {
if *rudas_first {
RudaMappingStrategyArgs::RudaFirst {
x_rudas: problem_count.x,
y_rudas: problem_count.y,
z_rudas: problem_count.z,
}
} else {
RudaMappingStrategyArgs::SmFirst {
x_rudas: problem_count.x,
y_rudas: problem_count.y,
z_rudas: problem_count.z,
}
}
}
RudaCountPlanKind::Flattened { problem_count, .. } => RudaMappingStrategyArgs::Flattened {
x_rudas: problem_count.x,
y_rudas: problem_count.y,
},
RudaCountPlanKind::Spread { problem_count, .. } => RudaMappingStrategyArgs::Spread {
x_rudas: problem_count.x,
y_rudas: problem_count.y,
z_rudas: problem_count.z,
},
}
}