use rlx_ir::DType;
use rlx_ir::kernel_schedule::{
Access, Action, Barrier, Feature, KernelSchedule, KernelScheduleError, Layout, Region, Role,
Space, Target, lower,
};
use crate::kernel_schedule_port::BLOCK_ROLE;
pub const SHIPPING_TILE: usize = 16;
pub const EMITTED_ENTRY: &str = "sgemm_sched";
pub fn shipping_msl() -> String {
crate::kernels::msl_source()
}
#[derive(Debug, Clone, PartialEq)]
pub enum EmitError {
Unsound(Vec<KernelScheduleError>),
MissingRegion { name: &'static str },
RegionTileMismatch {
region: &'static str,
declared: Vec<usize>,
from_tile: Vec<usize>,
},
ThreadCountMismatch { from_roles: usize, from_tile: usize },
MultiRole { roles: usize },
StageDisagreement {
schedule: usize,
region: &'static str,
declared: usize,
},
UnsupportedOnMetal { feature: &'static str },
}
impl std::fmt::Display for EmitError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Unsound(errors) => {
let joined = errors
.iter()
.map(|e| e.to_string())
.collect::<Vec<_>>()
.join("; ");
write!(f, "schedule does not verify: {joined}")
}
Self::MissingRegion { name } => {
write!(f, "GEMM emitter requires a region named `{name}`")
}
Self::RegionTileMismatch {
region,
declared,
from_tile,
} => write!(
f,
"region `{region}` declares {declared:?} but the tile implies {from_tile:?}"
),
Self::ThreadCountMismatch {
from_roles,
from_tile,
} => write!(
f,
"roles imply {from_roles} threads, the tile implies {from_tile}"
),
Self::MultiRole { roles } => write!(
f,
"{roles} roles declared; this emitter lowers single-role (uniform \
threadgroup) schedules only"
),
Self::StageDisagreement {
schedule,
region,
declared,
} => write!(
f,
"schedule declares {schedule} stages but region `{region}` declares {declared}"
),
Self::UnsupportedOnMetal { feature } => write!(
f,
"`{feature}` has no Metal equivalent — refused rather than lowered to \
something else"
),
}
}
}
impl std::error::Error for EmitError {}
fn gemm_regions(tile: usize, stages: usize) -> Vec<Region> {
vec![
Region {
name: "Asub".into(),
space: Space::Shared,
dims: vec![tile, tile],
dtype: DType::F32,
stages,
layout: Layout::row_major(&[tile, tile]),
},
Region {
name: "Bsub".into(),
space: Space::Shared,
dims: vec![tile, tile],
dtype: DType::F32,
stages,
layout: Layout::row_major(&[tile, tile]),
},
Region {
name: "sum".into(),
space: Space::Register,
dims: vec![1, 1],
dtype: DType::F32,
stages: 1,
layout: Layout::row_major(&[1, 1]),
},
]
}
fn gemm_role(tile: usize) -> Vec<Role> {
let simdgroups = (tile * tile).div_ceil(32) as u32;
vec![Role {
name: BLOCK_ROLE.into(),
warps: (0..simdgroups).collect(),
}]
}
pub fn sgemm_tiled_schedule(tile: usize) -> KernelSchedule {
let mut s = KernelSchedule::new(format!("sgemm_tiled_{tile}"));
s.stages = 1;
s.requires = vec![];
s.regions = gemm_regions(tile, 1);
s.roles = gemm_role(tile);
s.barriers = vec![
Barrier {
name: "tiles_filled".into(),
producers: vec![BLOCK_ROLE.into()],
consumers: vec![BLOCK_ROLE.into()],
count: 1,
},
Barrier {
name: "tiles_consumed".into(),
producers: vec![BLOCK_ROLE.into()],
consumers: vec![BLOCK_ROLE.into()],
count: 1,
},
];
s.body.insert(
BLOCK_ROLE.into(),
vec![
Action::Load {
access: Access::plain("Asub"),
stage: 0,
},
Action::Load {
access: Access::plain("Bsub"),
stage: 0,
},
Action::Arrive {
barrier: "tiles_filled".into(),
stage: 0,
},
Action::Wait {
barrier: "tiles_filled".into(),
stage: 0,
},
Action::Compute {
reads: vec![Access::plain("Asub"), Access::plain("Bsub")],
writes: vec![Access::plain("sum")],
stage: 0,
via: None,
},
Action::Arrive {
barrier: "tiles_consumed".into(),
stage: 0,
},
Action::Wait {
barrier: "tiles_consumed".into(),
stage: 0,
},
Action::Store {
access: Access::plain("sum"),
stage: 0,
},
],
);
s
}
pub fn sgemm_tiled_pipelined_schedule(tile: usize, stages: usize) -> KernelSchedule {
let mut s = KernelSchedule::new(format!("sgemm_pipe{stages}_{tile}"));
s.stages = stages;
s.requires = vec![];
s.regions = gemm_regions(tile, stages);
s.roles = gemm_role(tile);
s.barriers = vec![Barrier {
name: "tiles_filled".into(),
producers: vec![BLOCK_ROLE.into()],
consumers: vec![BLOCK_ROLE.into()],
count: 1,
}];
let next = stages - 1;
s.body.insert(
BLOCK_ROLE.into(),
vec![
Action::Wait {
barrier: "tiles_filled".into(),
stage: 0,
},
Action::Compute {
reads: vec![Access::plain("Asub"), Access::plain("Bsub")],
writes: vec![Access::plain("sum")],
stage: 0,
via: None,
},
Action::Load {
access: Access::plain("Asub"),
stage: next,
},
Action::Load {
access: Access::plain("Bsub"),
stage: next,
},
Action::Arrive {
barrier: "tiles_filled".into(),
stage: next,
},
Action::Store {
access: Access::plain("sum"),
stage: 0,
},
],
);
s
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct EmitFacts {
pub stages: usize,
pub barriers_per_k_iter: usize,
pub threadgroup_bytes: usize,
pub threads: usize,
}
pub fn schedule_for(p: &crate::apple_params::AppleKernelParams) -> KernelSchedule {
let mut s = if p.stages_value() > 1 {
sgemm_tiled_pipelined_schedule(p.tile_value(), p.stages_value())
} else {
sgemm_tiled_schedule(p.tile_value())
};
if p.precision_value() == crate::apple_params::Precision::F16Storage {
for r in s.regions.iter_mut() {
if r.space == Space::Shared {
r.dtype = DType::F16;
}
}
}
s
}
pub fn emit_for(
p: &crate::apple_params::AppleKernelParams,
target: Target,
) -> Result<(String, EmitFacts), EmitError> {
emit_msl_with(&schedule_for(p), p.tile_value(), target, Some(p))
}
pub fn emit_msl(
sched: &KernelSchedule,
tile: usize,
target: Target,
) -> Result<(String, EmitFacts), EmitError> {
emit_msl_with(sched, tile, target, None)
}
pub fn emit_msl_with(
sched: &KernelSchedule,
tile: usize,
target: Target,
params: Option<&crate::apple_params::AppleKernelParams>,
) -> Result<(String, EmitFacts), EmitError> {
if sched.roles.len() != 1 {
return Err(EmitError::MultiRole {
roles: sched.roles.len(),
});
}
if sched.requires.contains(&Feature::AsyncCopy) {
return Err(EmitError::UnsupportedOnMetal {
feature: "AsyncCopy",
});
}
let lowered = lower(sched, target).map_err(EmitError::Unsound)?;
if lowered.threads != tile * tile {
return Err(EmitError::ThreadCountMismatch {
from_roles: lowered.threads,
from_tile: tile * tile,
});
}
let region = |name: &'static str| -> Result<&Region, EmitError> {
sched
.regions
.iter()
.find(|r| r.name == name)
.ok_or(EmitError::MissingRegion { name })
};
let a = region("Asub")?;
let b = region("Bsub")?;
region("sum")?;
for (r, name) in [(a, "Asub"), (b, "Bsub")] {
if r.dims != [tile, tile] {
return Err(EmitError::RegionTileMismatch {
region: name,
declared: r.dims.clone(),
from_tile: vec![tile, tile],
});
}
}
let stages = sched.stages.max(1);
for (r, name) in [(a, "Asub"), (b, "Bsub")] {
if r.stages.max(1) != stages {
return Err(EmitError::StageDisagreement {
schedule: stages,
region: name,
declared: r.stages,
});
}
}
let body = sched.body.get(BLOCK_ROLE).map(Vec::as_slice).unwrap_or(&[]);
let barriers_per_k_iter = body
.iter()
.filter(|x| matches!(x, Action::Wait { .. }))
.count()
.max(1);
let facts = EmitFacts {
stages,
barriers_per_k_iter,
threadgroup_bytes: lowered.shared_bytes,
threads: lowered.threads,
};
use crate::apple_params::{Precision, SyncScope};
let stage_ty = match params.map(|p| p.precision_value()) {
Some(Precision::F16Storage) => "half",
_ => "float",
};
let barrier = match params.map(|p| p.sync_value()) {
Some(SyncScope::Simdgroup) => "simdgroup_barrier(mem_flags::mem_threadgroup);",
_ => "threadgroup_barrier(mem_flags::mem_threadgroup);",
};
let mut src = String::with_capacity(4096);
src.push_str(&format!(
"// @generated by rlx_metal::kernel_schedule_emit from KernelSchedule `{}`.\n\
// Do not edit; edit the schedule.\n\
//\n\
// Derived, not authored:\n\
// stages = {} (KernelSchedule::stages / Region::stages)\n\
// barriers per K iter = {} (count of Action::Wait in the role body)\n\
// threadgroup bytes = {} (lower().shared_bytes)\n\
// threads/group = {} (lower().threads, from Role::warps)\n\
// buffer offsets = {:?}\n\
#include <metal_stdlib>\n\
using namespace metal;\n\n\
constant uint TILE = {tile};\n\
constant uint STAGES = {};\n\n",
sched.name,
facts.stages,
facts.barriers_per_k_iter,
facts.threadgroup_bytes,
facts.threads,
lowered.region_offsets,
stages,
));
src.push_str(&format!(
"kernel void {EMITTED_ENTRY}(\n\
\x20 device const float* A [[buffer(0)]],\n\
\x20 device const float* B [[buffer(1)]],\n\
\x20 device float* C [[buffer(2)]],\n\
\x20 constant uint& M [[buffer(3)]],\n\
\x20 constant uint& K [[buffer(4)]],\n\
\x20 constant uint& N [[buffer(5)]],\n\
\x20 uint2 gid [[thread_position_in_grid]],\n\
\x20 uint2 tid [[thread_position_in_threadgroup]],\n\
\x20 uint2 tgid [[threadgroup_position_in_grid]]\n\
) {{\n"
));
src.push_str(&format!(
" threadgroup {stage_ty} Asub[STAGES][TILE][TILE];\n\
\x20 threadgroup {stage_ty} Bsub[STAGES][TILE][TILE];\n\n\
\x20 uint row = tgid.y * TILE + tid.y;\n\
\x20 uint col = tgid.x * TILE + tid.x;\n\n\
\x20 float sum = 0.0;\n\
\x20 uint num_tiles = (K + TILE - 1) / TILE;\n\n"
));
let stage_load = |t_expr: &str, buf_expr: &str| -> String {
format!(
" {{\n\
\x20 uint a_col = ({t_expr}) * TILE + tid.x;\n\
\x20 uint b_row = ({t_expr}) * TILE + tid.y;\n\
\x20 Asub[{buf_expr}][tid.y][tid.x] = (row < M && a_col < K) ? A[row * K + a_col] : 0.0;\n\
\x20 Bsub[{buf_expr}][tid.y][tid.x] = (b_row < K && col < N) ? B[b_row * N + col] : 0.0;\n\
\x20 }}\n"
)
};
let compute = |buf_expr: &str| -> String {
format!(
" for (uint k = 0; k < TILE; ++k) {{\n\
\x20 sum += Asub[{buf_expr}][tid.y][k] * Bsub[{buf_expr}][k][tid.x];\n\
\x20 }}\n"
)
};
if stages > 1 {
src.push_str(" // Prologue: fill STAGES-1 buffers before the first compute.\n");
src.push_str(" for (uint s = 0; s + 1 < STAGES; ++s) {\n");
src.push_str(" if (s < num_tiles) {\n");
src.push_str(&stage_load("s", "s"));
src.push_str(" }\n }\n\n");
src.push_str(" for (uint t = 0; t < num_tiles; ++t) {\n");
src.push_str(&format!(
" {barrier} // Action::Wait `tiles_filled` \
({barriers_per_k_iter}/iter declared)\n\n"
));
src.push_str(&compute("t % STAGES"));
src.push_str(
"\n // Stage the tile needed STAGES-1 iterations from now, into the\n\
\x20 // buffer iteration t-1 finished reading before the barrier above.\n\
\x20 // That ordering is why no second barrier is needed here.\n\
\x20 uint kt = t + STAGES - 1;\n\
\x20 if (kt < num_tiles) {\n",
);
src.push_str(&stage_load("kt", "kt % STAGES"));
src.push_str(" }\n }\n");
} else {
src.push_str(" for (uint t = 0; t < num_tiles; ++t) {\n");
src.push_str(&stage_load("t", "0"));
src.push_str(&format!(
" {barrier} // Action::Wait `tiles_filled`\n\n"
));
src.push_str(&compute("0"));
src.push_str(&format!(
"\n {barrier} // Action::Wait `tiles_consumed`\n }}\n"
));
}
src.push_str(
"\n if (row < M && col < N) {\n\
\x20 C[row * N + col] = sum;\n\
\x20 }\n}\n",
);
Ok((src, facts))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kernel_schedule_port::METAL_TARGET;
#[test]
fn the_shipping_schedule_emits_one_stage_and_two_barriers() {
let (src, facts) = emit_msl(
&sgemm_tiled_schedule(SHIPPING_TILE),
SHIPPING_TILE,
METAL_TARGET,
)
.unwrap();
assert_eq!(facts.stages, 1);
assert_eq!(facts.barriers_per_k_iter, 2);
assert_eq!(facts.threads, 256);
assert_eq!(facts.threadgroup_bytes, 2048);
assert_eq!(src.matches("threadgroup_barrier").count(), 2);
assert!(src.contains("constant uint STAGES = 1;"));
}
#[test]
fn the_pipelined_schedule_emits_rotation_and_one_barrier() {
let sched = sgemm_tiled_pipelined_schedule(SHIPPING_TILE, 2);
let (src, facts) = emit_msl(&sched, SHIPPING_TILE, METAL_TARGET).unwrap();
assert_eq!(facts.stages, 2);
assert_eq!(facts.barriers_per_k_iter, 1);
assert_eq!(facts.threadgroup_bytes, 4096);
assert_eq!(src.matches("threadgroup_barrier").count(), 1);
assert!(src.contains("t % STAGES"));
assert!(src.contains("kt % STAGES"));
}
#[test]
fn stage_depth_reaches_the_emitted_source() {
for stages in 2..=4 {
let sched = sgemm_tiled_pipelined_schedule(SHIPPING_TILE, stages);
let (src, facts) = emit_msl(&sched, SHIPPING_TILE, METAL_TARGET).unwrap();
assert!(src.contains(&format!("constant uint STAGES = {stages};")));
assert_eq!(facts.threadgroup_bytes, 2048 * stages);
}
}
#[test]
fn async_copy_is_refused_on_metal() {
let mut sched = sgemm_tiled_pipelined_schedule(SHIPPING_TILE, 2);
sched.requires = vec![Feature::AsyncCopy];
assert!(matches!(
emit_msl(&sched, SHIPPING_TILE, METAL_TARGET),
Err(EmitError::UnsupportedOnMetal { .. })
));
}
#[test]
fn an_over_deep_rotation_busts_the_threadgroup_budget() {
let sched = sgemm_tiled_pipelined_schedule(64, 2);
let err = emit_msl(&sched, 64, METAL_TARGET).unwrap_err();
assert!(
matches!(&err, EmitError::Unsound(es)
if es.iter().any(|e| matches!(
e, KernelScheduleError::SharedOverBudget { .. }))),
"expected a shared-budget finding, got {err}"
);
}
#[test]
fn params_drive_the_emitted_kernel_end_to_end() {
use crate::apple_params::{AppleKernelParams, Precision, SyncScope};
let p = AppleKernelParams::default()
.stages(2)
.precision(Precision::F16Storage)
.sync(SyncScope::Simdgroup);
let (src, facts) = emit_for(&p, METAL_TARGET).expect("params emit");
assert_eq!(facts.stages, 2, "stage depth did not reach the schedule");
assert!(
src.contains("threadgroup half Asub"),
"precision did not reach the MSL"
);
assert!(
src.contains("simdgroup_barrier"),
"sync scope did not reach the MSL"
);
assert!(
!src.contains("threadgroup_barrier(mem"),
"the old barrier survived"
);
assert_eq!(
facts.threadgroup_bytes,
p.threadgroup_bytes(),
"EmitFacts and AppleKernelParams disagree about threadgroup bytes"
);
assert_eq!(
facts.threadgroup_bytes, 2048,
"2 stages x 16x16 x f16 accounting"
);
}
#[test]
fn default_params_reproduce_the_shipping_structure() {
use crate::apple_params::AppleKernelParams;
let (src, facts) =
emit_for(&AppleKernelParams::default(), METAL_TARGET).expect("default emits");
assert_eq!(facts.stages, 1);
assert_eq!(facts.barriers_per_k_iter, 2);
assert!(src.contains("threadgroup float Asub"));
assert_eq!(src.matches("threadgroup_barrier").count(), 2);
}
#[test]
fn an_unverifiable_schedule_is_rejected_before_emission() {
let mut sched = sgemm_tiled_schedule(SHIPPING_TILE);
sched.barriers[0].producers.clear();
assert!(matches!(
emit_msl(&sched, SHIPPING_TILE, METAL_TARGET),
Err(EmitError::Unsound(_))
));
}
#[test]
fn the_emitted_staging_does_not_drift_from_sgemm_tiled() {
let (src, _) = emit_msl(
&sgemm_tiled_schedule(SHIPPING_TILE),
SHIPPING_TILE,
METAL_TARGET,
)
.unwrap();
let shipping = shipping_msl();
for fragment in [
"(row < M && a_col < K) ? A[row * K + a_col] : 0.0",
"(b_row < K && col < N) ? B[b_row * N + col] : 0.0",
"C[row * N + col] = sum;",
] {
assert!(
shipping.contains(fragment),
"sgemm_tiled no longer contains `{fragment}` — update the emitter"
);
assert!(
src.contains(fragment),
"the emitted kernel is missing `{fragment}`"
);
}
}
}