use crate::occupancy::{self, AppleGpuFamily, Verdict};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum SyncScope {
#[default]
Threadgroup,
Simdgroup,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Precision {
#[default]
F32,
F16Storage,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Encode {
#[default]
PerDispatch,
Indirect,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct AppleKernelParams {
stages: usize,
sync: SyncScope,
precision: Precision,
tile: usize,
encode: Encode,
}
impl Default for AppleKernelParams {
fn default() -> Self {
Self {
stages: 1,
sync: SyncScope::Threadgroup,
precision: Precision::F32,
tile: 16,
encode: Encode::PerDispatch,
}
}
}
impl AppleKernelParams {
pub fn stages(mut self, n: usize) -> Self {
self.stages = n.max(1);
self
}
pub fn sync(mut self, s: SyncScope) -> Self {
self.sync = s;
self
}
pub fn precision(mut self, p: Precision) -> Self {
self.precision = p;
self
}
pub fn tile(mut self, edge: usize) -> Self {
self.tile = edge.max(1);
self
}
pub fn encode(mut self, e: Encode) -> Self {
self.encode = e;
self
}
pub fn stages_value(&self) -> usize {
self.stages
}
pub fn sync_value(&self) -> SyncScope {
self.sync
}
pub fn precision_value(&self) -> Precision {
self.precision
}
pub fn tile_value(&self) -> usize {
self.tile
}
pub fn encode_value(&self) -> Encode {
self.encode
}
pub fn threadgroup_bytes(&self) -> usize {
let elem = match self.precision {
Precision::F32 => 4,
Precision::F16Storage => 2,
};
2 * self.tile * self.tile * elem * self.stages
}
pub fn predict(&self, chip: AppleGpuFamily) -> Verdict {
occupancy::predict(
chip,
Self::default().tile(self.tile).threadgroup_bytes(),
self.threadgroup_bytes(),
)
}
pub fn defines(&self) -> String {
format!(
"#define RLX_TS {}\n\
#define RLX_STAGES {}\n\
#define RLX_STAGE_T {}\n\
#define RLX_BARRIER() {}\n",
self.tile,
self.stages,
match self.precision {
Precision::F32 => "float",
Precision::F16Storage => "half",
},
match self.sync {
SyncScope::Threadgroup => "threadgroup_barrier(mem_flags::mem_threadgroup)",
SyncScope::Simdgroup => "simdgroup_barrier(mem_flags::mem_threadgroup)",
}
)
}
pub fn parse(spec: &str) -> Result<Self, String> {
let mut p = Self::default();
for field in spec.split(',').map(str::trim).filter(|s| !s.is_empty()) {
let (k, v) = field
.split_once('=')
.ok_or_else(|| format!("`{field}` is not `key=value`"))?;
let v = v.trim();
match k.trim() {
"stages" => p.stages = v.parse().map_err(|_| format!("stages={v:?}"))?,
"tile" => p.tile = v.parse().map_err(|_| format!("tile={v:?}"))?,
"sync" => {
p.sync = match v {
"threadgroup" | "tg" => SyncScope::Threadgroup,
"simd" | "simdgroup" => SyncScope::Simdgroup,
_ => return Err(format!("sync={v:?} (threadgroup|simd)")),
}
}
"precision" => {
p.precision = match v {
"f32" => Precision::F32,
"f16" => Precision::F16Storage,
_ => return Err(format!("precision={v:?} (f32|f16)")),
}
}
"encode" => {
p.encode = match v {
"dispatch" | "per-dispatch" => Encode::PerDispatch,
"icb" | "indirect" => Encode::Indirect,
_ => return Err(format!("encode={v:?} (dispatch|icb)")),
}
}
other => return Err(format!("unknown key `{other}`")),
}
}
if p.stages == 0 || p.tile == 0 {
return Err("stages and tile must be >= 1".into());
}
Ok(p)
}
pub fn for_shape(m: usize, _k: usize, _n: usize) -> Self {
let tile = if m < 32 { 8 } else { 32 };
Self::default().tile(tile)
}
pub const KEYS: &'static [(&'static str, &'static str)] = &[
("stages", "N (rotation depth; 1 disables)"),
("sync", "threadgroup|simd"),
("precision", "f32|f16"),
("tile", "N (threadgroup tile edge)"),
("encode", "dispatch|icb"),
];
pub fn help() -> String {
let mut out = String::from(
"Apple kernel parameters — the same keys in all three front doors:\n\
\x20 builder : AppleKernelParams::default().stages(2).sync(SyncScope::Simdgroup)\n\
\x20 config : RLX_METAL_PARAMS=\"stages=2,sync=simd\"\n\
\x20 cli : --metal-params stages=2,sync=simd OR --stages 2 --sync simd\n\n",
);
for (k, v) in Self::KEYS {
out.push_str(&format!(" {k:<10} {v}\n"));
}
out.push_str("\nprecedence: CLI > config > default\n");
out
}
pub fn spec_from_args<S: AsRef<str>>(args: &[S]) -> Result<Option<String>, String> {
let a: Vec<&str> = args.iter().map(|s| s.as_ref()).collect();
let mut parts: Vec<String> = Vec::new();
let mut i = 0;
while i < a.len() {
let tok = a[i];
if tok == "--metal-params" {
let v = a
.get(i + 1)
.ok_or("--metal-params needs a value".to_string())?;
parts.push((*v).to_string());
i += 2;
continue;
}
if let Some(key) = tok.strip_prefix("--")
&& Self::KEYS.iter().any(|(k, _)| *k == key)
{
let v = a
.get(i + 1)
.ok_or_else(|| format!("--{key} needs a value"))?;
parts.push(format!("{key}={v}"));
i += 2;
continue;
}
i += 1;
}
Ok((!parts.is_empty()).then(|| parts.join(",")))
}
pub fn from_args<S: AsRef<str>>(args: &[S]) -> Result<Option<Self>, String> {
match Self::spec_from_args(args)? {
None => Ok(None),
Some(spec) => Self::parse(&spec).map(Some),
}
}
pub fn from_env() -> Self {
match rlx_ir::env::var("RLX_METAL_PARAMS") {
None => Self::default(),
Some(spec) => Self::parse(&spec).unwrap_or_else(|why| {
eprintln!(
"rlx-metal: RLX_METAL_PARAMS {why} — using defaults.\n{}",
Self::help()
);
Self::default()
}),
}
}
pub fn resolve<S: AsRef<str>>(args: &[S]) -> Self {
match Self::from_args(args) {
Ok(Some(p)) => p,
Ok(None) => Self::from_env(),
Err(why) => {
eprintln!(
"rlx-metal: bad --metal-params {why} — falling back to config.\n{}",
Self::help()
);
Self::from_env()
}
}
}
pub fn spec(&self) -> String {
format!(
"stages={},sync={},precision={},tile={},encode={}",
self.stages,
match self.sync {
SyncScope::Threadgroup => "threadgroup",
SyncScope::Simdgroup => "simd",
},
match self.precision {
Precision::F32 => "f32",
Precision::F16Storage => "f16",
},
self.tile,
match self.encode {
Encode::PerDispatch => "dispatch",
Encode::Indirect => "icb",
}
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_default_is_the_configuration_that_measured_best() {
let d = AppleKernelParams::default();
assert_eq!(d.stages_value(), 1);
assert_eq!(d.sync_value(), SyncScope::Threadgroup);
assert_eq!(d.precision_value(), Precision::F32);
assert_eq!(d.threadgroup_bytes(), 2048);
}
#[test]
fn every_parameter_reaches_the_msl() {
let d = AppleKernelParams::default()
.stages(3)
.tile(8)
.precision(Precision::F16Storage)
.sync(SyncScope::Simdgroup)
.defines();
assert!(d.contains("#define RLX_STAGES 3"));
assert!(d.contains("#define RLX_TS 8"));
assert!(d.contains("#define RLX_STAGE_T half"));
assert!(d.contains("simdgroup_barrier"));
assert!(!d.contains("threadgroup_barrier(mem"));
}
#[test]
fn f16_staging_halves_the_threadgroup_footprint() {
let f32p = AppleKernelParams::default();
let f16p = AppleKernelParams::default().precision(Precision::F16Storage);
assert_eq!(f16p.threadgroup_bytes() * 2, f32p.threadgroup_bytes());
}
#[test]
fn a_deeper_rotation_is_predicted_slower_before_running() {
let p = AppleKernelParams::default().stages(3);
let v = p.predict(AppleGpuFamily::M4);
assert!(matches!(v, Verdict::Slower(_)), "got {v:?}");
assert!(!v.worth_measuring());
}
#[test]
fn f16_staging_is_not_penalised_by_occupancy() {
let p = AppleKernelParams::default().precision(Precision::F16Storage);
assert!(p.predict(AppleGpuFamily::M4).worth_measuring());
}
#[test]
fn for_shape_picks_the_measured_winner() {
for m in [1usize, 4, 16, 31] {
assert_eq!(AppleKernelParams::for_shape(m, 4096, 4096).tile_value(), 8);
}
for m in [32usize, 128, 512, 4096] {
assert_eq!(AppleKernelParams::for_shape(m, 4096, 4096).tile_value(), 32);
}
}
#[test]
fn for_shape_never_silently_costs_precision() {
for m in [1usize, 64, 4096] {
let p = AppleKernelParams::for_shape(m, 2048, 2048);
assert_eq!(p.precision_value(), Precision::F32);
assert_eq!(p.stages_value(), 1);
assert_eq!(p.sync_value(), SyncScope::Threadgroup);
}
}
#[test]
fn a_spec_round_trips() {
let p = AppleKernelParams::default()
.stages(2)
.sync(SyncScope::Simdgroup)
.precision(Precision::F16Storage)
.tile(32)
.encode(Encode::Indirect);
assert_eq!(AppleKernelParams::parse(&p.spec()), Ok(p));
}
#[test]
fn a_typo_is_an_error_rather_than_a_silent_default() {
for bad in [
"stages=two",
"sync=warp",
"precision=bf16",
"encode=graph",
"stagez=2",
"stages",
] {
assert!(
AppleKernelParams::parse(bad).is_err(),
"`{bad}` should not parse"
);
}
}
#[test]
fn an_empty_spec_is_the_default() {
assert_eq!(
AppleKernelParams::parse("").unwrap(),
AppleKernelParams::default()
);
}
#[test]
fn builder_config_and_cli_are_the_same_thing() {
let builder = AppleKernelParams::default()
.stages(2)
.sync(SyncScope::Simdgroup)
.precision(Precision::F16Storage);
let config = AppleKernelParams::parse("stages=2,sync=simd,precision=f16").unwrap();
let cli_compact =
AppleKernelParams::from_args(&["--metal-params", "stages=2,sync=simd,precision=f16"])
.unwrap()
.unwrap();
let cli_flags = AppleKernelParams::from_args(&[
"prog",
"--stages",
"2",
"--sync",
"simd",
"--precision",
"f16",
])
.unwrap()
.unwrap();
assert_eq!(builder, config, "builder != config");
assert_eq!(config, cli_compact, "config != --metal-params");
assert_eq!(cli_compact, cli_flags, "--metal-params != individual flags");
}
#[test]
fn arguments_that_mention_no_key_are_not_a_configuration() {
let none = AppleKernelParams::from_args(&["prog", "--verbose", "--out", "x.json"]).unwrap();
assert_eq!(none, None, "unrelated flags must not fabricate a config");
}
#[test]
fn a_cli_typo_is_an_error_too() {
assert!(AppleKernelParams::from_args(&["--sync", "warp"]).is_err());
assert!(AppleKernelParams::from_args(&["--stages"]).is_err());
}
#[test]
fn help_lists_exactly_the_keys_the_parser_accepts() {
let help = AppleKernelParams::help();
for (k, _) in AppleKernelParams::KEYS {
assert!(help.contains(k), "help omits `{k}`");
let probe = match *k {
"stages" | "tile" => format!("{k}=2"),
"sync" => "sync=simd".into(),
"precision" => "precision=f16".into(),
"encode" => "encode=icb".into(),
other => panic!("KEYS gained `{other}` with no probe"),
};
assert!(
AppleKernelParams::parse(&probe).is_ok(),
"`{probe}` rejected"
);
}
}
#[test]
fn the_whole_surface_is_one_environment_variable() {
let src = include_str!("apple_params.rs");
let code = src.split("#[cfg(test)]").next().expect("module body");
let count = code.matches("rlx_ir::env::var(").count();
assert_eq!(count, 1, "expected exactly one env read, found {count}");
assert!(code.contains("RLX_METAL_PARAMS"));
}
}