Downloads · 30 days
0
axjns/strix-halo-kernels
strix-halo-kernels is a machine learning model from axjns. Use it for the machine learning task on the model card, and read the license before you ship it in a product. The card lists the license as apache-2.0.
Fused Triton kernels for AMD RDNA3.5 / gfx1151 — Strix Halo, the Radeon 8050S/8060S iGPU in Ryzen AI Max parts.
Downloads · 30 days
0
Access
Public
Updated Aug 17, 2026
Repo size
—
Likes
1
Public
Click a slice to open those files.
.py26 KB · 74%
From the Hugging Face model README
Fused Triton kernels for AMD RDNA3.5 / gfx1151 — Strix Halo, the Radeon 8050S/8060S iGPU in Ryzen AI Max parts.
Strix Halo inverts the usual trade-off: up to 64 GiB addressable as VRAM out of 128 GB unified, but far lower memory bandwidth than a discrete card. I measured a ~21× gap between prefill and decode throughput on this part (benchmarks), which is the signature of a bandwidth-bound machine.
That changes which optimizations matter. Fusing away memory round-trips is worth more here than on hardware where compute is the constraint — so these kernels target the ops that are pure memory traffic in eager PyTorch.
Also, practically: the mainstream Hub kernels are CUDA-only. kernels-community/activation
offers 41 build variants and none of them resolve on ROCm — every one is rejected with
backend (cu128) does not match selected backend (rocm713). Triton sidesteps this
entirely by JIT-compiling for whatever GPU is present.
fp16, versus eager PyTorch, via triton.testing.do_bench:
| kernel | shape | eager | this kernel | speedup |
|---|---|---|---|---|
| RMSNorm | 512×4096 | 0.205 ms | 0.072 ms | 2.84× |
| RMSNorm | 2048×4096 | 1.592 ms | 0.278 ms | 5.73× |
| RMSNorm | 4096×5120 | 3.703 ms | 0.508 ms | 7.29× |
| RMSNorm | 8192×2048 | 3.028 ms | 0.387 ms | 7.82× |
| GEGLU | 512×8192 | 0.233 ms | 0.197 ms | 1.18× |
| GEGLU | 2048×8192 | 0.843 ms | 0.554 ms | 1.52× |
| GEGLU | 4096×4096 | 0.838 ms | 0.550 ms | 1.52× |
RMSNorm gains the most, and the gain grows with size — eager does several full passes over the activation, this does one. GEGLU's ceiling is lower because even the fused form must read two inputs and write one output; the win is only the eliminated intermediate.
PyTorch on gfx1151 reports no flash or memory-efficient SDPA backend by default, so
scaled_dot_product_attention falls back to math, which materializes the full S×S score
matrix. But there is an AOTriton path — it is gated behind an environment variable:
export TORCH_ROCM_AOTRITON_ENABLE_EXPERIMENTAL=1
If you run Strix Halo and change one thing after reading this, change that. Enabling it takes attention at (1,32,4096,128) from 204 ms to 12 ms — a ~17× win, for free, before any custom kernel is involved.
Against that properly enabled baseline, this kernel is still faster (fp16, causal):
| shape | default (math) | AOTriton flash | this kernel | vs AOTriton |
|---|---|---|---|---|
| (1,32,512,128) | 3.95 ms | 0.530 ms | 0.162 ms | 3.27× |
| (1,32,4096,128) | 204.3 ms | 12.14 ms | 5.84 ms | 2.08× |
| (2,24,4096,64) | 304.2 ms | 11.97 ms | 4.00 ms | 3.00× |
| (1,16,8192,64) | 342.9 ms | 15.43 ms | 6.06 ms | 2.55× |
Peak allocation at (1,32,4096,128): 5056 MiB → 160 MiB, because the S×S matrix is never written.
The honest summary: ~2–3× over the vendor path once it is enabled, and ~25–68× over what you get if you do not know the env var exists. Quote the first number.
Block sizes are chosen by triton.autotune per (seq_len, head_dim, causal); on gfx1151 it
picks BLOCK_M=32, BLOCK_N=64, num_warps=4, num_stages=2.
import torch.nn.functional as F
k = get_kernel("axjns/strix-halo-kernels", version=1)
F.scaled_dot_product_attention = k.sdpa # safe: falls back for unsupported inputs
sdpa defers to torch for explicit attn_mask, nonzero dropout_p, GQA/MQA, non-4-D
inputs, non-power-of-two head_dim, fp32, and CPU tensors — it never silently returns
wrong numbers for a case the kernel cannot handle.
Correctness: 24 checks against an fp32 reference (not the fp16 baseline) across causal and non-causal, fp16 and bf16, head_dim 64 and 128, and sequence lengths including non-multiples of the block size (77, 333). Max error 9.8e-4 (fp16), 8.4e-3 (bf16).
Attention is forward-only too — no backward pass.
import torch
from kernels import get_kernel
k = get_kernel("axjns/strix-halo-kernels", version=1)
# fused RMSNorm over the last dim
y = k.rms_norm(x, weight, eps=1e-6)
# gated activations
h = k.swiglu(gate, up) # silu(gate) * up — LLaMA/Qwen/Mistral MLPs
h = k.geglu(gate, up) # gelu(gate) * up — Flux / SD3 / T5 MLPs
h = k.geglu_chunked(proj_out) # splits (..., 2*d) then gates
layers provides nn.Module drop-ins that keep the original parameter names, so no
weight surgery or state-dict changes are needed:
from kernels import kernelize # swaps compatible layers in-place
layers.RMSNorm — reads self.weight and accepts either variance_epsilon or
eps, since transformers has used both names across versions.layers.SwiGLU — expects gate_proj / up_proj / down_proj (LLaMA-family MLP).layers.GEGLU — expects a proj emitting 2 * inner_dim, then chunks and gates.Validated against eager PyTorch references (also exported, as rms_norm_ref,
swiglu_ref, geglu_ref) across fp32 / fp16 / bf16, 2-D and 3-D inputs,
non-contiguous inputs, and shapes from 4×512 to 4096×5120. Sum-of-squares is
accumulated in fp32 regardless of input dtype — fp16 accumulation over a
4096-wide row loses precision badly.
Max observed absolute error: 1.9e-6 (fp32), 7.8e-3 (fp16), 6.3e-2 (bf16) — consistent with the dtype's own rounding.
torch.autograd-opaque, so
gradients will not flow through them. Inference and eval only; do not put these in a
training graph expecting it to work.triton.autotune and will re-tune on other GPUs.
RMSNorm and the gated activations use hand-picked block sizes and num_warps, chosen for
RDNA's 32-wide wavefront and an APU's low CU count — a discrete GPU would likely want
different values, and those two are not yet autotuned.python test_bench.py # RMSNorm + gated activations
python test_flash.py # attention correctness + headroom
python sdpa_probe.py # which SDPA backends your build actually has
sdpa_probe.py is worth running on any ROCm box before you optimize anything — it tells
you whether you are silently on the math fallback.