This document is relevant for: Trn2, Trn3

Kernel implementations#

Some of the performance in this plugin comes from NKI kernels that are vendored into the repository and modified for a specific model, rather than imported unchanged from the NKI Library. The pages in this section document those kernels: what each one computes, the design decisions behind its shape, its calling interface, and what changes when you map the same pattern onto a different model.

Kernel reference pages#

Const-max ring attention

Ring attention with a constant softmax max. Replaces softmax’s online row-max with a provable per-row bound, which frees the score matmul to stay K-stationary so the ring’s K/V rotation overlaps attention compute.

Fused adaptive LayerNorm and FP8 quantization

Fused adaptive LayerNorm. Collapses a gated residual add, a (Layer|RMS|no-)norm, a per-hidden affine/AdaLN modulation, and optional per-row FP8 quantization into one launch, hoisting the loop-invariant modulation broadcast out of the token loop.

QKV projection with QK Distributed RMSNorm and RoPE fusion

MXFP8 QKV projection kernel with an across-heads (Distributed) RMSNorm and in-kernel all-reduce, a bundled rotate-half RoPE, plus PSUM I-group tiling and graduated weight prefetch.

Wan2.2 vendors other kernels — the MLP variants and the distributed VAE components — that are not documented here yet.

When a kernel belongs here#

These pages are for kernels you are expected to read and adapt. If you only want to call an operation as it ships, the NKI Library’s API reference is the right place, and adapting a vendored copy is the wrong move — you would be forking away from upstream fixes for no benefit.

A page in this section exists because the kernel embodies a design decision worth reusing. It is written to be useful even if you never run the kernel itself.