const ARENA_LARGE_OFF: usize = 1usize << 32;
pub(crate) const WEIGHT_BUF_TAG: usize = 1usize << 63;
#[inline]
pub(crate) fn tag_weight_off(off: usize) -> usize {
debug_assert_eq!(off & WEIGHT_BUF_TAG, 0, "weight offset already tagged");
off | WEIGHT_BUF_TAG
}
#[inline]
pub(crate) fn is_weight_off(off: usize) -> bool {
off & WEIGHT_BUF_TAG != 0
}
#[inline]
pub(crate) fn raw_off(off: usize) -> usize {
off & !WEIGHT_BUF_TAG
}
#[inline]
fn arena_off_large(off: usize) -> bool {
!is_weight_off(off) && raw_off(off) >= ARENA_LARGE_OFF
}
#[inline]
fn metal_host_fallback_enabled() -> bool {
matches!(
std::env::var("RLX_METAL_HOST_SLICE").as_deref(),
Ok("1") | Ok("true") | Ok("on")
) || rlx_ir::env::flag("RLX_METAL_HOST_FALLBACK")
}
fn broadcast_strides(in_dims: &[usize], out_dims: &[usize]) -> Vec<u32> {
let r_out = out_dims.len();
let r_in = in_dims.len();
debug_assert!(r_in <= r_out, "broadcast in rank {r_in} > out rank {r_out}");
let pad = r_out - r_in;
let mut strides = vec![0u32; r_out];
let mut acc: usize = 1;
for d in (0..r_out).rev() {
let in_size = if d < pad { 1 } else { in_dims[d - pad] };
if in_size == 1 {
strides[d] = 0;
} else {
debug_assert_eq!(
in_size, out_dims[d],
"broadcast: dim {in_size} vs out {} at {d}",
out_dims[d]
);
strides[d] = acc as u32;
acc *= in_size;
}
}
strides
}
fn ada_lead_pack(x_dims: &[usize], mod_dims_in: &[usize]) -> [u32; 17] {
rlx_ir::ada_modulation_lead_pack(x_dims, mod_dims_in)
}
fn ada_mod_launch(x_dims: &[usize], mod_dims: &[usize]) -> (u32, u32) {
debug_assert!(!x_dims.is_empty() && !mod_dims.is_empty());
let xr = x_dims.len() - 1;
let mr = mod_dims.len() - 1;
let mut seq = 1u32;
let mut mods = 1u32;
for i in 0..xr {
let xd = x_dims[i] as u32;
let md = if i + mr >= xr {
mod_dims[i - (xr - mr)] as u32
} else {
1
};
if md == 1 && xd > 1 {
seq = seq.saturating_mul(xd);
} else {
mods = mods.saturating_mul(xd.max(1));
}
}
(mods.max(1), seq.max(1))
}
fn trailing_broadcast(lhs: &Shape, rhs: &Shape) -> bool {
if rhs.rank() > lhs.rank() {
return false;
}
let off = lhs.rank() - rhs.rank();
for i in 0..rhs.rank() {
let r = rhs.dim(i).unwrap_static();
let l = lhs.dim(off + i).unwrap_static();
if r != l {
return false;
}
}
true
}
use crate::op_registry::{MetalGpuKernel, MetalKernel};
use rlx_ir::op::{Activation, BinaryOp, CmpOp};
use rlx_ir::{DType, Shape};
use std::sync::Arc;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum HalfFlag {
F32,
F16,
}
impl From<DType> for HalfFlag {
fn from(d: DType) -> Self {
match d {
DType::F16 => HalfFlag::F16,
_ => HalfFlag::F32,
}
}
}
#[derive(Clone, Debug)]
pub enum Thunk {
Nop,
Cast {
src: usize,
dst: usize,
len: u32,
src_dt: HalfFlag,
dst_dt: HalfFlag,
},
CastHost {
src: usize,
dst: usize,
len: u32,
src_dt: DType,
dst_dt: DType,
},
CastTruncF32 {
src: usize,
dst: usize,
len: u32,
},
Sgemm {
a: usize,
b: usize,
c: usize,
m: u32,
k: u32,
n: u32,
dt: HalfFlag,
b_f16: bool,
},
BatchedSgemm {
a: usize,
b: usize,
c: usize,
batch: u32,
m: u32,
k: u32,
n: u32,
dt: HalfFlag,
a_bcast: bool,
b_bcast: bool,
},
FusedMmBiasAct {
a: usize,
w: usize,
bias: usize,
c: usize,
m: u32,
k: u32,
n: u32,
act: Option<Activation>,
dt: HalfFlag,
},
ActivationInPlace {
data: usize,
len: u32,
act: Activation,
dt: HalfFlag,
},
ActivationOut {
src: usize,
dst: usize,
len: u32,
act: Activation,
dt: HalfFlag,
},
GeluApproxOut {
src: usize,
dst: usize,
len: u32,
},
GeluApproxHost {
src: usize,
dst: usize,
len: u32,
},
FusedBinaryActivation {
lhs: usize,
rhs: usize,
dst: usize,
len: u32,
op: BinaryOp,
act: Activation,
dt: HalfFlag,
},
FusedTernaryActivation {
lhs: usize,
rhs0: usize,
rhs1: usize,
dst: usize,
len: u32,
op0: BinaryOp,
op1: BinaryOp,
act: Activation,
dt: HalfFlag,
},
LayerNorm {
src: usize,
g: usize,
b: usize,
dst: usize,
rows: u32,
h: u32,
eps: f32,
dt: HalfFlag,
},
GroupNorm {
src: usize,
g: usize,
b: usize,
dst: usize,
n: u32,
c: u32,
h: u32,
w: u32,
num_groups: u32,
eps: f32,
dt: HalfFlag,
},
LayerNorm2d {
src: usize,
g: usize,
b: usize,
dst: usize,
n: u32,
c: u32,
h: u32,
w: u32,
eps: f32,
dt: HalfFlag,
},
ConvTranspose2d {
src: usize,
weight: usize,
dst: usize,
n: u32,
c_in: u32,
h: u32,
w_in: u32,
c_out: u32,
h_out: u32,
w_out: u32,
kh: u32,
kw: u32,
sh: u32,
sw: u32,
ph: u32,
pw: u32,
dh: u32,
dw: u32,
groups: u32,
dt: HalfFlag,
},
ResizeNearest2x {
src: usize,
dst: usize,
n: u32,
c: u32,
h: u32,
w: u32,
dt: HalfFlag,
},
RmsNorm {
src: usize,
g: usize,
b: usize,
dst: usize,
rows: u32,
h: u32,
eps: f32,
dt: HalfFlag,
},
BinaryFull {
lhs: usize,
rhs: usize,
dst: usize,
len: u32,
op: BinaryOp,
dt: HalfFlag,
},
BinaryBroadcast {
lhs: usize,
rhs: usize,
dst: usize,
len: u32,
op: BinaryOp,
dt: HalfFlag,
rank: u32,
out_dims: Vec<u32>,
lhs_strides: Vec<u32>,
rhs_strides: Vec<u32>,
},
BiasAdd {
src: usize,
bias: usize,
dst: usize,
m: u32,
n: u32,
dt: HalfFlag,
},
FusedResidualLN {
x: usize,
res: usize,
bias: usize,
g: usize,
b: usize,
out: usize,
rows: u32,
h: u32,
eps: f32,
has_bias: bool,
dt: HalfFlag,
},
FusedResidualRmsNorm {
x: usize,
res: usize,
bias: usize,
g: usize,
b: usize,
out: usize,
rows: u32,
h: u32,
eps: f32,
has_bias: bool,
dt: HalfFlag,
},
AdaLayerNorm {
x: usize,
scale: usize,
shift: usize,
out: usize,
rows: u32,
h: u32,
eps: f32,
layer_norm: bool,
lead_pack: [u32; 17],
dt: HalfFlag,
},
GatedResidual {
x: usize,
y: usize,
gate: usize,
out: usize,
rows: u32,
h: u32,
lead_pack: [u32; 17],
dt: HalfFlag,
},
AdaLayerNormBackward {
x: usize,
scale: usize,
dy: usize,
out: usize,
h: u32,
eps: f32,
layer_norm: bool,
seq_per_mod: u32,
mod_rows: u32,
dt: HalfFlag,
},
GatedResidualBackward {
y: usize,
gate: usize,
dy: usize,
out: usize,
h: u32,
seq_per_mod: u32,
mod_rows: u32,
dt: HalfFlag,
},
FusedRmsNormMulSilu {
x: usize,
g: usize,
b: usize,
z: usize,
out: usize,
rows: u32,
h: u32,
eps: f32,
dt: HalfFlag,
},
FusedDepthwiseConv1dBsc {
src: usize,
weight: usize,
dst: usize,
batch: u32,
width: u32,
out_seq: u32,
channels: u32,
k: u32,
silu: bool,
},
Gather {
table: usize,
idx: usize,
dst: usize,
num_idx: u32,
trailing: u32,
dt: HalfFlag,
},
Narrow {
src: usize,
dst: usize,
outer: u32,
src_axis: u32,
start: u32,
len: u32,
dt: HalfFlag,
},
SplitLastAxis {
src: usize,
outer: u32,
src_axis: u32,
dt: HalfFlag,
segments: Vec<(usize, u32, u32)>,
},
Copy {
src: usize,
dst: usize,
len: u32,
dt: HalfFlag,
},
Attention {
q: usize,
k: usize,
v: usize,
mask: usize,
out: usize,
batch: u32,
seq: u32, kv_seq: u32, heads: u32,
kv_heads: u32,
head_dim: u32,
mask_kind: u32,
window: u32,
dt: HalfFlag,
bhsd: u32,
score_scale: f32,
attn_logit_softcap: f32,
},
FusedAttn {
qkv: usize,
mask: usize,
cos: usize,
sin: usize,
out: usize,
batch: u32,
seq: u32,
heads: u32,
head_dim: u32,
mask_kind: u32,
scale_bits: u32,
has_rope: u32,
},
AttentionBackward {
q: usize,
k: usize,
v: usize,
dy: usize,
mask: usize,
out: usize,
batch: u32,
seq: u32,
kv_seq: u32,
heads: u32,
head_dim: u32,
mask_kind: u32,
window: u32,
wrt: u32,
bhsd: u32,
},
Rope {
src: usize,
cos: usize,
sin: usize,
dst: usize,
batch: u32,
seq: u32,
hidden: u32,
head_dim: u32,
n_rot: u32,
dt: HalfFlag,
src_row_stride: u32,
cos_per_token: bool,
interleaved: bool,
},
Softmax {
data: usize,
rows: u32,
cols: u32,
dt: HalfFlag,
},
SoftmaxCrossEntropyDense {
logits: usize,
targets: usize,
dst: usize,
n: u32,
c: u32,
},
SoftmaxCrossEntropyWithLogits {
logits: usize,
labels: usize,
dst: usize,
n: u32,
c: u32,
},
SoftmaxCrossEntropyBackward {
logits: usize,
labels: usize,
d_loss: usize,
dlogits: usize,
n: u32,
c: u32,
},
Cumsum {
src: usize,
dst: usize,
rows: u32,
cols: u32,
exclusive: bool,
},
FusedSwiGLU {
src: usize,
dst: usize,
n_half: u32,
total: u32,
src_dt: HalfFlag,
dst_dt: HalfFlag,
gate_first: bool,
},
Concat {
dst: usize,
outer: u32,
dst_axis: u32,
inner: u32,
dt: HalfFlag,
inputs: Vec<(usize, u32)>,
},
Compare {
lhs: usize,
rhs: usize,
dst: usize,
len: u32,
op: CmpOp,
lhs_scalar: bool,
rhs_scalar: bool,
},
Reduce {
src: usize,
dst: usize,
outer: u32,
reduced: u32,
inner: u32,
op: rlx_ir::op::ReduceOp,
dt: HalfFlag,
},
TopK {
src: usize,
dst: usize,
outer: u32,
axis_dim: u32,
k: u32,
},
GroupedMatMul {
input: usize,
weight: usize,
expert_idx: usize,
dst: usize,
m: u32,
k_dim: u32,
n: u32,
num_experts: u32,
},
DequantGroupedMatMulGguf {
input: usize,
w_q: usize,
expert_idx: usize,
dst: usize,
m: u32,
k_dim: u32,
n: u32,
num_experts: u32,
scheme: rlx_ir::quant::QuantScheme,
},
ScatterAdd {
updates: usize,
indices: usize,
dst: usize,
num_updates: u32,
out_dim: u32,
trailing: u32,
},
Transpose {
src: usize,
dst: usize,
total: u32,
out_dims: Vec<u32>,
in_strides: Vec<u32>,
},
GatherAxis {
table: usize,
idx: usize,
dst: usize,
outer: u32,
axis_dim: u32,
num_idx: u32,
trailing: u32,
},
Pool2D {
src: usize,
dst: usize,
n: u32,
c: u32,
h: u32,
w: u32,
h_out: u32,
w_out: u32,
kh: u32,
kw: u32,
sh: u32,
sw: u32,
ph: u32,
pw: u32,
kind: rlx_ir::op::ReduceOp,
},
Conv2D {
src: usize,
weight: usize,
dst: usize,
n: u32,
c_in: u32,
h: u32,
w: u32,
c_out: u32,
h_out: u32,
w_out: u32,
kh: u32,
kw: u32,
sh: u32,
sw: u32,
ph: u32,
pw: u32,
dh: u32,
dw: u32,
groups: u32,
},
Where {
cond: usize,
on_true: usize,
on_false: usize,
dst: usize,
len: u32,
cond_scalar: bool,
true_scalar: bool,
false_scalar: bool,
},
Fma {
a: usize,
b: usize,
c: usize,
dst: usize,
len: u32,
},
ElementwiseRegion {
len: u32,
num_inputs: u32,
num_steps: u32,
dst: usize,
input_offs: [u32; 16],
chain: [u32; 128], scalar_input_mask: u32,
input_modulus: [u32; 16],
prologue: u32,
out_n: u32,
out_c: u32,
out_h: u32,
out_w: u32,
prologue_input: u32,
},
BatchElementwiseRegion {
slice_len: u32,
num_batch: u32,
num_steps: u32,
base_dst: usize,
slice_elems: u32,
batch_input_offs: [u32; 64],
chain: [u32; 128],
scalar_input_mask: u32,
input_modulus: [u32; 16],
},
GatedDeltaNet {
q: usize,
k: usize,
v: usize,
g: usize,
beta: usize,
state: usize,
dst: usize,
batch: u32,
seq: u32,
heads: u32,
state_size: u32,
f16: bool,
},
SelectiveScan {
x: usize,
delta: usize,
a: usize,
b: usize,
c: usize,
dst: usize,
batch: u32,
seq: u32,
hidden: u32,
state_size: u32,
},
Sample {
logits: usize,
dst: usize,
batch: u32,
vocab: u32,
top_k: u32,
top_p: f32,
temperature: f32,
seed: u64,
},
Reverse {
src: usize,
dst: usize,
dims: Vec<u32>,
rev_mask: Vec<bool>,
elem_bytes: u8,
},
ArgReduce {
src: usize,
dst: usize,
outer: u32,
reduced: u32,
inner: u32,
is_max: bool,
},
Lstm {
x: usize,
w_ih: usize,
w_hh: usize,
bias: usize,
h0: usize,
c0: usize,
dst: usize,
batch: u32,
seq: u32,
input_size: u32,
hidden: u32,
num_layers: u32,
bidirectional: bool,
carry: bool,
},
Gru {
x: usize,
w_ih: usize,
w_hh: usize,
b_ih: usize,
b_hh: usize,
h0: usize,
dst: usize,
batch: u32,
seq: u32,
input_size: u32,
hidden: u32,
num_layers: u32,
bidirectional: bool,
carry: bool,
},
Rnn {
x: usize,
w_ih: usize,
w_hh: usize,
bias: usize,
h0: usize,
dst: usize,
batch: u32,
seq: u32,
input_size: u32,
hidden: u32,
num_layers: u32,
bidirectional: bool,
carry: bool,
relu: bool,
},
Mamba2 {
x: usize,
dt: usize,
a: usize,
b: usize,
c: usize,
dst: usize,
batch: u32,
seq: u32,
heads: u32,
head_dim: u32,
state_size: u32,
},
DequantMatMulGguf {
x: usize,
w_q: usize,
dst: usize,
m: u32,
k: u32,
n: u32,
scheme: rlx_ir::quant::QuantScheme,
x_f16: bool,
dst_f16: bool,
},
DequantMatMulInt8 {
x: usize,
w_q: usize,
scale: usize,
zp: usize,
dst: usize,
m: u32,
k: u32,
n: u32,
block_size: u32,
is_asymmetric: bool,
},
DequantMatMulInt4 {
x: usize,
w_q: usize,
scale: usize,
zp: usize,
dst: usize,
m: u32,
k: u32,
n: u32,
block_size: u32,
is_asymmetric: bool,
},
DequantMatMulFp8 {
x: usize,
w_q: usize,
scale: usize,
dst: usize,
m: u32,
k: u32,
n: u32,
e5m2: bool,
},
DequantMatMulNvfp4 {
x: usize,
w_q: usize,
scale: usize,
global_scale: usize,
dst: usize,
m: u32,
k: u32,
n: u32,
},
FusedMlpGateUpSwiGLU {
x: usize,
gate_w: usize,
up_w: usize,
dst: usize,
k: u32,
n: u32,
scheme: rlx_ir::quant::QuantScheme,
x_f16: bool,
dst_f16: bool,
},
FusedMlpGateUpGelu {
x: usize,
gate_w: usize,
up_w: usize,
dst: usize,
k: u32,
n: u32,
scheme: rlx_ir::quant::QuantScheme,
},
FusedMlpDownResidual {
x: usize,
w: usize,
res: usize,
dst: usize,
k: u32,
n: u32,
scheme: rlx_ir::quant::QuantScheme,
x_f16: bool,
dst_f16: bool,
res_f16: bool,
},
ScaledMatMul {
lhs: usize,
rhs: usize,
lhs_scale: usize,
rhs_scale: usize,
bias: usize,
dst: usize,
m: u32,
k: u32,
n: u32,
lhs_fmt: rlx_ir::ScaledFormat,
rhs_fmt: rlx_ir::ScaledFormat,
layout: rlx_ir::ScaleLayout,
has_bias: bool,
},
ScaledQuantize {
x: usize,
scale: usize,
dst: usize,
rows: u32,
cols: u32,
fmt: rlx_ir::ScaledFormat,
layout: rlx_ir::ScaleLayout,
},
ScaledDequantize {
codes: usize,
scale: usize,
dst: usize,
rows: u32,
cols: u32,
fmt: rlx_ir::ScaledFormat,
layout: rlx_ir::ScaleLayout,
},
ScaledQuantScale {
x: usize,
dst: usize,
rows: u32,
cols: u32,
fmt: rlx_ir::ScaledFormat,
layout: rlx_ir::ScaleLayout,
},
RmsNormBackwardInput {
x: usize,
gamma: usize,
beta: usize,
dy: usize,
dx: usize,
rows: u32,
h: u32,
eps: f32,
},
RmsNormBackwardGamma {
x: usize,
gamma: usize,
beta: usize,
dy: usize,
dgamma: usize,
rows: u32,
h: u32,
eps: f32,
},
RmsNormBackwardBeta {
x: usize,
gamma: usize,
beta: usize,
dy: usize,
dbeta: usize,
rows: u32,
h: u32,
eps: f32,
},
RopeBackward {
dy: usize,
cos: usize,
sin: usize,
dx: usize,
batch: u32,
seq: u32,
hidden: u32,
head_dim: u32,
n_rot: u32,
cos_len: u32,
},
CumsumBackward {
dy: usize,
dx: usize,
rows: u32,
cols: u32,
exclusive: bool,
},
GatherBackward {
dy: usize,
indices: usize,
dst: usize,
outer: u32,
axis_dim: u32,
num_idx: u32,
trailing: u32,
},
MaxPool2dBackward {
x: usize,
dy: usize,
dx: usize,
n: u32,
c: u32,
h: u32,
w: u32,
h_out: u32,
w_out: u32,
kh: u32,
kw: u32,
sh: u32,
sw: u32,
ph: u32,
pw: u32,
},
Conv2dBackwardInput {
dy: usize,
w: usize,
dx: usize,
n: u32,
c_in: u32,
h: u32,
w_in: u32,
c_out: u32,
h_out: u32,
w_out: u32,
kh: u32,
kw: u32,
sh: u32,
sw: u32,
ph: u32,
pw: u32,
dh: u32,
dw: u32,
groups: u32,
},
Conv2dBackwardWeight {
x: usize,
dy: usize,
dw: usize,
n: u32,
c_in: u32,
h: u32,
w: u32,
c_out: u32,
h_out: u32,
w_out: u32,
kh: u32,
kw: u32,
sh: u32,
sw: u32,
ph: u32,
pw: u32,
dh: u32,
dw_dil: u32,
groups: u32,
},
CustomOp {
kernel: Arc<dyn MetalKernel>,
inputs: Vec<(usize, u32, Shape)>, output: (usize, u32, Shape), attrs: Vec<u8>,
},
CustomGpuOp {
kernel: Arc<dyn MetalGpuKernel>,
inputs: Vec<(usize, u32, Shape)>, output: (usize, u32, Shape), attrs: Vec<u8>,
},
SpdHost {
op: rlx_ir::Op,
inputs: Vec<(usize, u32, Shape)>, output: (usize, u32, Shape), },
GaussianSplatRender {
positions_off: usize,
positions_len: usize,
scales_off: usize,
scales_len: usize,
rotations_off: usize,
rotations_len: usize,
opacities_off: usize,
opacities_len: usize,
colors_off: usize,
colors_len: usize,
sh_coeffs_off: usize,
sh_coeffs_len: usize,
meta_off: usize,
dst_off: usize,
dst_len: usize,
width: u32,
height: u32,
tile_size: u32,
radius_scale: f32,
alpha_cutoff: f32,
max_splat_steps: u32,
transmittance_threshold: f32,
max_list_entries: u32,
},
GaussianSplatRenderBackward {
positions_off: usize,
positions_len: usize,
scales_off: usize,
scales_len: usize,
rotations_off: usize,
rotations_len: usize,
opacities_off: usize,
opacities_len: usize,
colors_off: usize,
colors_len: usize,
sh_coeffs_off: usize,
sh_coeffs_len: usize,
meta_off: usize,
d_loss_off: usize,
d_loss_len: usize,
packed_off: usize,
packed_len: usize,
width: u32,
height: u32,
tile_size: u32,
radius_scale: f32,
alpha_cutoff: f32,
max_splat_steps: u32,
transmittance_threshold: f32,
max_list_entries: u32,
loss_grad_clip: f32,
sh_band: u32,
max_anisotropy: f32,
},
GaussianSplatPrepare {
positions_off: usize,
positions_len: usize,
scales_off: usize,
scales_len: usize,
rotations_off: usize,
rotations_len: usize,
opacities_off: usize,
opacities_len: usize,
colors_off: usize,
colors_len: usize,
sh_coeffs_off: usize,
sh_coeffs_len: usize,
meta_off: usize,
meta_len: usize,
prep_off: usize,
prep_len: usize,
width: u32,
height: u32,
tile_size: u32,
radius_scale: f32,
alpha_cutoff: f32,
max_splat_steps: u32,
transmittance_threshold: f32,
max_list_entries: u32,
},
GaussianSplatRasterize {
prep_off: usize,
prep_len: usize,
meta_off: usize,
meta_len: usize,
dst_off: usize,
dst_len: usize,
count: usize,
width: u32,
height: u32,
tile_size: u32,
alpha_cutoff: f32,
max_splat_steps: u32,
transmittance_threshold: f32,
max_list_entries: u32,
},
AxialRope2dHost {
src: usize,
dst: usize,
batch: u32,
seq: u32,
hidden: u32,
end_x: u32,
end_y: u32,
head_dim: u32,
num_heads: u32,
theta: f32,
repeat_factor: u32,
},
Im2Col {
x: usize,
col: usize,
n: u32,
c_in: u32,
h: u32,
w: u32,
h_out: u32,
w_out: u32,
kh: u32,
kw: u32,
sh: u32,
sw: u32,
ph: u32,
pw: u32,
dh: u32,
dw_dil: u32,
},
Fft1d {
src: usize,
dst: usize,
outer: u32,
n_complex: u32,
inverse: bool,
norm_tag: u32,
dtype: rlx_ir::DType,
real_input: bool,
},
VqAssign {
x: usize,
cb: usize,
out: usize,
n: u32,
d: u32,
k: u32,
metric: u32,
},
ScanHost {
desc: rlx_cpu::thunk::ScanHostDesc,
},
HostOp {
desc: rlx_cpu::thunk::HostOpDesc,
},
CpuIndexing {
thunk: rlx_cpu::thunk::IndexingThunk,
},
LogMel {
spec: usize,
filters: usize,
dst: usize,
outer: u32,
n_fft: u32,
n_bins: u32,
n_mels: u32,
},
LogMelBackward {
spec: usize,
filters: usize,
dy: usize,
dst: usize,
outer: u32,
n_fft: u32,
n_bins: u32,
n_mels: u32,
},
WelchPeaks {
spec: usize,
dst: usize,
welch_batch: u32,
n_fft: u32,
n_segments: u32,
k: u32,
},
RngNormal {
dst: usize,
len: u32,
mean: f32,
scale: f32,
key: u64,
op_seed: Option<f32>,
},
RngUniform {
dst: usize,
len: u32,
low: f32,
high: f32,
key: u64,
op_seed: Option<f32>,
},
}
pub struct ThunkSchedule {
pub thunks: Vec<Thunk>,
pub rng: std::sync::Arc<std::sync::RwLock<rlx_ir::RngOptions>>,
}
pub fn thunk_name(t: &Thunk) -> &'static str {
match t {
Thunk::Nop => "nop",
Thunk::Cast { .. } => "cast",
Thunk::CastHost { .. } => "cast_host",
Thunk::CastTruncF32 { .. } => "cast_trunc_f32",
Thunk::ScaledMatMul { .. } => "scaled_matmul",
Thunk::ScaledQuantize { .. } => "scaled_quantize",
Thunk::ScaledDequantize { .. } => "scaled_dequantize",
Thunk::ScaledQuantScale { .. } => "scaled_quant_scale",
Thunk::Sgemm { .. } => "sgemm",
Thunk::BatchedSgemm { .. } => "batched_sgemm",
Thunk::FusedMmBiasAct { .. } => "fused_mm_bias_act",
Thunk::FusedBinaryActivation { .. } => "fused_binary_activation",
Thunk::FusedTernaryActivation { .. } => "fused_ternary_activation",
Thunk::ActivationInPlace { .. } => "activation",
Thunk::ActivationOut { .. } => "activation_out",
Thunk::GeluApproxOut { .. } => "gelu_approx_out",
Thunk::GeluApproxHost { .. } => "gelu_approx_host",
Thunk::LayerNorm { .. } => "layer_norm",
Thunk::GroupNorm { .. } => "group_norm",
Thunk::LayerNorm2d { .. } => "layer_norm2d",
Thunk::ConvTranspose2d { .. } => "conv_transpose2d",
Thunk::RmsNorm { .. } => "rms_norm",
Thunk::ResizeNearest2x { .. } => "resize_nearest_2x",
Thunk::BinaryFull { .. } => "binary",
Thunk::BinaryBroadcast { .. } => "binary_broadcast",
Thunk::BiasAdd { .. } => "bias_add",
Thunk::FusedResidualLN { .. } => "fused_residual_ln",
Thunk::FusedResidualRmsNorm { .. } => "fused_residual_rms_norm",
Thunk::AdaLayerNorm { .. } => "ada_layer_norm",
Thunk::GatedResidual { .. } => "gated_residual",
Thunk::AdaLayerNormBackward { .. } => "ada_layer_norm_backward",
Thunk::GatedResidualBackward { .. } => "gated_residual_backward",
Thunk::FusedRmsNormMulSilu { .. } => "fused_rms_norm_mul_silu",
Thunk::FusedDepthwiseConv1dBsc { .. } => "fused_depthwise_conv1d_bsc",
Thunk::Gather { .. } => "gather",
Thunk::Narrow { .. } => "narrow",
Thunk::SplitLastAxis { .. } => "split_lastax",
Thunk::Copy { .. } => "copy",
Thunk::Attention { .. } => "attention",
Thunk::FusedAttn { .. } => "fused_attn",
Thunk::AttentionBackward { .. } => "attention_bwd",
Thunk::RmsNormBackwardInput { .. } => "rms_norm_backward_input",
Thunk::RmsNormBackwardGamma { .. } => "rms_norm_backward_gamma",
Thunk::RmsNormBackwardBeta { .. } => "rms_norm_backward_beta",
Thunk::RopeBackward { .. } => "rope_backward",
Thunk::CumsumBackward { .. } => "cumsum_backward",
Thunk::GatherBackward { .. } => "gather_backward",
Thunk::MaxPool2dBackward { .. } => "maxpool2d_backward",
Thunk::Conv2dBackwardInput { .. } => "conv2d_backward_input",
Thunk::Conv2dBackwardWeight { .. } => "conv2d_backward_weight",
Thunk::Rope { .. } => "rope",
Thunk::Softmax { .. } => "softmax",
Thunk::SoftmaxCrossEntropyDense { .. } => "softmax_cross_entropy_dense",
Thunk::SoftmaxCrossEntropyWithLogits { .. } => "softmax_cross_entropy_with_logits",
Thunk::SoftmaxCrossEntropyBackward { .. } => "softmax_cross_entropy_backward",
Thunk::Cumsum { .. } => "cumsum",
Thunk::FusedSwiGLU { .. } => "fused_swiglu",
Thunk::Concat { .. } => "concat",
Thunk::Compare { .. } => "compare",
Thunk::Reduce { .. } => "reduce",
Thunk::TopK { .. } => "topk",
Thunk::GroupedMatMul { .. } => "grouped_matmul",
Thunk::ScatterAdd { .. } => "scatter_add",
Thunk::Transpose { .. } => "transpose",
Thunk::GatherAxis { .. } => "gather_axis",
Thunk::Pool2D { .. } => "pool2d",
Thunk::Conv2D { .. } => "conv2d",
Thunk::Where { .. } => "where",
Thunk::Fma { .. } => "fma",
Thunk::ElementwiseRegion { .. } => "elementwise_region",
Thunk::BatchElementwiseRegion { .. } => "batch_elementwise_region",
Thunk::CustomOp { .. } => "custom_op",
Thunk::CustomGpuOp { .. } => "custom_gpu_op",
Thunk::SpdHost { .. } => "spd_host",
Thunk::GaussianSplatRender { .. } => "gaussian_splat_render",
Thunk::GaussianSplatRenderBackward { .. } => "gaussian_splat_render_backward",
Thunk::GaussianSplatPrepare { .. } => "gaussian_splat_prepare",
Thunk::GaussianSplatRasterize { .. } => "gaussian_splat_rasterize",
Thunk::AxialRope2dHost { .. } => "axial_rope2d_host",
Thunk::Im2Col { .. } => "im2col",
Thunk::Fft1d { .. } => "fft1d",
Thunk::VqAssign { .. } => "vq_assign",
Thunk::ScanHost { .. } => "scan_host",
Thunk::HostOp { .. } => "host_op",
Thunk::CpuIndexing { .. } => "cpu_indexing",
Thunk::LogMel { .. } => "log_mel",
Thunk::LogMelBackward { .. } => "log_mel_backward",
Thunk::WelchPeaks { .. } => "welch_peaks",
Thunk::RngNormal { .. } => "rng_normal",
Thunk::RngUniform { .. } => "rng_uniform",
Thunk::GatedDeltaNet { .. } => "gated_delta_net",
Thunk::SelectiveScan { .. } => "selective_scan",
Thunk::Sample { .. } => "sample",
Thunk::Reverse { .. } => "reverse",
Thunk::ArgReduce { .. } => "argreduce",
Thunk::Lstm { .. } => "lstm",
Thunk::Gru { .. } => "gru",
Thunk::Rnn { .. } => "rnn",
Thunk::Mamba2 { .. } => "mamba2",
Thunk::DequantMatMulGguf { .. } => "dequant_matmul_gguf",
Thunk::DequantGroupedMatMulGguf { .. } => "dequant_grouped_matmul_gguf",
Thunk::DequantMatMulInt8 { .. } => "dequant_matmul_int8",
Thunk::DequantMatMulInt4 { .. } => "dequant_matmul_int4",
Thunk::DequantMatMulFp8 { .. } => "dequant_matmul_fp8",
Thunk::DequantMatMulNvfp4 { .. } => "dequant_matmul_nvfp4",
Thunk::FusedMlpGateUpSwiGLU { .. } => "fused_mlp_gate_up_swiglu",
Thunk::FusedMlpGateUpGelu { .. } => "fused_mlp_gate_up_gelu",
Thunk::FusedMlpDownResidual { .. } => "fused_mlp_down_residual",
}
}
impl Thunk {
pub fn safe_for_active_extent(&self) -> bool {
match self {
Thunk::Nop
| Thunk::Cast { .. }
| Thunk::CastHost { .. }
| Thunk::CastTruncF32 { .. }
| Thunk::Copy { .. }
| Thunk::ActivationInPlace { .. }
| Thunk::ActivationOut { .. }
| Thunk::GeluApproxOut { .. }
| Thunk::GeluApproxHost { .. }
| Thunk::FusedBinaryActivation { .. }
| Thunk::FusedTernaryActivation { .. }
| Thunk::Sgemm { .. }
| Thunk::BatchedSgemm { .. }
| Thunk::FusedMmBiasAct { .. }
| Thunk::BiasAdd { .. }
| Thunk::LayerNorm { .. }
| Thunk::RmsNorm { .. }
| Thunk::Softmax { .. }
| Thunk::SoftmaxCrossEntropyDense { .. }
| Thunk::Cumsum { .. }
| Thunk::FusedResidualLN { .. }
| Thunk::FusedResidualRmsNorm { .. }
| Thunk::AdaLayerNorm { .. }
| Thunk::GatedResidual { .. }
| Thunk::AdaLayerNormBackward { .. }
| Thunk::GatedResidualBackward { .. }
| Thunk::FusedRmsNormMulSilu { .. }
| Thunk::FusedDepthwiseConv1dBsc { .. }
| Thunk::Gather { .. }
| Thunk::Compare { .. }
| Thunk::Where { .. }
| Thunk::FusedSwiGLU { .. }
| Thunk::ElementwiseRegion { .. }
| Thunk::BatchElementwiseRegion { .. }
| Thunk::Narrow { .. }
| Thunk::SplitLastAxis { .. }
| Thunk::Reduce { .. }
| Thunk::TopK { .. }
| Thunk::GroupedMatMul { .. }
| Thunk::GatherAxis { .. }
| Thunk::Concat { .. }
| Thunk::Conv2D { .. }
| Thunk::Pool2D { .. } => true,
Thunk::Attention { .. } => true,
Thunk::AttentionBackward { .. } => true,
Thunk::RmsNormBackwardInput { .. }
| Thunk::RmsNormBackwardGamma { .. }
| Thunk::RmsNormBackwardBeta { .. }
| Thunk::RopeBackward { .. }
| Thunk::CumsumBackward { .. }
| Thunk::GatherBackward { .. }
| Thunk::MaxPool2dBackward { .. }
| Thunk::Conv2dBackwardInput { .. }
| Thunk::Conv2dBackwardWeight { .. } => true,
Thunk::Rope { .. } => true,
Thunk::GatedDeltaNet { .. }
| Thunk::SelectiveScan { .. }
| Thunk::Sample { .. }
| Thunk::Reverse { .. }
| Thunk::ArgReduce { .. }
| Thunk::Lstm { .. }
| Thunk::Gru { .. }
| Thunk::Rnn { .. }
| Thunk::Mamba2 { .. }
| Thunk::DequantMatMulGguf { .. }
| Thunk::DequantGroupedMatMulGguf { .. }
| Thunk::DequantMatMulInt8 { .. }
| Thunk::DequantMatMulInt4 { .. }
| Thunk::DequantMatMulFp8 { .. }
| Thunk::DequantMatMulNvfp4 { .. }
| Thunk::FusedMlpGateUpSwiGLU { .. }
| Thunk::FusedMlpGateUpGelu { .. }
| Thunk::FusedMlpDownResidual { .. } => true,
Thunk::ScatterAdd { .. } => true,
Thunk::Transpose {
out_dims,
in_strides,
..
} => {
if out_dims.is_empty() || in_strides.is_empty() {
return false;
}
let inner: u32 = out_dims[1..].iter().product();
in_strides[0] == inner
}
_ => false,
}
}
}
mod compile;
impl ThunkSchedule {}
fn strides_dense_contiguous(rank: usize, dims: &[u32], strides: &[u32]) -> bool {
if rank == 0 || dims.len() < rank || strides.len() < rank {
return rank == 0;
}
let mut expected = 1u32;
for ax in (0..rank).rev() {
if strides[ax] != expected {
return false;
}
expected = expected.saturating_mul(dims[ax].max(1));
}
true
}
fn rewrite_simple_elementwise_regions(thunks: &mut Vec<Thunk>) {
let mut i = 0;
while i < thunks.len() {
match try_rewrite_elementwise_region(&thunks[i]) {
RegionRewrite::Keep => {
i += 1;
}
RegionRewrite::One(t) => {
thunks[i] = t;
i += 1;
}
RegionRewrite::Many(ts) => {
if ts.is_empty() {
i += 1;
continue;
}
let n = ts.len();
thunks.splice(i..=i, ts);
i += n;
}
}
}
}
enum RegionRewrite {
Keep,
One(Thunk),
Many(Vec<Thunk>),
}
fn region_is_dense(n_in: usize, scalar_input_mask: u32, input_modulus: &[u32; 16]) -> bool {
scalar_input_mask == 0 && !input_modulus.iter().take(n_in).any(|&m| m != 0)
}
fn decode_input_operand(enc: u32) -> Option<usize> {
if enc & 0x8000_0000 != 0 {
None
} else {
Some(enc as usize)
}
}
fn decode_step_operand(enc: u32) -> Option<usize> {
if enc & 0x8000_0000 == 0 {
None
} else {
Some((enc & 0x7FFF_FFFF) as usize)
}
}
fn map_chain_binary_op(sub: u32) -> Option<rlx_ir::op::BinaryOp> {
use rlx_ir::op::BinaryOp;
Some(match sub {
0 => BinaryOp::Add,
1 => BinaryOp::Sub,
2 => BinaryOp::Mul,
3 => BinaryOp::Div,
4 => BinaryOp::Max,
5 => BinaryOp::Min,
6 => BinaryOp::Pow,
_ => return None,
})
}
fn map_chain_activation(sub: u32) -> Option<rlx_ir::op::Activation> {
use rlx_ir::op::Activation;
Some(match sub {
0 | 1 => Activation::Gelu,
2 => Activation::Silu,
3 => Activation::Relu,
4 => Activation::Sigmoid,
5 => Activation::Tanh,
6 => Activation::Exp,
7 => Activation::Log,
8 => Activation::Sqrt,
9 => Activation::Rsqrt,
10 => Activation::Neg,
11 => Activation::Abs,
_ => return None,
})
}
fn try_rewrite_elementwise_region(t: &Thunk) -> RegionRewrite {
let Thunk::ElementwiseRegion {
len,
num_inputs,
num_steps,
dst,
input_offs,
chain,
scalar_input_mask,
input_modulus,
prologue,
out_n: _,
out_c: _,
out_h: _,
out_w: _,
prologue_input: _,
} = t
else {
return RegionRewrite::Keep;
};
if *prologue != 0 {
return RegionRewrite::Keep;
}
let n_in = *num_inputs as usize;
if !region_is_dense(n_in, *scalar_input_mask, input_modulus) {
return RegionRewrite::Keep;
}
let input_byte = |idx: usize| input_offs[idx] as usize * 4;
if *num_steps == 1 && chain[0] == 2 && n_in == 2 {
let Some(lhs_idx) = decode_input_operand(chain[2]) else {
return RegionRewrite::Keep;
};
let Some(rhs_idx) = decode_input_operand(chain[3]) else {
return RegionRewrite::Keep;
};
if lhs_idx >= n_in || rhs_idx >= n_in {
return RegionRewrite::Keep;
}
let Some(op) = map_chain_binary_op(chain[1]) else {
return RegionRewrite::Keep;
};
return RegionRewrite::One(Thunk::BinaryFull {
lhs: input_byte(lhs_idx),
rhs: input_byte(rhs_idx),
dst: *dst,
len: *len,
op,
dt: HalfFlag::F32,
});
}
if *num_steps == 2 && chain[0] == 2 && chain[4] == 0 {
let Some(lhs_idx) = decode_input_operand(chain[2]) else {
return RegionRewrite::Keep;
};
let Some(rhs_idx) = decode_input_operand(chain[3]) else {
return RegionRewrite::Keep;
};
if decode_step_operand(chain[6]) != Some(0) {
return RegionRewrite::Keep;
}
let Some(op) = map_chain_binary_op(chain[1]) else {
return RegionRewrite::Keep;
};
let Some(act) = map_chain_activation(chain[5]) else {
return RegionRewrite::Keep;
};
if lhs_idx >= n_in || rhs_idx >= n_in {
return RegionRewrite::Keep;
}
return RegionRewrite::One(Thunk::FusedBinaryActivation {
lhs: input_byte(lhs_idx),
rhs: input_byte(rhs_idx),
dst: *dst,
len: *len,
op,
act,
dt: HalfFlag::F32,
});
}
if *num_steps == 2 && chain[0] == 2 && chain[4] == 2 {
let Some(lhs0) = decode_input_operand(chain[2]) else {
return RegionRewrite::Keep;
};
let Some(rhs0) = decode_input_operand(chain[3]) else {
return RegionRewrite::Keep;
};
if decode_step_operand(chain[6]) != Some(0) {
return RegionRewrite::Keep;
}
let Some(rhs1) = decode_input_operand(chain[7]) else {
return RegionRewrite::Keep;
};
let Some(op0) = map_chain_binary_op(chain[1]) else {
return RegionRewrite::Keep;
};
let Some(op1) = map_chain_binary_op(chain[5]) else {
return RegionRewrite::Keep;
};
if lhs0 >= n_in || rhs0 >= n_in || rhs1 >= n_in {
return RegionRewrite::Keep;
}
return RegionRewrite::Many(vec![
Thunk::BinaryFull {
lhs: input_byte(lhs0),
rhs: input_byte(rhs0),
dst: *dst,
len: *len,
op: op0,
dt: HalfFlag::F32,
},
Thunk::BinaryFull {
lhs: *dst,
rhs: input_byte(rhs1),
dst: *dst,
len: *len,
op: op1,
dt: HalfFlag::F32,
},
]);
}
if *num_steps == 3 && chain[0] == 2 && chain[4] == 2 && chain[8] == 0 {
let Some(lhs0) = decode_input_operand(chain[2]) else {
return RegionRewrite::Keep;
};
let Some(rhs0) = decode_input_operand(chain[3]) else {
return RegionRewrite::Keep;
};
if decode_step_operand(chain[6]) != Some(0) {
return RegionRewrite::Keep;
}
let Some(rhs1) = decode_input_operand(chain[7]) else {
return RegionRewrite::Keep;
};
if decode_step_operand(chain[10]) != Some(1) {
return RegionRewrite::Keep;
}
let Some(op0) = map_chain_binary_op(chain[1]) else {
return RegionRewrite::Keep;
};
let Some(op1) = map_chain_binary_op(chain[5]) else {
return RegionRewrite::Keep;
};
let Some(act) = map_chain_activation(chain[9]) else {
return RegionRewrite::Keep;
};
if lhs0 >= n_in || rhs0 >= n_in || rhs1 >= n_in {
return RegionRewrite::Keep;
}
return RegionRewrite::One(Thunk::FusedTernaryActivation {
lhs: input_byte(lhs0),
rhs0: input_byte(rhs0),
rhs1: input_byte(rhs1),
dst: *dst,
len: *len,
op0,
op1,
act,
dt: HalfFlag::F32,
});
}
if *num_steps == 3 && chain[0] == 2 && chain[4] == 2 && chain[8] == 2 {
let Some(lhs0) = decode_input_operand(chain[2]) else {
return RegionRewrite::Keep;
};
let Some(rhs0) = decode_input_operand(chain[3]) else {
return RegionRewrite::Keep;
};
if decode_step_operand(chain[6]) != Some(0) {
return RegionRewrite::Keep;
}
let Some(rhs1) = decode_input_operand(chain[7]) else {
return RegionRewrite::Keep;
};
if decode_step_operand(chain[10]) != Some(1) {
return RegionRewrite::Keep;
};
let Some(rhs2) = decode_input_operand(chain[11]) else {
return RegionRewrite::Keep;
};
let Some(op0) = map_chain_binary_op(chain[1]) else {
return RegionRewrite::Keep;
};
let Some(op1) = map_chain_binary_op(chain[5]) else {
return RegionRewrite::Keep;
};
let Some(op2) = map_chain_binary_op(chain[9]) else {
return RegionRewrite::Keep;
};
if lhs0 >= n_in || rhs0 >= n_in || rhs1 >= n_in || rhs2 >= n_in {
return RegionRewrite::Keep;
}
return RegionRewrite::Many(vec![
Thunk::BinaryFull {
lhs: input_byte(lhs0),
rhs: input_byte(rhs0),
dst: *dst,
len: *len,
op: op0,
dt: HalfFlag::F32,
},
Thunk::BinaryFull {
lhs: *dst,
rhs: input_byte(rhs1),
dst: *dst,
len: *len,
op: op1,
dt: HalfFlag::F32,
},
Thunk::BinaryFull {
lhs: *dst,
rhs: input_byte(rhs2),
dst: *dst,
len: *len,
op: op2,
dt: HalfFlag::F32,
},
]);
}
if *num_steps == 2 && chain[0] == 1 && chain[4] == 2 {
let Some(cast_in) = decode_input_operand(chain[2]) else {
return RegionRewrite::Keep;
};
if decode_step_operand(chain[6]) != Some(0) {
return RegionRewrite::Keep;
}
let Some(rhs_idx) = decode_input_operand(chain[7]) else {
return RegionRewrite::Keep;
};
let Some(op) = map_chain_binary_op(chain[5]) else {
return RegionRewrite::Keep;
};
if cast_in >= n_in || rhs_idx >= n_in {
return RegionRewrite::Keep;
}
return RegionRewrite::One(Thunk::BinaryFull {
lhs: input_byte(cast_in),
rhs: input_byte(rhs_idx),
dst: *dst,
len: *len,
op,
dt: HalfFlag::F32,
});
}
if *num_steps == 1 && chain[0] == 0 && n_in > 0 {
let Some(src_idx) = decode_input_operand(chain[2]) else {
return RegionRewrite::Keep;
};
if src_idx >= n_in {
return RegionRewrite::Keep;
}
let data = input_byte(src_idx);
if data != *dst {
return RegionRewrite::Keep;
}
let Some(act) = map_chain_activation(chain[1]) else {
return RegionRewrite::Keep;
};
return RegionRewrite::One(Thunk::ActivationInPlace {
data,
len: *len,
act,
dt: HalfFlag::F32,
});
}
RegionRewrite::Keep
}
fn rewrite_dense_binary_broadcast(thunks: &mut [Thunk]) {
for t in thunks.iter_mut() {
let Thunk::BinaryBroadcast {
lhs,
rhs,
dst,
len,
op,
dt,
rank,
out_dims,
lhs_strides,
rhs_strides,
} = t
else {
continue;
};
let rank = *rank as usize;
if rank == 0
|| !strides_dense_contiguous(rank, out_dims, lhs_strides)
|| !strides_dense_contiguous(rank, out_dims, rhs_strides)
{
continue;
}
*t = Thunk::BinaryFull {
lhs: *lhs,
rhs: *rhs,
dst: *dst,
len: *len,
op: *op,
dt: *dt,
};
}
}
fn narrow_segments_partition(src_axis: u32, segments: &[(u32, u32)]) -> bool {
let mut sorted = segments.to_vec();
sorted.sort_by_key(|(s, _)| *s);
let mut end = 0u32;
for (start, len) in sorted {
if start != end {
return false;
}
end = end.saturating_add(len);
}
end == src_axis
}
pub static FUSED_DECODE_MLP_BLOCKS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
pub fn fused_decode_mlp_blocks() -> usize {
FUSED_DECODE_MLP_BLOCKS.load(std::sync::atomic::Ordering::Relaxed)
}
pub static FUSED_RESIDUAL_RMS_BLOCKS: std::sync::atomic::AtomicUsize =
std::sync::atomic::AtomicUsize::new(0);
pub fn fused_residual_rms_blocks() -> usize {
FUSED_RESIDUAL_RMS_BLOCKS.load(std::sync::atomic::Ordering::Relaxed)
}
fn mlp_io(t: &Thunk) -> Option<(Vec<usize>, Vec<usize>)> {
use Thunk::*;
let io = match t {
Nop => (vec![], vec![]),
Cast { src, dst, .. } | CastHost { src, dst, .. } | CastTruncF32 { src, dst, .. } => {
(vec![*src], vec![*dst])
}
Copy { src, dst, .. } => (vec![*src], vec![*dst]),
ActivationInPlace { data, .. } => (vec![*data], vec![*data]),
ActivationOut { src, dst, .. } => (vec![*src], vec![*dst]),
GeluApproxOut { src, dst, .. } | GeluApproxHost { src, dst, .. } => {
(vec![*src], vec![*dst])
}
BinaryFull { lhs, rhs, dst, .. } => (vec![*lhs, *rhs], vec![*dst]),
BinaryBroadcast { lhs, rhs, dst, .. } => (vec![*lhs, *rhs], vec![*dst]),
FusedBinaryActivation { lhs, rhs, dst, .. } => (vec![*lhs, *rhs], vec![*dst]),
FusedTernaryActivation {
lhs,
rhs0,
rhs1,
dst,
..
} => (vec![*lhs, *rhs0, *rhs1], vec![*dst]),
BiasAdd { src, bias, dst, .. } => (vec![*src, *bias], vec![*dst]),
Fma { a, b, c, dst, .. } => (vec![*a, *b, *c], vec![*dst]),
Where {
cond,
on_true,
on_false,
dst,
..
} => (vec![*cond, *on_true, *on_false], vec![*dst]),
Compare { lhs, rhs, dst, .. } => (vec![*lhs, *rhs], vec![*dst]),
RmsNorm { src, g, b, dst, .. } => (vec![*src, *g, *b], vec![*dst]),
LayerNorm { src, g, b, dst, .. } => (vec![*src, *g, *b], vec![*dst]),
FusedResidualLN {
x,
res,
bias,
g,
b,
out,
..
}
| FusedResidualRmsNorm {
x,
res,
bias,
g,
b,
out,
..
} => (vec![*x, *res, *bias, *g, *b], vec![*out]),
AdaLayerNorm {
x,
scale,
shift,
out,
..
} => (vec![*x, *scale, *shift], vec![*out]),
GatedResidual {
x, y, gate, out, ..
} => (vec![*x, *y, *gate], vec![*out]),
FusedRmsNormMulSilu {
x, g, b, z, out, ..
} => (vec![*x, *g, *b, *z], vec![*out]),
FusedDepthwiseConv1dBsc {
src, weight, dst, ..
} => (vec![*src, *weight], vec![*dst]),
Conv2D {
src, weight, dst, ..
} => (vec![*src, *weight], vec![*dst]),
FusedSwiGLU { src, dst, .. } => (vec![*src], vec![*dst]),
Softmax { data, .. } => (vec![*data], vec![*data]),
Rope {
src, cos, sin, dst, ..
} => (vec![*src, *cos, *sin], vec![*dst]),
Attention {
q, k, v, mask, out, ..
} => (vec![*q, *k, *v, *mask], vec![*out]),
FusedAttn {
qkv,
mask,
cos,
sin,
out,
has_rope,
..
} => {
let mut r = vec![*qkv, *mask];
if *has_rope != 0 {
r.push(*cos);
r.push(*sin);
}
(r, vec![*out])
}
Concat { dst, inputs, .. } => (inputs.iter().map(|(o, _)| *o).collect(), vec![*dst]),
SplitLastAxis { src, segments, .. } => {
(vec![*src], segments.iter().map(|(o, _, _)| *o).collect())
}
Narrow { src, dst, .. } => (vec![*src], vec![*dst]),
Gather {
table, idx, dst, ..
}
| GatherAxis {
table, idx, dst, ..
} => (vec![*table, *idx], vec![*dst]),
Sgemm { a, b, c, .. } => (vec![*a, *b], vec![*c]),
BatchedSgemm { a, b, c, .. } => (vec![*a, *b], vec![*c]),
FusedMmBiasAct { a, w, bias, c, .. } => (vec![*a, *w, *bias], vec![*c]),
DequantMatMulGguf { x, w_q, dst, .. } => (vec![*x, *w_q], vec![*dst]),
FusedMlpGateUpSwiGLU {
x,
gate_w,
up_w,
dst,
..
}
| FusedMlpGateUpGelu {
x,
gate_w,
up_w,
dst,
..
} => (vec![*x, *gate_w, *up_w], vec![*dst]),
FusedMlpDownResidual { x, w, res, dst, .. } => (vec![*x, *w, *res], vec![*dst]),
Reduce { src, dst, .. } => (vec![*src], vec![*dst]),
Transpose { src, dst, .. } => (vec![*src], vec![*dst]),
_ => return None,
};
Some(io)
}
const SENTINEL_OFF: usize = usize::MAX;
fn mlp_last_writer(thunks: &[Thunk], before: usize, off: usize) -> Result<Option<usize>, ()> {
if off == SENTINEL_OFF {
return Ok(None);
}
for i in (0..before).rev() {
match mlp_io(&thunks[i]) {
None => return Err(()),
Some((_, writes)) => {
if writes.contains(&off) {
return Ok(Some(i));
}
}
}
}
Ok(None)
}
fn mlp_find_forward<F: Fn(&Thunk) -> bool>(
thunks: &[Thunk],
after: usize,
pred: F,
) -> Option<usize> {
(after + 1..thunks.len()).find(|&i| pred(&thunks[i]))
}
fn mlp_value_dead_in_range(
thunks: &[Thunk],
producer: usize,
off: usize,
allowed: &[usize],
until: usize,
) -> bool {
if off == SENTINEL_OFF {
return false;
}
let until = until.min(thunks.len());
for (i, t) in thunks
.iter()
.enumerate()
.skip(producer + 1)
.take(until.saturating_sub(producer + 1))
{
let Some((reads, writes)) = mlp_io(t) else {
return false;
};
if writes.contains(&off) {
break;
}
if reads.contains(&off) && !allowed.contains(&i) {
return false;
}
}
true
}
fn mlp_unwritten_in_range(thunks: &[Thunk], start: usize, until: usize, off: usize) -> bool {
if off == SENTINEL_OFF {
return false;
}
for t in thunks.iter().take(until.min(thunks.len())).skip(start) {
match mlp_io(t) {
Some((_, writes)) => {
if writes.contains(&off) {
return false;
}
}
None => return false,
}
}
true
}
fn mlp_f32_ranges_overlap(a_off: usize, a_elems: u32, b_off: usize, b_elems: u32) -> bool {
if a_elems == 0 || b_elems == 0 {
return false;
}
if a_off == b_off {
return true;
}
let a_end = a_off.saturating_add(a_elems as usize * 4);
let b_end = b_off.saturating_add(b_elems as usize * 4);
a_off < b_end && b_off < a_end
}
fn mlp_gate_up_row_bytes(k: u32, scheme: rlx_ir::quant::QuantScheme) -> usize {
use rlx_ir::quant::QuantScheme;
match scheme {
QuantScheme::GgufQ4K => (k as usize / 256) * 144,
QuantScheme::GgufQ5_0 => (k as usize).div_ceil(32) * 22,
QuantScheme::GgufQ6K => (k as usize / 256) * 210,
QuantScheme::GgufQ1_0 => (k as usize / 128) * 18,
QuantScheme::GgufQ2_0 => (k as usize / 128) * 34,
_ => 0,
}
}
fn fuse_decode_mlp_combined_gate_up(
thunks: &mut [Thunk],
output_offsets: &std::collections::HashSet<usize>,
) {
if rlx_ir::env::var("RLX_METAL_FUSE_DECODE").as_deref() == Some("0") {
return;
}
use rlx_ir::quant::QuantScheme;
let verbose = rlx_ir::env::flag("RLX_METAL_FUSE_DECODE_LOG");
let as_packed_gate_up_mm = |t: &Thunk| -> Option<(usize, usize, usize, u32, u32, QuantScheme)> {
if let Thunk::DequantMatMulGguf {
x,
w_q,
dst,
m,
k,
n,
scheme,
..
} = *t
{
if m == 1 && matches!(scheme, QuantScheme::GgufQ4K | QuantScheme::GgufQ5_0) {
return Some((x, w_q, dst, k, n, scheme));
}
}
None
};
let as_narrow = |t: &Thunk| -> Option<(usize, usize, u32, u32)> {
if let Thunk::Narrow {
src,
dst,
start,
len,
..
} = *t
{
Some((src, dst, start, len))
} else {
None
}
};
let is_silu = |t: &Thunk| {
matches!(
t,
Thunk::ActivationInPlace {
act: Activation::Silu,
..
} | Thunk::ActivationOut {
act: Activation::Silu,
..
}
)
};
let is_gelu = |t: &Thunk| {
matches!(
t,
Thunk::ActivationInPlace {
act: Activation::GeluApprox,
..
} | Thunk::ActivationOut {
act: Activation::GeluApprox,
..
} | Thunk::GeluApproxOut { .. }
| Thunk::GeluApproxHost { .. }
)
};
let n_thunks = thunks.len();
let mut i = 0;
while i < n_thunks {
let (mul_lhs, mul_rhs, prod) = match &thunks[i] {
Thunk::BinaryFull {
lhs,
rhs,
dst,
op: BinaryOp::Mul,
..
} => (*lhs, *rhs, *dst),
_ => {
i += 1;
continue;
}
};
let mul_idx = i;
let silu_on = |off: usize| -> Option<usize> {
match mlp_last_writer(thunks, mul_idx, off) {
Ok(Some(idx)) if is_silu(&thunks[idx]) => Some(idx),
_ => None,
}
};
let gelu_on = |off: usize| -> Option<usize> {
match mlp_last_writer(thunks, mul_idx, off) {
Ok(Some(idx)) if is_gelu(&thunks[idx]) => Some(idx),
_ => None,
}
};
let (act_idx, up_off, use_gelu) = if let Some(idx) = silu_on(mul_lhs) {
(idx, mul_rhs, false)
} else if let Some(idx) = silu_on(mul_rhs) {
(idx, mul_lhs, false)
} else if let Some(idx) = gelu_on(mul_lhs) {
(idx, mul_rhs, true)
} else if let Some(idx) = gelu_on(mul_rhs) {
(idx, mul_lhs, true)
} else {
i += 1;
continue;
};
if use_gelu && rlx_ir::env::var("RLX_METAL_FUSE_DECODE_GELU").as_deref() == Some("0") {
i += 1;
continue;
}
let gate_src_off = match &thunks[act_idx] {
Thunk::ActivationInPlace { data, .. } => *data,
Thunk::ActivationOut { src, .. } => *src,
Thunk::GeluApproxOut { src, .. } | Thunk::GeluApproxHost { src, .. } => *src,
_ => {
i += 1;
continue;
}
};
let gate_narrow_idx = match mlp_last_writer(thunks, act_idx, gate_src_off) {
Ok(Some(idx)) if as_narrow(&thunks[idx]).is_some() => idx,
Ok(Some(copy_idx)) if matches!(&thunks[copy_idx], Thunk::Copy { .. }) => {
let Thunk::Copy { src, .. } = &thunks[copy_idx] else {
i += 1;
continue;
};
match mlp_last_writer(thunks, copy_idx, *src) {
Ok(Some(ni)) if as_narrow(&thunks[ni]).is_some() => ni,
_ => {
i += 1;
continue;
}
}
}
_ => {
i += 1;
continue;
}
};
let (combined_off, _gate_dst, gate_start, gate_len) =
as_narrow(&thunks[gate_narrow_idx]).unwrap();
let up_narrow_idx = match mlp_last_writer(thunks, mul_idx, up_off) {
Ok(Some(idx)) if as_narrow(&thunks[idx]).is_some() => idx,
_ => {
i += 1;
continue;
}
};
let (combined_up, _up_dst, up_start, up_len) = as_narrow(&thunks[up_narrow_idx]).unwrap();
if combined_off != combined_up || gate_len != up_len || up_start != gate_start + gate_len {
i += 1;
continue;
}
let n_half = gate_len;
let combined_mm_idx = match mlp_last_writer(thunks, gate_narrow_idx, combined_off) {
Ok(Some(idx)) if as_packed_gate_up_mm(&thunks[idx]).is_some() => idx,
_ => {
i += 1;
continue;
}
};
let (comb_x, comb_w, _comb_dst, comb_k, comb_n, comb_scheme) =
as_packed_gate_up_mm(&thunks[combined_mm_idx]).unwrap();
let (comb_x_f16, comb_dst_f16) = match &thunks[combined_mm_idx] {
Thunk::DequantMatMulGguf { x_f16, dst_f16, .. } => (*x_f16, *dst_f16),
_ => (false, false),
};
if comb_x_f16 || comb_dst_f16 {
i += 1;
continue;
}
if comb_n != 2 * n_half {
i += 1;
continue;
}
if mlp_f32_ranges_overlap(comb_x, comb_k, prod, n_half) {
i += 1;
continue;
}
let row_bytes = mlp_gate_up_row_bytes(comb_k, comb_scheme);
if row_bytes == 0 {
i += 1;
continue;
}
let gate_w = comb_w;
let up_w = comb_w + (n_half as usize) * row_bytes;
let down_mm_idx = mlp_find_forward(thunks, mul_idx, |t| {
matches!(
t,
Thunk::DequantMatMulGguf {
x,
m: 1,
scheme: QuantScheme::GgufQ4K
| QuantScheme::GgufQ5_0
| QuantScheme::GgufQ6K,
..
} if *x == prod
)
});
let Some(down_mm_idx) = down_mm_idx else {
i += 1;
continue;
};
let (down_w, down_dst, down_k, down_n, down_scheme) = match &thunks[down_mm_idx] {
Thunk::DequantMatMulGguf {
w_q,
dst,
k,
n,
scheme,
..
} => (*w_q, *dst, *k, *n, *scheme),
_ => unreachable!(),
};
let down_add_tail = {
let add_idx = mlp_find_forward(thunks, down_mm_idx, |t| {
matches!(
t,
Thunk::BinaryFull {
lhs,
rhs,
op: BinaryOp::Add,
..
} if *lhs == down_dst || *rhs == down_dst
)
});
add_idx.map(|add_idx| {
let (res_off, out_off) = match &thunks[add_idx] {
Thunk::BinaryFull { lhs, rhs, dst, .. } => {
let res = if *lhs == down_dst { *rhs } else { *lhs };
(res, *dst)
}
_ => unreachable!(),
};
(add_idx, res_off, out_off)
})
};
let layer_until = down_mm_idx + 1;
let mut dead_ok =
mlp_value_dead_in_range(
thunks,
combined_mm_idx,
combined_off,
&[gate_narrow_idx, up_narrow_idx],
layer_until,
) && mlp_value_dead_in_range(
thunks,
gate_narrow_idx,
gate_src_off,
&[act_idx],
layer_until,
) && mlp_value_dead_in_range(thunks, up_narrow_idx, up_off, &[mul_idx], layer_until)
&& mlp_value_dead_in_range(thunks, act_idx, gate_src_off, &[mul_idx], layer_until);
if let Some((add_idx, _, _)) = down_add_tail {
dead_ok &=
mlp_value_dead_in_range(thunks, down_mm_idx, down_dst, &[add_idx], layer_until);
}
let no_output_clash = ![combined_off, gate_src_off, up_off]
.iter()
.any(|o| output_offsets.contains(o));
let gelu_graph_ok = if use_gelu {
let prod_readers = [down_mm_idx];
mlp_value_dead_in_range(thunks, mul_idx, prod, &prod_readers, thunks.len())
&& mlp_value_dead_in_range(
thunks,
combined_mm_idx,
combined_off,
&[gate_narrow_idx, up_narrow_idx],
thunks.len(),
)
&& match mlp_last_writer(thunks, combined_mm_idx, comb_x) {
Ok(Some(w)) => mlp_unwritten_in_range(thunks, w + 1, mul_idx, comb_x),
_ => false,
}
} else {
true
};
if !dead_ok || !no_output_clash || !gelu_graph_ok {
i += 1;
continue;
}
thunks[combined_mm_idx] = Thunk::Nop;
thunks[gate_narrow_idx] = Thunk::Nop;
thunks[up_narrow_idx] = Thunk::Nop;
thunks[act_idx] = Thunk::Nop;
thunks[mul_idx] = if use_gelu {
Thunk::FusedMlpGateUpGelu {
x: comb_x,
gate_w,
up_w,
dst: prod,
k: comb_k,
n: n_half,
scheme: comb_scheme,
}
} else {
Thunk::FusedMlpGateUpSwiGLU {
x: comb_x,
gate_w,
up_w,
dst: prod,
k: comb_k,
n: n_half,
scheme: comb_scheme,
x_f16: false,
dst_f16: false,
}
};
if let Some((add_idx, res_off, out_off)) = down_add_tail {
thunks[down_mm_idx] = Thunk::Nop;
thunks[add_idx] = Thunk::FusedMlpDownResidual {
x: prod,
w: down_w,
res: res_off,
dst: out_off,
k: down_k,
n: down_n,
scheme: down_scheme,
x_f16: false,
dst_f16: false,
res_f16: false,
};
}
FUSED_DECODE_MLP_BLOCKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if verbose {
if down_add_tail.is_some() {
eprintln!(
"[rlx-metal] fuse_decode_mlp_combined: gate_up {comb_scheme:?} k={comb_k} n={n_half} \
act={} down {down_scheme:?} — 7 dispatches → 2",
if use_gelu { "gelu" } else { "silu" }
);
} else {
eprintln!(
"[rlx-metal] fuse_decode_mlp_combined: gate_up {comb_scheme:?} k={comb_k} n={n_half} \
act={} (post_ffn norm blocks down fuse) — 5 dispatches → 1",
if use_gelu { "gelu" } else { "silu" }
);
}
}
i += 1;
}
}
fn fuse_decode_mlp(thunks: &mut [Thunk], output_offsets: &std::collections::HashSet<usize>) {
if rlx_ir::env::var("RLX_METAL_FUSE_DECODE").as_deref() == Some("0") {
return;
}
use rlx_ir::quant::QuantScheme;
let verbose = rlx_ir::env::flag("RLX_METAL_FUSE_DECODE_LOG");
let as_packed_gate_up_mm = |t: &Thunk| -> Option<(usize, usize, usize, u32, u32, QuantScheme)> {
if let Thunk::DequantMatMulGguf {
x,
w_q,
dst,
m,
k,
n,
scheme,
..
} = *t
{
if m == 1
&& matches!(
scheme,
QuantScheme::GgufQ4K
| QuantScheme::GgufQ5_0
| QuantScheme::GgufQ1_0
| QuantScheme::GgufQ2_0
)
{
if matches!(scheme, QuantScheme::GgufQ1_0 | QuantScheme::GgufQ2_0)
&& !k.is_multiple_of(128)
{
return None;
}
if matches!(scheme, QuantScheme::GgufQ2_0)
&& rlx_ir::env::flag("RLX_METAL_Q2_0_FUSED_DISABLE")
{
return None;
}
return Some((x, w_q, dst, k, n, scheme));
}
}
None
};
let n_thunks = thunks.len();
let mut i = 0;
while i < n_thunks {
let (mul_lhs, mul_rhs, prod) = match &thunks[i] {
Thunk::BinaryFull {
lhs,
rhs,
dst,
op: BinaryOp::Mul,
..
} => (*lhs, *rhs, *dst),
_ => {
i += 1;
continue;
}
};
let mul_idx = i;
let is_silu = |t: &Thunk| {
matches!(
t,
Thunk::ActivationInPlace {
act: Activation::Silu,
..
} | Thunk::ActivationOut {
act: Activation::Silu,
..
}
)
};
let is_gelu = |t: &Thunk| {
matches!(
t,
Thunk::ActivationInPlace {
act: Activation::GeluApprox,
..
} | Thunk::ActivationOut {
act: Activation::GeluApprox,
..
} | Thunk::GeluApproxOut { .. }
| Thunk::GeluApproxHost { .. }
)
};
let act_writer = |gate_act: usize| -> Option<(usize, bool)> {
match mlp_last_writer(thunks, mul_idx, gate_act) {
Ok(Some(idx)) if is_silu(&thunks[idx]) => Some((idx, false)),
Ok(Some(idx)) if is_gelu(&thunks[idx]) => Some((idx, true)),
_ => None,
}
};
let (gate_act_off, up_off, act_idx, use_gelu) =
if let Some((idx, gelu)) = act_writer(mul_lhs) {
(mul_lhs, mul_rhs, idx, gelu)
} else if let Some((idx, gelu)) = act_writer(mul_rhs) {
(mul_rhs, mul_lhs, idx, gelu)
} else {
i += 1;
continue;
};
if use_gelu && rlx_ir::env::var("RLX_METAL_FUSE_DECODE_GELU").as_deref() == Some("0") {
i += 1;
continue;
}
let gate_src_off = match &thunks[act_idx] {
Thunk::ActivationInPlace { data, .. } => *data,
Thunk::ActivationOut { src, .. } => *src,
Thunk::GeluApproxOut { src, .. } | Thunk::GeluApproxHost { src, .. } => *src,
_ => {
i += 1;
continue;
}
};
let gate_producer = match mlp_last_writer(thunks, act_idx, gate_src_off) {
Ok(Some(idx)) => idx,
_ => {
i += 1;
continue;
}
};
let (copy_idx, gate_mm_idx, gate_mm_off) = match &thunks[gate_producer] {
Thunk::Copy { src, .. } => match mlp_last_writer(thunks, gate_producer, *src) {
Ok(Some(gm)) if as_packed_gate_up_mm(&thunks[gm]).is_some() => {
(Some(gate_producer), gm, *src)
}
_ => {
i += 1;
continue;
}
},
t if as_packed_gate_up_mm(t).is_some() => (None, gate_producer, gate_src_off),
_ => {
i += 1;
continue;
}
};
let (gate_x, gate_w, _g_dst, gate_k, gate_n, gate_scheme) =
as_packed_gate_up_mm(&thunks[gate_mm_idx]).unwrap();
let (gate_x_f16, gate_dst_f16) = match &thunks[gate_mm_idx] {
Thunk::DequantMatMulGguf { x_f16, dst_f16, .. } => (*x_f16, *dst_f16),
_ => (false, false),
};
if (gate_x_f16 || gate_dst_f16)
&& !matches!(gate_scheme, QuantScheme::GgufQ1_0 | QuantScheme::GgufQ2_0)
{
i += 1;
continue;
}
if mlp_f32_ranges_overlap(gate_x, gate_k, prod, gate_n) {
i += 1;
continue;
}
let up_mm_idx = match mlp_last_writer(thunks, mul_idx, up_off) {
Ok(Some(idx)) => idx,
_ => {
i += 1;
continue;
}
};
let Some((up_x, up_w, _u_dst, up_k, up_n, up_scheme)) =
as_packed_gate_up_mm(&thunks[up_mm_idx])
else {
i += 1;
continue;
};
if up_x != gate_x || up_k != gate_k || up_n != gate_n || up_scheme != gate_scheme {
i += 1;
continue;
}
let down_mm_idx = mlp_find_forward(thunks, mul_idx, |t| match t {
Thunk::DequantMatMulGguf {
x,
m: 1,
scheme: QuantScheme::GgufQ4K | QuantScheme::GgufQ5_0 | QuantScheme::GgufQ6K,
..
} if *x == prod => true,
Thunk::DequantMatMulGguf {
x,
m: 1,
k,
scheme: QuantScheme::GgufQ1_0 | QuantScheme::GgufQ2_0,
..
} if *x == prod && k.is_multiple_of(128) => true,
_ => false,
});
let Some(down_mm_idx) = down_mm_idx else {
i += 1;
continue;
};
let (down_w, down_dst, down_k, down_n, down_scheme, down_x_f16, down_dst_f16) =
match &thunks[down_mm_idx] {
Thunk::DequantMatMulGguf {
w_q,
dst,
k,
n,
scheme,
x_f16,
dst_f16,
..
} => (*w_q, *dst, *k, *n, *scheme, *x_f16, *dst_f16),
_ => unreachable!(),
};
if matches!(down_scheme, QuantScheme::GgufQ2_0)
&& rlx_ir::env::flag("RLX_METAL_Q2_0_FUSED_DISABLE")
{
i += 1;
continue;
}
if (down_x_f16 || down_dst_f16)
&& !matches!(down_scheme, QuantScheme::GgufQ1_0 | QuantScheme::GgufQ2_0)
{
i += 1;
continue;
}
let add_idx = mlp_find_forward(
thunks,
down_mm_idx,
|t| matches!(t, Thunk::BinaryFull { lhs, rhs, op: BinaryOp::Add, .. } if *lhs == down_dst || *rhs == down_dst),
);
let Some(add_idx) = add_idx else {
i += 1;
continue;
};
let (res_off, out_off, add_f16) = match &thunks[add_idx] {
Thunk::BinaryFull {
lhs, rhs, dst, dt, ..
} => {
let res = if *lhs == down_dst { *rhs } else { *lhs };
(res, *dst, matches!(dt, HalfFlag::F16))
}
_ => unreachable!(),
};
let layer_until = down_mm_idx + 1;
let dead_ok =
mlp_value_dead_in_range(
thunks,
gate_mm_idx,
gate_mm_off,
&[copy_idx.unwrap_or(act_idx)],
layer_until,
) && mlp_value_dead_in_range(thunks, up_mm_idx, up_off, &[mul_idx], layer_until)
&& mlp_value_dead_in_range(thunks, act_idx, gate_act_off, &[mul_idx], layer_until)
&& mlp_value_dead_in_range(thunks, down_mm_idx, down_dst, &[add_idx], layer_until);
let no_output_clash = ![gate_mm_off, up_off, gate_act_off, down_dst]
.iter()
.any(|o| output_offsets.contains(o));
let gelu_graph_ok = if use_gelu {
mlp_value_dead_in_range(thunks, mul_idx, prod, &[down_mm_idx, add_idx], thunks.len())
&& match mlp_last_writer(thunks, gate_mm_idx, gate_x) {
Ok(Some(w)) => mlp_unwritten_in_range(thunks, w + 1, mul_idx, gate_x),
_ => false,
}
} else {
true
};
if !dead_ok || !no_output_clash || !gelu_graph_ok {
i += 1;
continue;
}
thunks[gate_mm_idx] = Thunk::Nop;
thunks[up_mm_idx] = Thunk::Nop;
thunks[act_idx] = Thunk::Nop;
thunks[down_mm_idx] = Thunk::Nop;
if let Some(c) = copy_idx {
thunks[c] = Thunk::Nop;
}
thunks[mul_idx] = if use_gelu {
Thunk::FusedMlpGateUpGelu {
x: gate_x,
gate_w,
up_w,
dst: prod,
k: gate_k,
n: gate_n,
scheme: gate_scheme,
}
} else {
Thunk::FusedMlpGateUpSwiGLU {
x: gate_x,
gate_w,
up_w,
dst: prod,
k: gate_k,
n: gate_n,
scheme: gate_scheme,
x_f16: gate_x_f16,
dst_f16: gate_dst_f16,
}
};
thunks[add_idx] = Thunk::FusedMlpDownResidual {
x: prod,
w: down_w,
res: res_off,
dst: out_off,
k: down_k,
n: down_n,
scheme: down_scheme,
x_f16: down_x_f16,
dst_f16: add_f16,
res_f16: add_f16,
};
FUSED_DECODE_MLP_BLOCKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if verbose {
eprintln!(
"[rlx-metal] fuse_decode_mlp: block fused (gate/up {gate_scheme:?} k={gate_k} n={gate_n}, \
down {down_scheme:?} k={down_k} n={down_n}, act={}) — 6 dispatches → 2",
if use_gelu { "gelu" } else { "silu" }
);
}
i += 1;
}
}
fn fuse_residual_rms_norm(thunks: &mut [Thunk], output_offsets: &std::collections::HashSet<usize>) {
if rlx_ir::env::var("RLX_METAL_FUSE_RESIDUAL_RMS").as_deref() == Some("0")
|| rlx_ir::env::var("RLX_METAL_FUSE_DECODE").as_deref() == Some("0")
{
return;
}
let verbose = rlx_ir::env::flag("RLX_METAL_FUSE_DECODE_LOG");
let n_thunks = thunks.len();
let mut fused = 0usize;
let mut i = 0;
let as_add =
|t: &Thunk, expect_dst: usize, expect_len: u32, dt: HalfFlag| -> Option<(usize, usize)> {
match t {
Thunk::BinaryFull {
lhs,
rhs,
dst,
len,
op: BinaryOp::Add,
dt: add_dt,
} if *dst == expect_dst && *add_dt == dt && *len == expect_len => {
Some((*lhs, *rhs))
}
_ => None,
}
};
let last_writer_skip_opaque = |thunks: &[Thunk], before: usize, off: usize| -> Option<usize> {
for j in (0..before).rev() {
match mlp_io(&thunks[j]) {
None => continue,
Some((_, writes)) if writes.contains(&off) => return Some(j),
_ => {}
}
}
None
};
let is_proj_branch = |thunks: &[Thunk], mut before: usize, off: usize| -> bool {
let mut cur = off;
for _ in 0..4 {
let Some(idx) = last_writer_skip_opaque(thunks, before, cur) else {
return false;
};
match &thunks[idx] {
Thunk::DequantMatMulGguf { .. }
| Thunk::Sgemm { .. }
| Thunk::BatchedSgemm { .. }
| Thunk::FusedMmBiasAct { .. }
| Thunk::FusedMlpDownResidual { .. } => return true,
Thunk::Copy { src, dst, .. } if *dst == cur => {
cur = *src;
before = idx;
}
Thunk::BiasAdd { src, dst, .. } if *dst == cur => {
cur = *src;
before = idx;
}
_ => return false,
}
}
false
};
while i < n_thunks {
let (mut src, g_off, b_off, dst, rows, h, eps, dt) = match &thunks[i] {
Thunk::RmsNorm {
src,
g,
b,
dst,
rows,
h,
eps,
dt,
} => (*src, *g, *b, *dst, *rows, *h, *eps, *dt),
_ => {
i += 1;
continue;
}
};
if h < 1024 || rows != 1 || !matches!(dt, HalfFlag::F32) {
i += 1;
continue;
}
let expect_len = rows.saturating_mul(h);
if expect_len == 0 || output_offsets.contains(&src) {
i += 1;
continue;
}
let mut copy_i = None;
if let Ok(Some(idx)) = mlp_last_writer(thunks, i, src) {
if let Thunk::Copy {
src: csrc,
dst: cdst,
len,
dt: cdt,
} = &thunks[idx]
{
if *cdst == src
&& *cdt == dt
&& *len == expect_len
&& !output_offsets.contains(csrc)
{
copy_i = Some(idx);
src = *csrc;
}
}
}
let add_before = copy_i.unwrap_or(i);
let add_i = match mlp_last_writer(thunks, add_before, src) {
Ok(Some(idx)) => idx,
_ => {
i += 1;
continue;
}
};
let gap = (add_i + 1..add_before)
.filter(|&j| !matches!(thunks[j], Thunk::Nop))
.count();
if gap > 2 {
i += 1;
continue;
}
let Some((x, res)) = as_add(&thunks[add_i], src, expect_len, dt) else {
i += 1;
continue;
};
if !(is_proj_branch(thunks, add_i, x) || is_proj_branch(thunks, add_i, res)) {
i += 1;
continue;
}
if mlp_f32_ranges_overlap(dst, expect_len, x, expect_len)
|| mlp_f32_ranges_overlap(dst, expect_len, res, expect_len)
{
i += 1;
continue;
}
let add_readers: Vec<usize> = match copy_i {
Some(c) => vec![c, i],
None => vec![i],
};
if !mlp_value_dead_in_range(thunks, add_i, src, &add_readers, thunks.len()) {
i += 1;
continue;
}
if let Some(c) = copy_i {
let copy_dst = match &thunks[c] {
Thunk::Copy { dst, .. } => *dst,
_ => src,
};
if !mlp_value_dead_in_range(thunks, c, copy_dst, &[i], thunks.len()) {
i += 1;
continue;
}
}
thunks[add_i] = Thunk::Nop;
if let Some(c) = copy_i {
thunks[c] = Thunk::Nop;
}
thunks[i] = Thunk::FusedResidualRmsNorm {
x,
res,
bias: 0,
g: g_off,
b: b_off,
out: dst,
rows,
h,
eps,
has_bias: false,
dt,
};
fused += 1;
FUSED_RESIDUAL_RMS_BLOCKS.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
if verbose {
eprintln!("[rlx-metal] fuse_residual_rms_norm: rows={rows} h={h} — add(+copy)+rms → 1");
}
i += 1;
}
if verbose && fused > 0 {
eprintln!("[rlx-metal] fuse_residual_rms_norm: {fused} blocks fused");
}
}
fn fuse_gdn_gated_norm(thunks: &mut [Thunk], output_offsets: &std::collections::HashSet<usize>) {
if rlx_ir::env::var("RLX_METAL_FUSE_DECODE").as_deref() == Some("0")
|| rlx_ir::env::var("RLX_METAL_FUSE_GDN_NORM").as_deref() == Some("0")
{
return;
}
let verbose = rlx_ir::env::flag("RLX_METAL_FUSE_DECODE_LOG");
let n_thunks = thunks.len();
let mut fused = 0usize;
let mut i = 0;
while i < n_thunks {
let (mul_lhs, mul_rhs, mul_dst, mul_len, mul_dt) = match &thunks[i] {
Thunk::BinaryFull {
lhs,
rhs,
dst,
len,
op: BinaryOp::Mul,
dt,
} => (*lhs, *rhs, *dst, *len, *dt),
_ => {
i += 1;
continue;
}
};
if matches!(mul_dt, HalfFlag::F16) {
i += 1;
continue;
}
let try_rms = |off: usize| -> Option<(usize, usize, usize, usize, u32, u32, f32)> {
match mlp_last_writer(thunks, i, off) {
Ok(Some(idx)) => match &thunks[idx] {
Thunk::RmsNorm {
src,
g,
b,
dst,
rows,
h,
eps,
dt,
} if *dst == off && *dt == mul_dt => Some((idx, *src, *g, *b, *rows, *h, *eps)),
_ => None,
},
_ => None,
}
};
let (rms_i, x, g_off, b_off, rows, h, eps, z_side) = if let Some(t) = try_rms(mul_lhs) {
(t.0, t.1, t.2, t.3, t.4, t.5, t.6, mul_rhs)
} else if let Some(t) = try_rms(mul_rhs) {
(t.0, t.1, t.2, t.3, t.4, t.5, t.6, mul_lhs)
} else {
i += 1;
continue;
};
if rows.saturating_mul(h) != mul_len {
i += 1;
continue;
}
let rms_dst = match &thunks[rms_i] {
Thunk::RmsNorm { dst, .. } => *dst,
_ => {
i += 1;
continue;
}
};
if output_offsets.contains(&rms_dst) {
i += 1;
continue;
}
let silu_i = match mlp_last_writer(thunks, i, z_side) {
Ok(Some(idx)) => idx,
_ => {
i += 1;
continue;
}
};
let (z_src, copy_i) = match &thunks[silu_i] {
Thunk::ActivationInPlace {
data,
act: Activation::Silu,
dt,
..
} if *data == z_side && *dt == mul_dt => {
let copy = match mlp_last_writer(thunks, silu_i, z_side) {
Ok(Some(cidx)) => match &thunks[cidx] {
Thunk::Copy {
src,
dst,
len,
dt: cdt,
} if *dst == z_side && *cdt == mul_dt && *len == mul_len => {
Some((cidx, *src))
}
_ => None,
},
_ => None,
};
match copy {
Some((cidx, src)) => (src, Some(cidx)),
None => (z_side, None), }
}
Thunk::ActivationOut {
src,
dst,
act: Activation::Silu,
dt,
..
} if *dst == z_side && *dt == mul_dt => (*src, None),
_ => {
i += 1;
continue;
}
};
if !mlp_value_dead_in_range(thunks, rms_i, rms_dst, &[i], i + 1) {
i += 1;
continue;
}
if !mlp_value_dead_in_range(thunks, silu_i, z_side, &[i], i + 1) {
i += 1;
continue;
}
if let Some(c) = copy_i {
if output_offsets.contains(&z_src) {
i += 1;
continue;
}
let _ = c;
}
if let Some(c) = copy_i {
thunks[c] = Thunk::Nop;
}
thunks[silu_i] = Thunk::Nop;
thunks[rms_i] = Thunk::Nop;
thunks[i] = Thunk::FusedRmsNormMulSilu {
x,
g: g_off,
b: b_off,
z: z_src,
out: mul_dst,
rows,
h,
eps,
dt: mul_dt,
};
fused += 1;
if verbose {
eprintln!("[rlx-metal] fuse_gdn_gated_norm: rows={rows} h={h} — silu+rms+mul → 1");
}
i += 1;
}
if verbose && fused > 0 {
eprintln!("[rlx-metal] fuse_gdn_gated_norm: {fused} blocks fused");
}
}
fn fuse_depthwise_conv1d_bsc(
thunks: &mut [Thunk],
output_offsets: &std::collections::HashSet<usize>,
) {
if rlx_ir::env::var("RLX_METAL_FUSE_DECODE").as_deref() == Some("0")
|| rlx_ir::env::var("RLX_METAL_FUSE_DEPTHWISE").as_deref() == Some("0")
{
return;
}
let verbose = rlx_ir::env::flag("RLX_METAL_FUSE_DECODE_LOG");
let n_thunks = thunks.len();
let mut fused = 0usize;
let mut i = 0;
while i < n_thunks {
let (conv_src, weight, conv_dst, batch, c_in, width, out_seq, kw) = match &thunks[i] {
Thunk::Conv2D {
src,
weight,
dst,
n,
c_in,
h: 1,
w,
c_out,
h_out: 1,
w_out,
kh: 1,
kw,
sh: 1,
sw: 1,
ph: 0,
pw: 0,
dh: 1,
dw: 1,
groups,
} if *groups == *c_in && *c_in == *c_out && *kw >= 1 && *w_out == 1 => {
(*src, *weight, *dst, *n, *c_in, *w, *w_out, *kw)
}
_ => {
i += 1;
continue;
}
};
let mut copy_in_i: Option<usize> = None;
let mut bcw = conv_src;
let mut cursor = i;
if let Ok(Some(idx)) = mlp_last_writer(thunks, cursor, bcw) {
match &thunks[idx] {
Thunk::Copy {
src,
dst,
len,
dt: HalfFlag::F32,
} if *dst == bcw && *len == batch.saturating_mul(c_in).saturating_mul(width) => {
copy_in_i = Some(idx);
bcw = *src;
cursor = idx;
}
Thunk::Nop => {}
_ => {}
}
}
let transpose_in_i = match mlp_last_writer(thunks, cursor, bcw) {
Ok(Some(idx)) => idx,
_ => {
i += 1;
continue;
}
};
let bsc_src = match &thunks[transpose_in_i] {
Thunk::Transpose {
src,
dst,
out_dims,
in_strides,
..
} if *dst == bcw
&& out_dims.len() == 3
&& in_strides.len() == 3
&& out_dims[0] == batch
&& out_dims[1] == c_in
&& out_dims[2] == width
&& in_strides[0] == width.saturating_mul(c_in)
&& in_strides[1] == 1
&& in_strides[2] == c_in =>
{
*src
}
_ => {
i += 1;
continue;
}
};
if output_offsets.contains(&bcw) || output_offsets.contains(&conv_src) {
i += 1;
continue;
}
let mut copy_out_i: Option<usize> = None;
let mut bcs = conv_dst;
let mut scan_from = i + 1;
if let Some(idx) = (scan_from..n_thunks).find(|&j| {
!matches!(thunks[j], Thunk::Nop)
&& matches!(
&thunks[j],
Thunk::Copy {
src,
dt: HalfFlag::F32,
..
} if *src == conv_dst
)
}) {
match &thunks[idx] {
Thunk::Copy { dst, len, .. }
if *len == batch.saturating_mul(c_in).saturating_mul(out_seq) =>
{
copy_out_i = Some(idx);
bcs = *dst;
scan_from = idx + 1;
}
_ => {}
}
}
let transpose_out_i = match (scan_from..n_thunks).find(|&j| {
!matches!(thunks[j], Thunk::Nop)
&& matches!(&thunks[j], Thunk::Transpose { src, .. } if *src == bcs)
}) {
Some(idx) => idx,
None => {
match (i + 1..n_thunks).find(|&j| {
!matches!(thunks[j], Thunk::Nop)
&& matches!(
&thunks[j],
Thunk::Transpose { src, .. } if *src == conv_dst
)
}) {
Some(idx) => {
bcs = conv_dst;
idx
}
None => {
i += 1;
continue;
}
}
}
};
let bsc_dst = match &thunks[transpose_out_i] {
Thunk::Transpose {
dst,
out_dims,
in_strides,
..
} if out_dims.len() == 3
&& in_strides.len() == 3
&& out_dims[0] == batch
&& out_dims[1] == out_seq
&& out_dims[2] == c_in
&& in_strides[0] == c_in.saturating_mul(out_seq)
&& in_strides[1] == 1
&& in_strides[2] == out_seq =>
{
*dst
}
_ => {
i += 1;
continue;
}
};
if output_offsets.contains(&bcs) && bcs != conv_dst {
i += 1;
continue;
}
if output_offsets.contains(&conv_dst) && copy_out_i.is_some() {
i += 1;
continue;
}
let mut silu_i: Option<usize> = None;
let mut silu_copy_i: Option<usize> = None;
let mut silu_dst = bsc_dst;
for j in transpose_out_i + 1..(transpose_out_i + 1 + 32).min(n_thunks) {
match &thunks[j] {
Thunk::Nop => continue,
Thunk::ActivationInPlace {
data,
act: Activation::Silu,
..
} if *data == silu_dst => {
silu_i = Some(j);
break;
}
Thunk::ActivationOut {
src,
dst,
act: Activation::Silu,
..
} if *src == silu_dst => {
silu_dst = *dst;
silu_i = Some(j);
break;
}
Thunk::Copy {
src,
dst,
len,
dt: HalfFlag::F32,
} if *src == bsc_dst
&& silu_copy_i.is_none()
&& *len == batch.saturating_mul(out_seq).saturating_mul(c_in) =>
{
silu_copy_i = Some(j);
silu_dst = *dst;
}
Thunk::Narrow { .. } | Thunk::SplitLastAxis { .. } => continue,
_ if silu_copy_i.is_none() => continue,
_ => break,
}
}
let do_silu = silu_i.is_some();
let final_dst = if do_silu { silu_dst } else { bsc_dst };
if width != out_seq.saturating_add(kw).saturating_sub(1) {
i += 1;
continue;
}
let until = silu_i.unwrap_or(transpose_out_i) + 1;
let mut allowed_bcw = vec![i];
if let Some(c) = copy_in_i {
allowed_bcw.push(c);
}
if !mlp_value_dead_in_range(thunks, transpose_in_i, bcw, &allowed_bcw, until) {
i += 1;
continue;
}
if let Some(c) = copy_in_i {
if !mlp_value_dead_in_range(thunks, c, conv_src, &[i], until) {
i += 1;
continue;
}
}
let mut allowed_conv = Vec::new();
if let Some(c) = copy_out_i {
allowed_conv.push(c);
} else {
allowed_conv.push(transpose_out_i);
}
if !mlp_value_dead_in_range(thunks, i, conv_dst, &allowed_conv, until) {
i += 1;
continue;
}
if let Some(c) = copy_out_i {
if !mlp_value_dead_in_range(thunks, c, bcs, &[transpose_out_i], until) {
i += 1;
continue;
}
}
thunks[transpose_in_i] = Thunk::Nop;
if let Some(c) = copy_in_i {
thunks[c] = Thunk::Nop;
}
if let Some(c) = copy_out_i {
thunks[c] = Thunk::Nop;
}
thunks[transpose_out_i] = Thunk::Nop;
if let Some(c) = silu_copy_i {
thunks[c] = Thunk::Nop;
}
if let Some(si) = silu_i {
thunks[si] = Thunk::Nop;
}
thunks[i] = Thunk::FusedDepthwiseConv1dBsc {
src: bsc_src,
weight,
dst: final_dst,
batch,
width,
out_seq,
channels: c_in,
k: kw,
silu: do_silu,
};
fused += 1;
if verbose {
eprintln!(
"[rlx-metal] fuse_depthwise_conv1d_bsc: B={batch} W={width} out={out_seq} \
C={c_in} k={kw} silu={do_silu} — chain → 1"
);
}
i += 1;
}
if verbose && fused > 0 {
eprintln!("[rlx-metal] fuse_depthwise_conv1d_bsc: {fused} blocks fused");
}
}
fn fuse_narrow_clusters(thunks: &mut [Thunk]) {
use std::collections::HashMap;
#[derive(Hash, PartialEq, Eq, Clone, Copy)]
struct NarrowKey {
src: usize,
outer: u32,
src_axis: u32,
dt: u8,
}
let mut groups: HashMap<NarrowKey, Vec<(usize, usize, u32, u32)>> = HashMap::new();
for (i, t) in thunks.iter().enumerate() {
let Thunk::Narrow {
src,
dst,
outer,
src_axis,
start,
len,
dt,
} = t
else {
continue;
};
let key = NarrowKey {
src: *src,
outer: *outer,
src_axis: *src_axis,
dt: match dt {
HalfFlag::F32 => 0,
HalfFlag::F16 => 1,
},
};
groups.entry(key).or_default().push((i, *dst, *start, *len));
}
let mut groups_fused = 0usize;
let mut narrows_fused = 0usize;
for (key, mut items) in groups {
if items.len() < 2 || key.dt != 0 {
continue;
}
let meta: Vec<(u32, u32)> = items.iter().map(|(_, _, s, l)| (*s, *l)).collect();
if !narrow_segments_partition(key.src_axis, &meta) {
continue;
}
items.sort_by_key(|(i, _, _, _)| *i);
let dt = HalfFlag::F32;
let segments: Vec<(usize, u32, u32)> =
items.iter().map(|(_, d, s, l)| (*d, *s, *l)).collect();
let n = items.len();
let first = items[0].0;
thunks[first] = Thunk::SplitLastAxis {
src: key.src,
outer: key.outer,
src_axis: key.src_axis,
dt,
segments,
};
for (i, _, _, _) in items.into_iter().skip(1) {
thunks[i] = Thunk::Nop;
}
groups_fused += 1;
narrows_fused += n;
}
let verbose = rlx_ir::env::var("RLX_VERBOSE")
.and_then(|v| v.parse::<u8>().ok())
.unwrap_or(0)
>= 1;
if verbose && groups_fused > 0 {
eprintln!(
"[rlx-metal] fuse_narrow_clusters: {groups_fused} split groups ({narrows_fused} narrows merged)"
);
}
}
fn metal_thunk_read_offsets(t: &Thunk) -> Vec<usize> {
match t {
Thunk::Sgemm { a, b, .. } => vec![*a, *b],
Thunk::BatchedSgemm { a, b, .. } => vec![*a, *b],
Thunk::FusedMmBiasAct { a, w, bias, .. } => vec![*a, *w, *bias],
Thunk::BinaryFull { lhs, rhs, .. } => vec![*lhs, *rhs],
Thunk::FusedBinaryActivation { lhs, rhs, .. } => vec![*lhs, *rhs],
Thunk::FusedTernaryActivation {
lhs, rhs0, rhs1, ..
} => vec![*lhs, *rhs0, *rhs1],
Thunk::BinaryBroadcast { lhs, rhs, .. } => vec![*lhs, *rhs],
Thunk::ActivationInPlace { data, .. } => vec![*data],
Thunk::ActivationOut { src, .. } => vec![*src],
Thunk::GeluApproxOut { src, .. } | Thunk::GeluApproxHost { src, .. } => vec![*src],
Thunk::LayerNorm { src, g, b, .. } | Thunk::GroupNorm { src, g, b, .. } => {
vec![*src, *g, *b]
}
Thunk::ResizeNearest2x { src, .. } => vec![*src],
Thunk::RmsNorm { src, g, b, .. } => vec![*src, *g, *b],
Thunk::FusedResidualLN {
x, res, bias, g, b, ..
} => vec![*x, *res, *bias, *g, *b],
Thunk::FusedResidualRmsNorm {
x, res, bias, g, b, ..
} => vec![*x, *res, *bias, *g, *b],
Thunk::AdaLayerNorm {
x, scale, shift, ..
} => vec![*x, *scale, *shift],
Thunk::GatedResidual { x, y, gate, .. } => vec![*x, *y, *gate],
Thunk::AdaLayerNormBackward { x, scale, dy, .. } => vec![*x, *scale, *dy],
Thunk::GatedResidualBackward { y, gate, dy, .. } => vec![*y, *gate, *dy],
Thunk::FusedRmsNormMulSilu { x, g, b, z, .. } => vec![*x, *g, *b, *z],
Thunk::FusedDepthwiseConv1dBsc { src, weight, .. } => vec![*src, *weight],
Thunk::Softmax { data, .. } => vec![*data],
Thunk::SoftmaxCrossEntropyDense {
logits, targets, ..
} => vec![*logits, *targets],
Thunk::SoftmaxCrossEntropyWithLogits { logits, labels, .. } => vec![*logits, *labels],
Thunk::SoftmaxCrossEntropyBackward {
logits,
labels,
d_loss,
..
} => vec![*logits, *labels, *d_loss],
Thunk::Cumsum { src, .. } => vec![*src],
Thunk::SpdHost { inputs, .. } => inputs.iter().map(|(off, _, _)| *off).collect(),
Thunk::CustomGpuOp { inputs, .. } => inputs.iter().map(|(off, _, _)| *off).collect(),
Thunk::Attention { q, k, v, mask, .. } => vec![*q, *k, *v, *mask],
Thunk::FusedAttn {
qkv,
mask,
cos,
sin,
has_rope,
..
} => {
let mut r = vec![*qkv, *mask];
if *has_rope != 0 {
r.push(*cos);
r.push(*sin);
}
r
}
Thunk::AttentionBackward {
q, k, v, dy, mask, ..
} => {
let mut v = vec![*q, *k, *v, *dy];
if *mask != *q {
v.push(*mask);
}
v
}
Thunk::Rope { src, cos, sin, .. } => vec![*src, *cos, *sin],
Thunk::RmsNormBackwardInput {
x, gamma, beta, dy, ..
} => {
vec![*x, *gamma, *beta, *dy]
}
Thunk::RmsNormBackwardGamma {
x, gamma, beta, dy, ..
} => {
vec![*x, *gamma, *beta, *dy]
}
Thunk::RmsNormBackwardBeta {
x, gamma, beta, dy, ..
} => {
vec![*x, *gamma, *beta, *dy]
}
Thunk::RopeBackward { dy, cos, sin, .. } => vec![*dy, *cos, *sin],
Thunk::CumsumBackward { dy, .. } => vec![*dy],
Thunk::GatherBackward { dy, indices, .. } => vec![*dy, *indices],
Thunk::MaxPool2dBackward { x, dy, .. } => vec![*x, *dy],
Thunk::Conv2dBackwardInput { dy, w, .. } => vec![*dy, *w],
Thunk::Conv2dBackwardWeight { x, dy, .. } => vec![*x, *dy],
Thunk::FusedSwiGLU { src, .. } => vec![*src],
Thunk::FusedMlpGateUpSwiGLU {
x, gate_w, up_w, ..
}
| Thunk::FusedMlpGateUpGelu {
x, gate_w, up_w, ..
} => vec![*x, *gate_w, *up_w],
Thunk::FusedMlpDownResidual { x, w, res, .. } => vec![*x, *w, *res],
Thunk::Concat { inputs, .. } => inputs.iter().map(|(o, _)| *o).collect(),
Thunk::Narrow { src, .. } | Thunk::SplitLastAxis { src, .. } => vec![*src],
Thunk::Copy { src, .. } => vec![*src],
_ => vec![],
}
}
fn concat_axis_extent(input: &rlx_ir::Shape, axis: usize, out_rank: usize) -> usize {
let in_rank = input.rank();
if axis >= out_rank {
return 1;
}
if axis < in_rank {
input.dim(axis).unwrap_static()
} else {
1
}
}
#[cfg(test)]
mod region_rewrite_tests {
use super::*;
use rlx_ir::op::{Activation, BinaryOp};
fn empty_modulus() -> [u32; 16] {
[0; 16]
}
fn region(
len: u32,
n_in: u32,
num_steps: u32,
dst: usize,
input_offs: [u32; 16],
chain: [u32; 128],
) -> Thunk {
Thunk::ElementwiseRegion {
len,
num_inputs: n_in,
num_steps,
dst,
input_offs,
chain,
scalar_input_mask: 0,
input_modulus: empty_modulus(),
prologue: 0,
out_n: 0,
out_c: 0,
out_h: 0,
out_w: 0,
prologue_input: 0,
}
}
#[test]
fn rewrite_single_binary_to_binary_full() {
let mut chain = [0u32; 128];
chain[0] = 2;
chain[1] = 0; chain[2] = 0;
chain[3] = 1;
let t = region(
128,
2,
1,
4096,
[256, 512, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
chain,
);
match try_rewrite_elementwise_region(&t) {
RegionRewrite::One(Thunk::BinaryFull { op, len, .. }) => {
assert_eq!(op, BinaryOp::Add);
assert_eq!(len, 128);
}
_ => panic!("expected BinaryFull"),
}
}
#[test]
fn rewrite_binary_then_activation_to_fused() {
let mut chain = [0u32; 128];
chain[0] = 2;
chain[1] = 2; chain[2] = 0;
chain[3] = 1;
chain[4] = 0;
chain[5] = 2; chain[6] = 0x8000_0000;
let t = region(
64,
2,
2,
8192,
[128, 256, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0],
chain,
);
match try_rewrite_elementwise_region(&t) {
RegionRewrite::One(Thunk::FusedBinaryActivation { op, act, .. }) => {
assert_eq!(op, BinaryOp::Mul);
assert_eq!(act, Activation::Silu);
}
_ => panic!("expected fused binary+activation"),
}
}
#[test]
fn rewrite_binary_then_binary_to_pair() {
let mut chain = [0u32; 128];
chain[0] = 2;
chain[1] = 0; chain[2] = 0;
chain[3] = 1;
chain[4] = 2;
chain[5] = 2; chain[6] = 0x8000_0000;
chain[7] = 2;
let offs = [128, 256, 384, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0];
let t = region(32, 3, 2, 4096, offs, chain);
match try_rewrite_elementwise_region(&t) {
RegionRewrite::Many(ts) if ts.len() == 2 => {
assert!(matches!(
ts[0],
Thunk::BinaryFull {
op: BinaryOp::Add,
..
}
));
assert!(matches!(
ts[1],
Thunk::BinaryFull {
op: BinaryOp::Mul,
..
}
));
}
_ => panic!("expected binary+binary pair"),
}
}
#[test]
fn rewrite_binary_binary_activation_to_fused_ternary() {
let mut chain = [0u32; 128];
chain[0] = 2;
chain[1] = 0; chain[2] = 0;
chain[3] = 1;
chain[4] = 2;
chain[5] = 2; chain[6] = 0x8000_0000;
chain[7] = 2;
chain[8] = 0;
chain[9] = 2; chain[10] = 0x8000_0001;
let offs = [128, 256, 384, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0];
let t = region(32, 3, 3, 4096, offs, chain);
match try_rewrite_elementwise_region(&t) {
RegionRewrite::One(Thunk::FusedTernaryActivation { op0, op1, act, .. }) => {
assert_eq!(op0, BinaryOp::Add);
assert_eq!(op1, BinaryOp::Mul);
assert_eq!(act, Activation::Silu);
}
_ => panic!("expected fused ternary+activation"),
}
}
}