use crate::device::metal_device;
use std::sync::OnceLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AppleGpuFamily {
Unknown,
M1, M1Pro, M2, M3, M4, }
impl AppleGpuFamily {
fn from_name(name: &str) -> Self {
let lower = name.to_lowercase();
if lower.contains("m4") {
Self::M4
} else if lower.contains("m3") {
Self::M3
} else if lower.contains("m2") {
Self::M2
} else if lower.contains("m1 pro") || lower.contains("m1 max") || lower.contains("m1 ultra")
{
Self::M1Pro
} else if lower.contains("m1") {
Self::M1
} else {
Self::Unknown
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SgemmVariant {
Mps,
Simd4x4,
Simd,
SimdPadded,
Tiled,
Naive,
}
pub struct MetalHwModel {
pub gpu_family: AppleGpuFamily,
pub gpu_name: String,
pub sgemm_simd_flops: f64,
pub sgemm_simd_4x4_flops: f64,
pub sgemm_padded_flops: f64,
pub sgemm_tiled_flops: f64,
pub dispatch_overhead_ns: f64,
pub roundtrip_overhead_ns: f64,
pub threadgroup_mem_bytes: usize,
pub unified_memory: bool,
pub mps_threshold_flop: u64,
}
impl MetalHwModel {
fn detect() -> Self {
let dev = metal_device();
let (name, unified) = match dev {
Some(d) => (d.name.clone(), d.has_unified_memory),
None => ("unknown".to_string(), false),
};
let family = AppleGpuFamily::from_name(&name);
let (simd_flops, padded_flops, tiled_flops) = match family {
AppleGpuFamily::M4 => (600e9, 350e9, 100e9),
AppleGpuFamily::M3 => (500e9, 300e9, 90e9),
AppleGpuFamily::M2 => (400e9, 240e9, 75e9),
AppleGpuFamily::M1Pro => (350e9, 200e9, 65e9),
AppleGpuFamily::M1 => (200e9, 110e9, 40e9),
AppleGpuFamily::Unknown => (300e9, 180e9, 60e9),
};
let mut simd_4x4_flops = simd_flops * 3.5;
let mut simd_flops = simd_flops;
let mut padded_flops = padded_flops;
let mut tiled_flops = tiled_flops;
let mut roundtrip_ns = 800_000.0_f64;
let dev_id = dev.map(|d| d.registry_id).unwrap_or(0);
if let Some(cal) = crate::calibrate::Calibration::load(dev_id) {
simd_4x4_flops = cal.sgemm_simd_4x4_flops;
simd_flops = cal.sgemm_simd_flops;
padded_flops = cal.sgemm_padded_flops;
tiled_flops = cal.sgemm_tiled_flops;
roundtrip_ns = cal.roundtrip_overhead_ns;
}
let mps_threshold_flop = rlx_ir::env::var("RLX_MPS_THRESHOLD_FLOP")
.and_then(|s| s.parse::<u64>().ok())
.unwrap_or(16_000_000);
Self {
gpu_family: family,
gpu_name: name,
sgemm_simd_flops: simd_flops,
sgemm_simd_4x4_flops: simd_4x4_flops,
sgemm_padded_flops: padded_flops,
sgemm_tiled_flops: tiled_flops,
dispatch_overhead_ns: 8_000.0,
roundtrip_overhead_ns: roundtrip_ns,
threadgroup_mem_bytes: 32 * 1024,
unified_memory: unified,
mps_threshold_flop,
}
}
pub fn pick_sgemm(&self, m: usize, k: usize, n: usize) -> SgemmVariant {
if let Some(forced) = sgemm_variant_override() {
return forced;
}
let aligned_8 = m.is_multiple_of(8) && k.is_multiple_of(8) && n.is_multiple_of(8);
let mps_enabled = rlx_ir::env::flag("RLX_METAL_SGEMM_MPS");
let mps_disabled = rlx_ir::env::var("RLX_DISABLE_MPS")
.map(|v| v == "1")
.unwrap_or(false);
let flop = (m as u64) * (k as u64) * (n as u64);
if mps_enabled
&& !mps_disabled
&& crate::mps_blas::mps_supports_matmul()
&& flop >= self.mps_threshold_flop
{
return SgemmVariant::Mps;
}
if k.is_multiple_of(32) && n.is_multiple_of(32) && m.is_multiple_of(32) {
SgemmVariant::Simd4x4
} else if m < 32 {
SgemmVariant::Naive
} else if aligned_8 && m >= 8 && n >= 8 {
SgemmVariant::Simd
} else if k.is_multiple_of(8) && n >= 8 && m >= 1 {
SgemmVariant::SimdPadded
} else if m >= 16 && n >= 16 {
SgemmVariant::Tiled
} else {
SgemmVariant::Naive
}
}
pub fn sgemm_cost_ns(&self, m: usize, k: usize, n: usize) -> f64 {
let flops = 2.0 * m as f64 * k as f64 * n as f64;
let throughput = match self.pick_sgemm(m, k, n) {
SgemmVariant::Mps => self.sgemm_simd_4x4_flops * 2.0,
SgemmVariant::Simd4x4 => self.sgemm_simd_4x4_flops,
SgemmVariant::Simd => self.sgemm_simd_flops,
SgemmVariant::SimdPadded => self.sgemm_padded_flops,
SgemmVariant::Tiled => self.sgemm_tiled_flops,
SgemmVariant::Naive => self.sgemm_tiled_flops * 0.3,
};
let compute_ns = flops / throughput;
compute_ns + self.dispatch_overhead_ns
}
pub fn prefer_fused_matmul_bias(&self, _m: usize, _k: usize, _n: usize) -> bool {
true
}
pub fn fits_threadgroup_mem(
&self,
batch: usize,
seq: usize,
hidden: usize,
intermediate: usize,
) -> bool {
let m = batch * seq;
let bytes = m * (hidden + 3 * hidden + hidden + intermediate) * 4;
bytes <= self.threadgroup_mem_bytes
}
pub fn estimate_transformer_forward_ns(
&self,
batch: usize,
seq: usize,
hidden: usize,
intermediate: usize,
num_heads: usize,
num_layers: usize,
) -> f64 {
let m = batch * seq;
let _ = num_heads;
let qkv = self.sgemm_cost_ns(m, hidden, 3 * hidden);
let out = self.sgemm_cost_ns(m, hidden, hidden);
let fc1 = self.sgemm_cost_ns(m, hidden, intermediate);
let fc2 = self.sgemm_cost_ns(m, intermediate, hidden);
let attn = (seq * seq * hidden) as f64 / self.sgemm_simd_flops + self.dispatch_overhead_ns;
let elem = 4.0 * self.dispatch_overhead_ns;
let per_layer = qkv + out + fc1 + fc2 + attn + elem;
per_layer * num_layers as f64 + self.roundtrip_overhead_ns
}
}
pub(crate) fn sgemm_variant_override() -> Option<SgemmVariant> {
if let Some(raw) = rlx_ir::env::var("RLX_METAL_SGEMM_VARIANT") {
match raw.to_ascii_lowercase().as_str() {
"mps" => return Some(SgemmVariant::Mps),
"simd4x4" | "simd_4x4" | "4x4" => return Some(SgemmVariant::Simd4x4),
"simd" | "simd8" | "simd_8" => return Some(SgemmVariant::Simd),
"padded" | "simd_padded" | "simdpadded" => return Some(SgemmVariant::SimdPadded),
"tiled" => return Some(SgemmVariant::Tiled),
"naive" => return Some(SgemmVariant::Naive),
_ => {}
}
}
if let Some(raw) = rlx_ir::env::var("RLX_METAL_PRECISE") {
match raw.to_ascii_lowercase().as_str() {
"1" | "true" | "yes" | "on" => return Some(SgemmVariant::Naive),
_ => {}
}
}
None
}
pub fn hw_model() -> &'static MetalHwModel {
static MODEL: OnceLock<MetalHwModel> = OnceLock::new();
MODEL.get_or_init(MetalHwModel::detect)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn detects_some_gpu() {
let hw = hw_model();
assert!(!hw.gpu_name.is_empty());
assert!(hw.sgemm_simd_flops > 0.0);
}
#[test]
fn picks_simd_for_aligned() {
rlx_ir::env::set("RLX_DISABLE_MPS", "1");
rlx_ir::env::unset("RLX_METAL_SGEMM_VARIANT");
let hw = MetalHwModel::detect();
assert_eq!(hw.pick_sgemm(64, 768, 2304), SgemmVariant::Simd4x4);
assert_eq!(hw.pick_sgemm(750, 768, 2304), SgemmVariant::SimdPadded);
assert_eq!(hw.pick_sgemm(8, 16, 16), SgemmVariant::Naive);
assert_eq!(hw.pick_sgemm(6, 768, 2304), SgemmVariant::Naive);
assert_eq!(hw.pick_sgemm(6, 768, 2300), SgemmVariant::Naive);
assert_eq!(hw.pick_sgemm(6, 7, 7), SgemmVariant::Naive);
rlx_ir::env::unset("RLX_DISABLE_MPS");
}
}