Skip to content

Latest commit

 

History

9 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

torchnorm

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.

Installation

pip install torchnorm            # PyTorch fallback only
pip install "torchnorm[triton]"  # with Triton fused kernels

From source:

git clone https://github.com/liodon-ai/torchnorm
cd torchnorm
pip install -e ".[triton]"

Modules

RMSNorm

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 elsewhere

The fused kernel reads each element once, computes the RMS, and applies the weight in a single pass — no intermediate allocation.

FusedAddRMSNorm

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.

LayerNorm

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)

GroupNorm

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)

InstanceNorm

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

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)

Benchmarks

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).

RMSNorm — forward + backward

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.

FusedAddRMSNorm — residual add + norm

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 comparison

Dispatch logic

Every 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.

Compatibility

Python ≥ 3.9
PyTorch ≥ 2.0
Triton ≥ 2.1 (optional)
CUDA any version supported by your PyTorch build

Related projects

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.

Part of the Liodon AI stack

  • torchembed — fused kernels for embedding layers (RoPE, patch, categorical)
  • torchnorm — fused kernels for normalisation layers

About

Fused Triton kernels for normalisation layers (RMSNorm, LayerNorm, GroupNorm, FusedAddRMSNorm)

Topics

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages