use cubecl_core as cubecl;
use cubecl_core::prelude::*;
use cubecl_std::{CubeOption, CubeOptionExpand, tensor::layout::Coords2d};
use crate::components::{AccS, global::MaxGlobalReaderPlanes};
use crate::components::{
AvailableLineSizes, LhsS, MatmulLineSizes, MatmulSelection, RhsS, StageIdent,
};
use crate::components::{
MatmulPrecision, MatmulProblem, MatrixLayout, TilingScheme,
global::{self, PlaneRoleConfig, RoleRuleConfig},
tile::TileConfig,
};
use crate::components::{
error::MatmulSetupError, global::WriteEventListener, stage::StageMemoryConfig,
};
use crate::components::{
stage::{NumStages, PartitionScheduler, PartitionSchedulerScheme},
tile::io::TileKind,
};
use std::{fmt::Debug, hash::Hash};
use super::{StageEventListener, TilingLayout};
pub trait StageMatmulFamily: Send + Sync + 'static {
type Matmul<MP: MatmulPrecision, TL: TilingLayout, TR: TilingLayout, TA: TilingLayout, TO: TilingLayout>: StageMatmul<
MP,
Config = Self::Config,
LhsStage = <Self::LhsStage as StageFamily>::Stage<LhsS<MP>, TL>,
RhsStage = <Self::RhsStage as StageFamily>::Stage<RhsS<MP>, TR>,
AccStage = <Self::AccStage as StageFamily>::Stage<AccS<MP>, TA>,
OutStage = <Self::OutStage as StageFamily<ReadWrite>>::Stage<AccS<MP>, TO>,
>;
type LhsStage: StageFamily;
type RhsStage: StageFamily;
type AccStage: StageFamily;
type OutStage: StageFamily<ReadWrite>;
type Config: StageConfig;
fn setup<MP: MatmulPrecision, R: Runtime>(
client: &ComputeClient<R::Server>,
problem: &MatmulProblem,
selection: &MatmulSelection,
line_sizes: &MatmulLineSizes,
num_stages: NumStages,
max_global_readers: Option<MaxGlobalReaderPlanes>,
ordered: bool,
) -> Result<Self::Config, MatmulSetupError>;
fn filter_line_sizes(available_line_sizes: AvailableLineSizes) -> AvailableLineSizes {
available_line_sizes
}
}
#[cube]
pub trait StageMatmul<MP: MatmulPrecision>: 'static + Send + Sync {
type Config: StageConfig;
type Accumulators: CubeType;
type LhsStage: CubeType;
type RhsStage: CubeType;
type AccStage: CubeType;
type OutStage: CubeType;
type LhsTile: CubeType;
type RhsTile: CubeType;
fn execute(
lhs: &Self::LhsStage,
rhs: &Self::RhsStage,
instruction_lhs: &mut Self::LhsTile,
instruction_rhs: &mut Self::RhsTile,
acc: &mut Self::Accumulators,
#[comptime] config: Self::Config,
partition_scheduler: &PartitionScheduler,
);
fn execute_with_listener<SEL: StageEventListener<Self::Config>>(
lhs: &Self::LhsStage,
rhs: &Self::RhsStage,
instruction_lhs: &mut Self::LhsTile,
instruction_rhs: &mut Self::RhsTile,
acc: &mut Self::Accumulators,
#[comptime] config: Self::Config,
listener: SEL,
partition_scheduler: &PartitionScheduler,
);
fn init_tile_inputs(#[comptime] config: Self::Config) -> (Self::LhsTile, Self::RhsTile);
fn init_accumulators(#[comptime] config: Self::Config) -> Self::Accumulators;
fn load_accumulators(
reader: &Self::AccStage,
acc: &mut Self::Accumulators,
#[comptime] config: Self::Config,
);
fn write_results<W: WriteEventListener, G: global::GlobalConfig>(
acc: &Self::Accumulators,
stage: &mut Self::OutStage,
listener: &mut W,
partition_scheduler: &PartitionScheduler,
#[comptime] stage_config: Self::Config,
#[comptime] global_config: G,
);
fn init_scheduler(#[comptime] config: Self::Config) -> PartitionScheduler;
}
pub trait StageConfig:
Copy + Clone + Eq + PartialEq + Hash + Debug + Send + Sync + 'static
{
type TileConfig: TileConfig;
fn tile_config(self) -> Self::TileConfig;
fn stage_memory_config(self, ident: StageIdent) -> StageMemoryConfig {
let tiling = self.tiling_scheme();
StageMemoryConfig {
num_main_flow_planes: self.num_main_flow_planes(),
elements_in_tile_row: tiling.elements_in_tile_row(ident),
elements_in_tile_col: tiling.elements_in_tile_col(ident),
tiles_in_stage_row: tiling.tiles_in_stage_row(ident),
tiles_in_stage_col: tiling.tiles_in_stage_col(ident),
stage_line_size: self.stage_line_size(ident),
matrix_layout: self.matrix_layout(ident),
num_stages: self.num_stages(ident),
}
}
fn stage_line_size(&self, ident: StageIdent) -> u32;
fn global_line_size(&self, ident: StageIdent) -> u32;
fn matrix_layout(&self, ident: StageIdent) -> MatrixLayout;
fn plane_dim(&self) -> u32;
fn partition_buffering(&self) -> PartitionBuffering;
fn tiling_scheme(&self) -> TilingScheme;
fn plane_role_config(&self) -> PlaneRoleConfig;
fn role_rule_config(&self) -> RoleRuleConfig;
fn num_main_flow_planes(&self) -> u32;
fn quantized(&self) -> bool;
fn must_sync_plane_after_execution(&self) -> bool;
fn partition_schedule_scheme(&self) -> PartitionSchedulerScheme;
fn num_stages(&self, ident: StageIdent) -> u32;
}
#[derive(Default, Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub enum PartitionBuffering {
Single,
#[default]
Double,
}
#[cube]
pub trait Stage<ES: Numeric, IO: SliceVisibility = ReadOnly>:
CubeType + Send + Sync + 'static
{
type TileKind: TileKind<IO>;
fn tile(this: &Self, tile: Coords2d) -> <Self::TileKind as TileKind<IO>>::Tile<ES>;
}
pub trait StageFamily<IO: SliceVisibility = ReadOnly>: Send + Sync + 'static {
type TileKind: TileKind<IO>;
type Stage<ES: Numeric, T: TilingLayout>: Stage<ES, IO, TileKind = Self::TileKind>;
}
#[cube]
impl<ES: Numeric, IO: SliceVisibility, Inner: Stage<ES, IO>> Stage<ES, IO> for CubeOption<Inner> {
type TileKind = CubeOption<Inner::TileKind>;
fn tile(this: &Self, tile: Coords2d) -> <Self::TileKind as TileKind<IO>>::Tile<ES> {
match this {
CubeOption::Some(stage) => CubeOption::new_Some(Inner::tile(stage, tile)),
CubeOption::None => CubeOption::new_None(),
}
}
}
impl<IO: SliceVisibility, Inner: StageFamily<IO>> StageFamily<IO> for Option<Inner> {
type TileKind = CubeOption<Inner::TileKind>;
type Stage<ES: Numeric, T: TilingLayout> = CubeOption<Inner::Stage<ES, T>>;
}