#![cfg(all(target_vendor = "apple", not(target_os = "watchos")))]
use std::collections::HashMap;
use std::sync::{Arc, OnceLock, RwLock};
use rlx_ir::Shape;
pub trait MetalKernel: Send + Sync + std::fmt::Debug {
fn name(&self) -> &str;
fn execute(
&self,
inputs: &[(&[u8], &Shape)],
output: (&mut [u8], &Shape),
attrs: &[u8],
) -> Result<(), String>;
}
pub struct MetalKernelRegistry {
kernels: RwLock<HashMap<String, Arc<dyn MetalKernel>>>,
}
impl MetalKernelRegistry {
pub fn new() -> Self {
Self {
kernels: RwLock::new(HashMap::new()),
}
}
pub fn register(&self, k: Arc<dyn MetalKernel>) {
let name = k.name().to_string();
let mut g = self.kernels.write().unwrap();
if g.contains_key(&name) {
eprintln!(
"rlx-metal: MetalKernel '{name}' was already registered — \
replacing the previous entry"
);
}
g.insert(name, k);
}
pub fn lookup(&self, name: &str) -> Option<Arc<dyn MetalKernel>> {
self.kernels.read().unwrap().get(name).cloned()
}
}
impl Default for MetalKernelRegistry {
fn default() -> Self {
Self::new()
}
}
pub fn global_metal_kernels() -> &'static MetalKernelRegistry {
static R: OnceLock<MetalKernelRegistry> = OnceLock::new();
R.get_or_init(MetalKernelRegistry::new)
}
pub fn register_metal_kernel(k: Arc<dyn MetalKernel>) {
global_metal_kernels().register(k);
}
fn ensure_builtins_registered() {
static ONCE: OnceLock<()> = OnceLock::new();
ONCE.get_or_init(|| {
crate::ms_deform_attn::register();
crate::collective::register();
rlx_cpu::onnx_ref::register_onnx_reference_kernels();
});
}
#[derive(Debug)]
struct OnnxHostDelegate {
name: String,
}
impl MetalKernel for OnnxHostDelegate {
fn name(&self) -> &str {
&self.name
}
fn execute(
&self,
inputs: &[(&[u8], &Shape)],
output: (&mut [u8], &Shape),
attrs: &[u8],
) -> Result<(), String> {
if std::env::var("RLX_DBG_CUSTOM").is_ok() {
eprintln!(
"[custom] {} in={:?} out={:?}",
self.name,
inputs
.iter()
.map(|(_, s)| (s.dtype(), s.dims().to_vec()))
.collect::<Vec<_>>(),
(output.1.dtype(), output.1.dims().to_vec()),
);
}
rlx_cpu::op_registry::run_custom_op_host(&self.name, inputs, output, attrs)
}
}
pub fn lookup_metal_kernel(name: &str) -> Option<Arc<dyn MetalKernel>> {
ensure_builtins_registered();
if let Some(k) = global_metal_kernels().lookup(name) {
return Some(k);
}
if rlx_cpu::op_registry::lookup_cpu_kernel(name).is_some() {
return Some(Arc::new(OnnxHostDelegate {
name: name.to_string(),
}));
}
None
}
pub struct MetalGpuDispatch<'a> {
pub encoder: &'a metal::ComputeCommandEncoderRef,
pub arena: &'a metal::BufferRef,
pub inputs: &'a [(usize, u32, Shape)],
pub output: &'a (usize, u32, Shape),
pub attrs: &'a [u8],
}
pub trait MetalGpuKernel: Send + Sync + std::fmt::Debug {
fn name(&self) -> &str;
fn encode(&self, d: &MetalGpuDispatch) -> Result<(), String>;
}
struct MetalGpuKernelRegistry {
kernels: RwLock<HashMap<String, Arc<dyn MetalGpuKernel>>>,
}
impl MetalGpuKernelRegistry {
fn new() -> Self {
Self {
kernels: RwLock::new(HashMap::new()),
}
}
fn register(&self, k: Arc<dyn MetalGpuKernel>) {
let name = k.name().to_string();
let mut g = self.kernels.write().unwrap();
if g.contains_key(&name) {
eprintln!(
"rlx-metal: MetalGpuKernel '{name}' was already registered — \
replacing the previous entry"
);
}
g.insert(name, k);
}
fn lookup(&self, name: &str) -> Option<Arc<dyn MetalGpuKernel>> {
self.kernels.read().unwrap().get(name).cloned()
}
}
fn global_metal_gpu_kernels() -> &'static MetalGpuKernelRegistry {
static R: OnceLock<MetalGpuKernelRegistry> = OnceLock::new();
R.get_or_init(MetalGpuKernelRegistry::new)
}
pub fn register_metal_gpu_kernel(k: Arc<dyn MetalGpuKernel>) {
global_metal_gpu_kernels().register(k);
}
pub fn lookup_metal_gpu_kernel(name: &str) -> Option<Arc<dyn MetalGpuKernel>> {
global_metal_gpu_kernels().lookup(name)
}
#[cfg(test)]
mod tests {
use super::*;
use rlx_ir::DType;
#[derive(Debug)]
struct StubKernel;
impl MetalKernel for StubKernel {
fn name(&self) -> &str {
"stub.metal"
}
fn execute(
&self,
_inputs: &[(&[u8], &Shape)],
_output: (&mut [u8], &Shape),
_attrs: &[u8],
) -> Result<(), String> {
Ok(())
}
}
#[test]
fn register_and_lookup_round_trips() {
let reg = MetalKernelRegistry::new();
reg.register(Arc::new(StubKernel));
let k = reg
.lookup("stub.metal")
.expect("registered kernel must be findable");
assert_eq!(k.name(), "stub.metal");
}
#[test]
fn execute_signature_compiles_and_runs() {
let k: Arc<dyn MetalKernel> = Arc::new(StubKernel);
let in_shape = Shape::new(&[4], DType::F32);
let out_shape = Shape::new(&[4], DType::F32);
let in_bytes = vec![0u8; 16];
let mut out_bytes = vec![0u8; 16];
k.execute(&[(&in_bytes, &in_shape)], (&mut out_bytes, &out_shape), &[])
.expect("stub kernel must succeed");
}
}