const ARENA_LARGE_OFF: usize = 1usize << 32;
#[inline]
fn arena_off_large(off: usize) -> bool {
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 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::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,
},
Sgemm {
a: usize,
b: usize,
c: usize,
m: u32,
k: u32,
n: u32,
dt: HalfFlag,
},
BatchedSgemm {
a: usize,
b: usize,
c: usize,
batch: u32,
m: u32,
k: u32,
n: u32,
dt: HalfFlag,
},
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,
},
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,
},
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,
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,
},
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,
},
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,
},
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,
},
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,
},
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>,
},
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 {
plan: std::sync::Arc<rlx_cpu::thunk::ScanBodyPlan>,
outer_init_off: usize,
outer_final_off: usize,
length: u32,
save_trajectory: bool,
xs_outer: Vec<(usize, usize)>,
bcast_outer: Vec<(usize, usize)>,
},
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::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::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::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::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::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::Copy { .. }
| Thunk::ActivationInPlace { .. }
| 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::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)
}
fn mlp_io(t: &Thunk) -> Option<(Vec<usize>, Vec<usize>)> {
use Thunk::*;
let io = match t {
Nop => (vec![], vec![]),
Cast { src, dst, .. } => (vec![*src], vec![*dst]),
Copy { src, dst, .. } => (vec![*src], vec![*dst]),
ActivationInPlace { data, .. } => (vec![*data], vec![*data]),
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]),
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,
_ => 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,
..
}
)
};
let is_gelu = |t: &Thunk| {
matches!(
t,
Thunk::ActivationInPlace {
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::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();
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,
}
};
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,
};
}
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) {
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,
..
}
)
};
let is_gelu = |t: &Thunk| {
matches!(
t,
Thunk::ActivationInPlace {
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::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();
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| 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 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) = match &thunks[add_idx] {
Thunk::BinaryFull { lhs, rhs, dst, .. } => {
let res = if *lhs == down_dst { *rhs } else { *lhs };
(res, *dst)
}
_ => 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,
}
};
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,
};
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_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::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::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::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"),
}
}
}