use cubecl::prelude::*;
use crate::PartitionSize;
#[derive(Copy, Clone, Debug, Hash, PartialEq, Eq)]
pub enum PartitionSchedulerScheme {
Offset,
Naive,
}
#[derive(CubeType)]
pub struct PartitionScheduler {
pub m: AxisScheduler,
pub n: AxisScheduler,
pub k: AxisScheduler,
}
#[cube]
impl PartitionScheduler {
pub fn new(
partition_index_m: u32,
partition_index_n: u32,
#[comptime] partition_size: PartitionSize,
#[comptime] partition_schedule_scheme: PartitionSchedulerScheme,
) -> PartitionScheduler {
match partition_schedule_scheme {
PartitionSchedulerScheme::Offset => {
let m_offset = (partition_index_n / partition_size.k()) % partition_size.m();
let n_offset = (partition_index_m / partition_size.k()) % partition_size.n();
let k_offset = (partition_index_m + partition_index_n) % partition_size.k();
PartitionScheduler {
m: AxisScheduler::new_Offset(OffsetAxisScheduler::new(
m_offset,
partition_index_m,
partition_size.m(),
)),
n: AxisScheduler::new_Offset(OffsetAxisScheduler::new(
n_offset,
partition_index_n,
partition_size.n(),
)),
k: AxisScheduler::new_Offset(OffsetAxisScheduler::new(
k_offset,
0u32,
partition_size.k(),
)),
}
}
PartitionSchedulerScheme::Naive => PartitionScheduler {
m: AxisScheduler::new_Naive(NaiveAxisScheduler::new(
partition_index_m,
partition_size.m(),
)),
n: AxisScheduler::new_Naive(NaiveAxisScheduler::new(
partition_index_n,
partition_size.n(),
)),
k: AxisScheduler::new_Naive(NaiveAxisScheduler::new(0u32, partition_size.k())),
},
}
}
pub fn map_m(&self, i: u32) -> u32 {
self.m.map(i)
}
pub fn map_n(&self, i: u32) -> u32 {
self.n.map(i)
}
pub fn map_k(&self, i: u32) -> u32 {
self.k.map(i)
}
}
#[derive(CubeType)]
#[allow(unused)]
pub enum AxisScheduler {
Offset(OffsetAxisScheduler),
Naive(NaiveAxisScheduler),
}
#[derive(CubeType)]
pub struct OffsetAxisScheduler {
inner_offset: u32,
outer_offset: u32,
#[cube(comptime)]
len: u32,
}
#[derive(CubeType)]
pub struct NaiveAxisScheduler {
outer_offset: u32,
}
#[cube]
impl AxisScheduler {
pub fn map(&self, i: u32) -> u32 {
match self {
AxisScheduler::Offset(offset_axis_scheduler) => offset_axis_scheduler.map(i),
AxisScheduler::Naive(naive_axis_scheduler) => naive_axis_scheduler.map(i),
}
}
}
#[cube]
impl OffsetAxisScheduler {
pub fn new(
inner_offset: u32,
partition_index: u32,
#[comptime] len: u32,
) -> OffsetAxisScheduler {
let outer_offset = partition_index * len;
OffsetAxisScheduler {
inner_offset,
outer_offset,
len,
}
}
pub fn map(&self, i: u32) -> u32 {
let relative = (i + self.inner_offset) % self.len;
relative + self.outer_offset
}
}
#[cube]
impl NaiveAxisScheduler {
pub fn new(partition_index: u32, #[comptime] len: u32) -> NaiveAxisScheduler {
let outer_offset = partition_index * len;
NaiveAxisScheduler { outer_offset }
}
pub fn map(&self, i: u32) -> u32 {
i + self.outer_offset
}
}
#[derive(Default, Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub enum PartitionBuffering {
Single,
#[default]
Double,
}