Fused Triton kernels for normalisation layers. Drop-in replacements for PyTorch's norm modules with automatic dispatch to optimised CUDA kernels when triton is available.
pip install torchnorm # PyTorch fallback only
pip install "torchnorm[triton]" # with Triton fused kernelsFrom source:
git clone https://github.com/liodon-ai/torchnorm
cd torchnorm
pip install -e ".[triton]"Root Mean Square Layer Normalisation (Zhang & Sennrich, 2019). Used in LLaMA, Mistral, Gemma, Qwen, Falcon, and most modern LLMs.
from torchnorm import RMSNorm
norm = RMSNorm(dim=4096).cuda()
x = torch.randn(4, 512, 4096, device="cuda", dtype=torch.float16)
out = norm(x) # fused Triton kernel on CUDA, PyTorch fallback elsewhereThe fused kernel reads each element once, computes the RMS, and applies the weight in a single pass — no intermediate allocation.
Fuses the residual add and RMSNorm into one kernel. Canonical pattern in transformer pre-norm blocks:
from torchnorm import FusedAddRMSNorm
norm = FusedAddRMSNorm(dim=4096).cuda()
# Instead of:
# residual = x + residual
# hidden = rms_norm(residual)
# Write:
hidden, residual = norm(x, residual)Returns both the normalised output and the updated residual so the next block can use it directly.
Layer Normalisation (Ba et al., 2016). Identical interface to torch.nn.LayerNorm.
from torchnorm import LayerNorm
norm = LayerNorm(768, bias=False).cuda() # BERT-style, no bias
out = norm(x)Group Normalisation (Wu & He, 2018). Standard for diffusion UNets and CNNs.
from torchnorm import GroupNorm
norm = GroupNorm(num_groups=32, num_channels=512).cuda()
out = norm(x) # x: (N, C, H, W)Instance Normalisation (Ulyanov et al., 2016). Normalises each sample–channel pair over spatial dims. Standard in style transfer.
from torchnorm import InstanceNorm
norm = InstanceNorm(num_features=512, affine=True).cuda()
out = norm(x) # x: (N, C, H, W) or (N, C)ScaleNorm (Nguyen & Salazar, 2019). L2-normalises and scales by a single learnable scalar. Cheaper than LayerNorm — no mean subtraction, no per-element parameters.
from torchnorm import ScaleNorm
norm = ScaleNorm(dim=512).cuda()
out = norm(x)Measured on NVIDIA GB10 (Grace Hopper) · PyTorch 2.11 · Triton 3.6 · bfloat16 · B=16 seq=1024 · warmup=60 reps=300.
All speedups are relative to vanilla F.rms_norm / unfused x + residual → rms_norm(x).
| Hidden dim | Vanilla (ms) | torch.compile (ms) | torchnorm (ms) | compile × | torchnorm × |
|---|---|---|---|---|---|
| 512 | 0.987 | 0.977 | 1.150 | 1.01× | 0.86× |
| 1024 | 2.232 | 2.263 | 2.793 | 0.99× | 0.80× |
| 2048 | 4.177 | 4.575 | 4.565 | 0.91× | 0.92× |
| 4096 | 9.494 | 8.842 | 8.844 | 1.07× | 1.07× |
| 8192 | 22.252 | 17.130 | 17.057 | 1.30× | 1.30× |
At D ≥ 4096, torchnorm matches torch.compile exactly and both give ~1.30× over vanilla. torchnorm runs on the first call with no compilation step.
| Hidden dim | Vanilla (ms) | torch.compile (ms) | torchnorm (ms) | compile × | torchnorm × |
|---|---|---|---|---|---|
| 512 | 0.716 | 0.464 | 0.646 | 1.54× | 1.11× |
| 1024 | 1.586 | 0.817 | 1.211 | 1.94× | 1.31× |
| 2048 | 2.980 | 1.877 | 2.548 | 1.59× | 1.17× |
| 4096 | 6.259 | 3.728 | 5.091 | 1.68× | 1.23× |
| 8192 | 12.742 | 7.122 | 9.447 | 1.79× | 1.35× |
torch.compile wins here by fusing the add and norm at the compiler level. torchnorm improvement planned for v0.1.3.
To reproduce:
python benchmarks/bench_torchnorm.py # kernel micro-benchmarks
python benchmarks/bench_compile.py # vs torch.compile comparisonEvery module auto-selects the fastest available path:
| Condition | Path |
|---|---|
| CUDA tensor + triton installed | Fused Triton kernel |
| CPU tensor or triton not installed | PyTorch (F.rms_norm, F.layer_norm, …) |
No flags to set. Installing triton is the only opt-in required.
| Python | ≥ 3.9 |
| PyTorch | ≥ 2.0 |
| Triton | ≥ 2.1 (optional) |
| CUDA | any version supported by your PyTorch build |
| Project | Focus | Install |
|---|---|---|
| Liger Kernel | Full training suite: RMSNorm + RoPE + SwiGLU + CrossEntropy + model patching | pip install liger-kernel |
| flash-attn | Flash Attention + fused ops including LayerNorm | build from source |
| torchembed | Companion library: RoPE (rotate-half & adjacent-pairs), ALiBi, patch, categorical embeddings | pip install torchembed |
torchnorm vs Liger Kernel: Liger patches entire model classes at import time — ideal when you want the full optimization suite with use_liger_kernel=True. torchnorm provides individual nn.Module drop-ins without model patching — useful when building custom architectures or when you only need the norm layers.
- torchembed — fused kernels for embedding layers (RoPE, patch, categorical)
- torchnorm — fused kernels for normalisation layers