use metal::{
Buffer, ComputeCommandEncoderRef, ComputePipelineDescriptor, ComputePipelineState, Device,
IndirectCommandBuffer, IndirectCommandBufferDescriptor, MTLIndirectCommandType,
MTLResourceOptions, MTLSize, NSRange,
};
use objc::{msg_send, sel, sel_impl};
use std::sync::OnceLock;
use crate::thunk::Thunk;
pub struct IcbKernels {
pub bias_add: ComputePipelineState,
pub gelu_inplace: ComputePipelineState,
pub silu_inplace: ComputePipelineState,
pub elem_add: ComputePipelineState,
pub elem_mul: ComputePipelineState,
pub copy_f32: ComputePipelineState,
pub layer_norm: ComputePipelineState,
pub fused_residual_ln: ComputePipelineState,
pub narrow_lastax: ComputePipelineState,
pub rope: ComputePipelineState,
}
unsafe impl Send for IcbKernels {}
unsafe impl Sync for IcbKernels {}
impl IcbKernels {
fn new(dev: &Device, library: &metal::LibraryRef) -> Self {
let trace = rlx_ir::env::flag("RLX_ICB_TRACE");
let make = |name: &str| -> ComputePipelineState {
let f = library.get_function(name, None).expect(name);
let desc = ComputePipelineDescriptor::new();
desc.set_support_indirect_command_buffers(true);
desc.set_compute_function(Some(&f));
let flag = desc.support_indirect_command_buffers();
if trace {
eprintln!("[icb-kernels] {name}: support_icb={flag}");
}
dev.new_compute_pipeline_state(&desc).expect(name)
};
Self {
bias_add: make("bias_add"),
gelu_inplace: make("gelu_inplace"),
silu_inplace: make("silu_inplace"),
elem_add: make("elem_add"),
elem_mul: make("elem_mul"),
copy_f32: make("copy_f32"),
layer_norm: make("layer_norm"),
fused_residual_ln: make("fused_residual_ln"),
narrow_lastax: make("narrow_lastax"),
rope: make("rope"),
}
}
}
pub fn icb_kernels() -> &'static IcbKernels {
static K: OnceLock<IcbKernels> = OnceLock::new();
K.get_or_init(|| {
use crate::device::metal_device;
let dev = metal_device().expect("Metal device required");
let opts = metal::CompileOptions::new();
let library = dev
.device
.new_library_with_source(crate::kernels::RLX_KERNELS_MSL, &opts)
.expect("MSL compilation for ICB kernels failed");
IcbKernels::new(&dev.device, &library)
})
}
pub struct IcbSegment {
pub icb: IndirectCommandBuffer,
pub command_count: u64,
pub constants: Buffer,
}
unsafe impl Send for IcbSegment {}
unsafe impl Sync for IcbSegment {}
const RESOURCE_USAGE_READ: u64 = 1 << 0;
const RESOURCE_USAGE_WRITE: u64 = 1 << 1;
impl IcbSegment {
pub fn execute_on(&self, enc: &ComputeCommandEncoderRef, arena: &Buffer) {
if self.command_count == 0 {
return;
}
unsafe {
let arena_ref: &metal::BufferRef = arena;
let const_ref: &metal::BufferRef = &self.constants;
let _: () = msg_send![enc, useResource: arena_ref
usage: (RESOURCE_USAGE_READ | RESOURCE_USAGE_WRITE)];
let _: () = msg_send![enc, useResource: const_ref
usage: RESOURCE_USAGE_READ];
let range = NSRange {
location: 0,
length: self.command_count,
};
let _: () = msg_send![enc,
executeCommandsInBuffer: &*self.icb
withRange: range];
}
}
}
fn thunk_kind(t: &Thunk) -> &'static str {
match t {
Thunk::Nop => "Nop",
Thunk::Cast { .. } => "Cast",
Thunk::Sgemm { .. } => "Sgemm",
Thunk::FusedMmBiasAct { .. } => "FusedMmBiasAct",
Thunk::BiasAdd { .. } => "BiasAdd",
Thunk::ActivationInPlace { .. } => "ActivationInPlace",
Thunk::BinaryFull { .. } => "BinaryFull",
Thunk::Copy { .. } => "Copy",
Thunk::LayerNorm { .. } => "LayerNorm",
Thunk::RmsNorm { .. } => "RmsNorm",
Thunk::Softmax { .. } => "Softmax",
Thunk::SoftmaxCrossEntropyDense { .. } => "SoftmaxCrossEntropy",
Thunk::Reduce { .. } => "Reduce",
Thunk::Gather { .. } => "Gather",
Thunk::Narrow { .. } => "Narrow",
Thunk::Transpose { .. } => "Transpose",
Thunk::Concat { .. } => "Concat",
Thunk::Attention { .. } => "Attention",
Thunk::Rope { .. } => "Rope",
Thunk::FusedResidualLN { .. } => "FusedResidualLN",
_ => "Other",
}
}
fn is_icb_compatible(t: &Thunk) -> bool {
use crate::thunk::HalfFlag;
let f32_dt = |dt: &HalfFlag| matches!(dt, HalfFlag::F32);
match t {
Thunk::BiasAdd { dt, .. } => f32_dt(dt),
Thunk::ActivationInPlace { dt, .. } => f32_dt(dt),
Thunk::BinaryFull { dt, .. } => f32_dt(dt),
Thunk::Copy { dt, .. } => f32_dt(dt),
Thunk::LayerNorm { dt, .. } => f32_dt(dt),
Thunk::FusedResidualLN { dt, .. } => f32_dt(dt),
Thunk::Narrow { dt, .. } => f32_dt(dt),
Thunk::Rope { dt, .. } => f32_dt(dt),
_ => false,
}
}
const CONSTANTS_BYTES_PER_CMD: usize = 64;
pub struct IcbRange {
pub start: usize,
pub end: usize,
pub segment: IcbSegment,
}
const MIN_ICB_RUN: usize = 2;
pub fn compile_segments(thunks: &[Thunk], arena: &Buffer, dev: &Device) -> Vec<IcbRange> {
let trace = rlx_ir::env::flag("RLX_ICB_TRACE");
if trace {
let mut hist: std::collections::HashMap<&'static str, usize> = Default::default();
for t in thunks {
*hist.entry(thunk_kind(t)).or_default() += 1;
}
let mut entries: Vec<_> = hist.into_iter().collect();
entries.sort_by_key(|(_, c)| std::cmp::Reverse(*c));
eprintln!("[icb] thunk histogram: {:?}", entries);
let order: String = thunks
.iter()
.filter(|t| !matches!(t, Thunk::Nop))
.map(|t| if is_icb_compatible(t) { "+" } else { "-" })
.collect();
eprintln!("[icb] schedule (skipping Nops, +=ICB / -=other): {order}");
}
let mut out = Vec::new();
let mut i = 0;
while i < thunks.len() {
while i < thunks.len() && !is_icb_compatible(&thunks[i]) && !matches!(thunks[i], Thunk::Nop)
{
i += 1;
}
let start = i;
let mut run: Vec<&Thunk> = Vec::new();
while i < thunks.len() && (is_icb_compatible(&thunks[i]) || matches!(thunks[i], Thunk::Nop))
{
if !matches!(thunks[i], Thunk::Nop) {
run.push(&thunks[i]);
}
i += 1;
}
let end = i;
if run.len() >= MIN_ICB_RUN
&& let Some(seg) = build_segment(&run, arena, dev)
{
if trace {
eprintln!(
"[icb] segment thunks {}..{} ({} cmds)",
start, end, seg.command_count
);
}
out.push(IcbRange {
start,
end,
segment: seg,
});
}
}
if trace {
eprintln!(
"[icb] compile_segments: {} segments over {} thunks",
out.len(),
thunks.len()
);
}
out
}
pub fn try_compile(thunks: &[Thunk], arena: &Buffer, dev: &Device) -> Option<IcbSegment> {
let compute_thunks: Vec<&Thunk> = thunks.iter().filter(|t| !matches!(t, Thunk::Nop)).collect();
if compute_thunks.is_empty() {
return None;
}
if !compute_thunks.iter().all(|t| is_icb_compatible(t)) {
return None;
}
build_segment(&compute_thunks, arena, dev)
}
fn build_segment(icb_thunks: &[&Thunk], arena: &Buffer, dev: &Device) -> Option<IcbSegment> {
let trace = rlx_ir::env::flag("RLX_ICB_TRACE");
if icb_thunks.is_empty() {
return None;
}
let n = icb_thunks.len() as u64;
if trace {
eprintln!("[icb] build_segment n={n}");
}
let desc = IndirectCommandBufferDescriptor::new();
if trace {
eprintln!("[icb] descriptor created");
}
desc.set_command_types(
MTLIndirectCommandType::ConcurrentDispatch
| MTLIndirectCommandType::ConcurrentDispatchThreads,
);
desc.set_inherit_buffers(false);
desc.set_inherit_pipeline_state(false);
desc.set_max_kernel_buffer_bind_count(8);
if trace {
eprintln!("[icb] descriptor configured");
}
let icb =
dev.new_indirect_command_buffer_with_descriptor(&desc, n, MTLResourceOptions::empty());
if trace {
eprintln!("[icb] icb allocated");
}
let constants = dev.new_buffer(
(n as usize * CONSTANTS_BYTES_PER_CMD) as u64,
MTLResourceOptions::StorageModeShared,
);
if trace {
eprintln!("[icb] constants buffer allocated");
}
let kk = icb_kernels();
for (idx, t) in icb_thunks.iter().enumerate() {
if trace {
eprintln!("[icb] encode cmd {idx}");
}
let cmd = icb.indirect_compute_command_at_index(idx as u64);
let cb_off = idx * CONSTANTS_BYTES_PER_CMD;
encode_thunk_into_icb(cmd, t, arena, &constants, cb_off, kk);
cmd.set_barrier();
}
if trace {
eprintln!("[icb] build_segment done");
}
Some(IcbSegment {
icb,
command_count: n,
constants,
})
}
fn encode_thunk_into_icb(
cmd: &metal::IndirectComputeCommandRef,
thunk: &Thunk,
arena: &Buffer,
constants_buf: &Buffer,
cb_off: usize,
k: &IcbKernels,
) {
use rlx_ir::op::{Activation, BinaryOp};
let write_u32s = |vals: &[u32]| unsafe {
let p = constants_buf.contents() as *mut u8;
let dst = p.add(cb_off) as *mut u32;
for (i, &v) in vals.iter().enumerate() {
*dst.add(i) = v;
}
};
let cb_arg = |slot_idx: usize| -> u64 { (cb_off + slot_idx * 4) as u64 };
match thunk {
Thunk::BiasAdd {
src,
bias,
dst,
m,
n,
..
} => {
if *src != *dst {
return;
}
write_u32s(&[*m, *n]);
cmd.set_compute_pipeline_state(&k.bias_add);
cmd.set_kernel_buffer(0, Some(&**arena), *dst as u64);
cmd.set_kernel_buffer(1, Some(&**arena), *bias as u64);
cmd.set_kernel_buffer(2, Some(&**constants_buf), cb_arg(0));
cmd.set_kernel_buffer(3, Some(&**constants_buf), cb_arg(1));
let grid = MTLSize {
width: *n as u64,
height: *m as u64,
depth: 1,
};
let tg = MTLSize {
width: 16.min(*n as u64),
height: 16.min(*m as u64),
depth: 1,
};
cmd.concurrent_dispatch_threads(grid, tg);
}
Thunk::ActivationInPlace { data, len, act, .. } => {
let trace = rlx_ir::env::flag("RLX_ICB_TRACE");
if trace {
eprintln!(" [act] write u32");
}
write_u32s(&[*len]);
let pipeline = match act {
Activation::Gelu => &k.gelu_inplace,
Activation::Silu => &k.silu_inplace,
_ => return,
};
if trace {
eprintln!(" [act] set pipeline");
}
cmd.set_compute_pipeline_state(pipeline);
if trace {
eprintln!(" [act] set buf 0");
}
let arena_ref: &metal::BufferRef = arena;
cmd.set_kernel_buffer(0, Some(arena_ref), *data as u64);
if trace {
eprintln!(" [act] set buf 1");
}
let cb_ref: &metal::BufferRef = constants_buf;
cmd.set_kernel_buffer(1, Some(cb_ref), cb_arg(0));
if trace {
eprintln!(" [act] dispatch");
}
let tg_w = pipeline.thread_execution_width().min(*len as u64);
cmd.concurrent_dispatch_threads(
MTLSize {
width: *len as u64,
height: 1,
depth: 1,
},
MTLSize {
width: tg_w,
height: 1,
depth: 1,
},
);
if trace {
eprintln!(" [act] done");
}
}
Thunk::BinaryFull {
lhs,
rhs,
dst,
len,
op,
..
} => {
write_u32s(&[*len]);
let pipeline = match op {
BinaryOp::Add => &k.elem_add,
BinaryOp::Mul => &k.elem_mul,
_ => return,
};
cmd.set_compute_pipeline_state(pipeline);
cmd.set_kernel_buffer(0, Some(&**arena), *lhs as u64);
cmd.set_kernel_buffer(1, Some(&**arena), *rhs as u64);
cmd.set_kernel_buffer(2, Some(&**arena), *dst as u64);
cmd.set_kernel_buffer(3, Some(&**constants_buf), cb_arg(0));
let tg_w = pipeline.thread_execution_width().min(*len as u64);
cmd.concurrent_dispatch_threads(
MTLSize {
width: *len as u64,
height: 1,
depth: 1,
},
MTLSize {
width: tg_w,
height: 1,
depth: 1,
},
);
}
Thunk::Copy { src, dst, len, .. } => {
write_u32s(&[*len]);
cmd.set_compute_pipeline_state(&k.copy_f32);
cmd.set_kernel_buffer(0, Some(&**arena), *src as u64);
cmd.set_kernel_buffer(1, Some(&**arena), *dst as u64);
cmd.set_kernel_buffer(2, Some(&**constants_buf), cb_arg(0));
let tg_w = k.copy_f32.thread_execution_width().min(*len as u64);
cmd.concurrent_dispatch_threads(
MTLSize {
width: *len as u64,
height: 1,
depth: 1,
},
MTLSize {
width: tg_w,
height: 1,
depth: 1,
},
);
}
Thunk::LayerNorm {
src,
g,
b,
dst,
rows,
h,
eps,
..
} => {
write_u32s(&[*h]);
unsafe {
let p = constants_buf.contents() as *mut u8;
let f = p.add(cb_off + 4) as *mut f32;
*f = *eps;
}
cmd.set_compute_pipeline_state(&k.layer_norm);
cmd.set_kernel_buffer(0, Some(&**arena), *src as u64);
cmd.set_kernel_buffer(1, Some(&**arena), *g as u64);
cmd.set_kernel_buffer(2, Some(&**arena), *b as u64);
cmd.set_kernel_buffer(3, Some(&**arena), *dst as u64);
cmd.set_kernel_buffer(4, Some(&**constants_buf), cb_arg(0));
cmd.set_kernel_buffer(5, Some(&**constants_buf), cb_arg(1));
let tg_w = 256u64.min(*h as u64);
cmd.concurrent_dispatch_threads(
MTLSize {
width: tg_w,
height: *rows as u64,
depth: 1,
},
MTLSize {
width: tg_w,
height: 1,
depth: 1,
},
);
}
Thunk::FusedResidualLN {
x,
res,
g,
b,
out,
rows,
h,
eps,
..
} => {
write_u32s(&[*h]);
unsafe {
let p = constants_buf.contents() as *mut u8;
let f = p.add(cb_off + 4) as *mut f32;
*f = *eps;
}
cmd.set_compute_pipeline_state(&k.fused_residual_ln);
cmd.set_kernel_buffer(0, Some(&**arena), *x as u64);
cmd.set_kernel_buffer(1, Some(&**arena), *res as u64);
cmd.set_kernel_buffer(2, Some(&**arena), *g as u64);
cmd.set_kernel_buffer(3, Some(&**arena), *b as u64);
cmd.set_kernel_buffer(4, Some(&**arena), *out as u64);
cmd.set_kernel_buffer(5, Some(&**constants_buf), cb_arg(0));
cmd.set_kernel_buffer(6, Some(&**constants_buf), cb_arg(1));
let tg_w = 256u64.min(*h as u64);
cmd.concurrent_dispatch_threadgroups(
MTLSize {
width: 1,
height: *rows as u64,
depth: 1,
},
MTLSize {
width: tg_w,
height: 1,
depth: 1,
},
);
}
Thunk::Narrow {
src,
dst,
outer,
src_axis,
start,
len,
..
} => {
write_u32s(&[*outer, *src_axis, *start, *len]);
cmd.set_compute_pipeline_state(&k.narrow_lastax);
cmd.set_kernel_buffer(0, Some(&**arena), *src as u64);
cmd.set_kernel_buffer(1, Some(&**arena), *dst as u64);
cmd.set_kernel_buffer(2, Some(&**constants_buf), cb_arg(0));
cmd.set_kernel_buffer(3, Some(&**constants_buf), cb_arg(1));
cmd.set_kernel_buffer(4, Some(&**constants_buf), cb_arg(2));
cmd.set_kernel_buffer(5, Some(&**constants_buf), cb_arg(3));
let grid = MTLSize {
width: *len as u64,
height: *outer as u64,
depth: 1,
};
let tg = MTLSize {
width: 16.min(*len as u64),
height: 16.min(*outer as u64),
depth: 1,
};
cmd.concurrent_dispatch_threads(grid, tg);
}
Thunk::Rope {
src,
cos,
sin,
dst,
batch,
seq,
hidden,
head_dim,
n_rot,
src_row_stride,
cos_per_token,
interleaved,
..
} => {
write_u32s(&[
*batch,
*seq,
*hidden,
*head_dim,
*src_row_stride,
*seq,
*n_rot,
*cos_per_token as u32,
*interleaved as u32,
]);
cmd.set_compute_pipeline_state(&k.rope);
cmd.set_kernel_buffer(0, Some(&**arena), *src as u64);
cmd.set_kernel_buffer(1, Some(&**arena), *cos as u64);
cmd.set_kernel_buffer(2, Some(&**arena), *sin as u64);
cmd.set_kernel_buffer(3, Some(&**arena), *dst as u64);
cmd.set_kernel_buffer(4, Some(&**constants_buf), cb_arg(0));
cmd.set_kernel_buffer(5, Some(&**constants_buf), cb_arg(1));
cmd.set_kernel_buffer(6, Some(&**constants_buf), cb_arg(2));
cmd.set_kernel_buffer(7, Some(&**constants_buf), cb_arg(3));
cmd.set_kernel_buffer(8, Some(&**constants_buf), cb_arg(4));
cmd.set_kernel_buffer(9, Some(&**constants_buf), cb_arg(5));
cmd.set_kernel_buffer(10, Some(&**constants_buf), cb_arg(6));
cmd.set_kernel_buffer(11, Some(&**constants_buf), cb_arg(7));
cmd.set_kernel_buffer(12, Some(&**constants_buf), cb_arg(8));
let nh = *hidden / *head_dim;
let grid = MTLSize {
width: *head_dim as u64,
height: nh as u64,
depth: (*batch * *seq) as u64,
};
let tg = MTLSize {
width: (*head_dim).min(16) as u64,
height: nh.min(8) as u64,
depth: 1,
};
cmd.concurrent_dispatch_threads(grid, tg);
}
_ => {} }
}