--- license: apache-2.0 library_name: kernels tags: - kernel - xpu - sycl - intel - rmsnorm - rope - decode --- # decode-ops-xpu Fused decode-step glue for LLM inference on Intel GPUs, loadable through `kernels`: RMSNorm with a fused residual add, in-place RoPE, SiLU-gate, and greedy argmax. The reference baselines are the eager PyTorch equivalents, matched to 1e-6 or bit-exactly. This is the non-matmul residue that owns a decode step once attention and the linear layers are fast. Left in eager PyTorch it is several launches and several full passes over the hidden state per layer, and on an integrated GPU those passes are the step. Each op here is one kernel that touches the data once. ![Per-op bars where the fused kernel runs a fraction of eager PyTorch on an Intel GPU](https://huggingface.co/kernels/phanerozoic/decode-ops-xpu/resolve/main/media/hero.gif) *Measured live on an Intel Iris Xe: RMSNorm at 0.037 ms against eager's 0.100 (2.7x), argmax over a 128,256 vocabulary at 0.246 against 0.355 (1.4x), matching eager to 1e-6 and bit-exactly for argmax.* ## Usage ```python from kernels import get_kernel d = get_kernel("phanerozoic/decode-ops-xpu", version=1, trust_remote_code=True) h = d.rms_norm(hidden, norm_weight, 1e-6) d.fused_add_rms_norm(hidden, residual, norm_weight, 1e-6) d.rope_inplace(q, k, cos, sin) gated = d.silu_and_mul(mlp_out) tokens = d.argmax(logits) ``` `layers.RMSNorm` and `layers.SiluAndMul` are exported for use as Hub layer replacements. `version` selects the release branch; `trust_remote_code` is required by `kernels` for publishers without the trusted-publisher mark. ## API | Symbol | Purpose | |---|---| | `rms_norm(x, weight, eps)` | RMS-normalize the last dim, scale by `weight` | | `fused_add_rms_norm(x, residual, weight, eps)` | `residual += x`, then `x = rms_norm(residual) * weight`, both in place | | `rope_inplace(query, key, cos, sin)` | NeoX-style rotary on both tensors in place, grouped-query aware | | `silu_and_mul(x)` | `silu(x[..., :n]) * x[..., n:]` | | `argmax(logits)` | greedy token selection, ties resolve to the lowest index | `rope_inplace` takes `query` as `[tokens, n_q_heads, head_dim]` and `key` as `[tokens, n_kv_heads, head_dim]` with differing head counts, and rotates the pair `(i, i + head_dim // 2)` using `cos`/`sin` of shape `[tokens, head_dim // 2]`. ## Method Normalization and argmax use one work-group per row with a group reduction, so a row of any length is handled without scratch memory. The normalization kernels hold a row in registers between the sum-of-squares pass and the scaling pass, so it is read from memory once; rows wider than the register budget fall back to re-reading. SiLU-gate moves four floats per access where the width allows it. Argmax carries the index alongside the value so the row is read once. ## Measured Against the eager PyTorch equivalent on Intel Iris Xe (Gen12LP, 96 EUs), torch 2.13.0+xpu, 30 iterations after 5 warmup: | op | shape | kernel | eager | ratio | |---|---|---|---|---| | `rms_norm` | [8, 4096] | 0.037 ms | 0.093 ms | 2.5x | | `silu_and_mul` | [8, 22016] | 0.023 ms | 0.043 ms | 1.9x | | `argmax` | [8, 128256] | 0.246 ms | 0.353 ms | 1.4x | These are memory-bound, so the gains come from traffic rather than arithmetic. The table is reproduced by `benchmarks/bench.py`, which loads this kernel from the Hub and times it against the eager equivalent on the local device. ## Validation Thirteen checks against PyTorch references on Intel Iris Xe, torch 2.13.0+xpu. RMSNorm matches at 4.77e-07 to 9.54e-07 across shapes `(4, 512)`, `(2, 7, 1024)`, `(1, 4096)`, and `(3, 333)`, the last covering a row that is not a multiple of the work-group size. The residual output of `fused_add_rms_norm` is bit-exact and its normalized output matches at 4.77e-07. RoPE matches at 2.38e-07 for both query and a grouped-query key with a different head count. SiLU-gate matches at 4.77e-07. Argmax is bit-exact against `torch.argmax` at vocabularies of 1000, 32000, and 7, and a deliberate two-way tie is verified to resolve to the lower index. ## Requirements and limits - float32 or float16; reductions run in float32 regardless, so the float16 path keeps float32 accumulation accuracy. float16 requires `aspect::fp16`, which is checked before dispatch. - Intel GPU with a working Level Zero driver and a torch XPU build. - The ops ship fake implementations, so they trace under `torch.compile` without breaking the graph; each is verified to compile into a single frame and to match eager output. ## References Zhang and Sennrich, "Root Mean Square Layer Normalization" (2019); Su et al., "RoFormer" (2021); the NeoX rotary layout. ## License Apache-2.0.