linkedin / linkedin/Liger-Kernel

Add fused Modulated RMSNorm for DiT-style scale/shift conditioning

Open
#1,224 5 comments 0 reactions 0 assignees View on GitHub

Nobody has claimed this yet.

Dominant language
Python
Stars
6.6k
Forks
603
Avg merge
1d 20h
Merged PRs (30d)
47

Description

## The feature, motivation

I would like to add a fused Modulated RMSNorm kernel for AdaLN/FiLM-style conditioning patterns used in diffusion Transformers.

The operation is:

```python
y = rms_norm(x, weight, eps) * (1 + scale) + shift
```

with `shift` optional:

```python
y = rms_norm(x, weight, eps) * (1 + scale)
```

This pattern shows up in DiT-style blocks where a conditioning vector produces per-layer modulation parameters. In current PyTorch-style implementations, RMSNorm and the scale/shift modulation are usually separate ops, which means extra kernel launches and an intermediate normalized tensor.

Liger already has a fast `RMSNorm` kernel, so this seems like a small and useful extension for diffusion/video Transformer workloads.

Proposed API:

```python
liger_modulated_rms_norm(
X,
W,
scale,
shift=None,
eps=1e-6,
offset=0.0,
casting_mode="llama",
in_place=True,
)
```

and a module wrapper:

```python
LigerModulatedRMSNorm(hidden_size, eps=1e-6, ...)
```

Initial scope:

- Fuse RMSNorm + modulation scale + optional shift
- Support `scale/shift` shaped per-row or per-batch, broadcast over tokens
- Support `W=None` / `elementwise_affine=False`
- Keep existing RMSNorm options where possible: `offset`, `casting_mode`, `in_place`
- Add correctness tests for forward and backward against a PyTorch reference
- Add a benchmark comparing:
- PyTorch/HF RMSNorm + modulation
- `LigerRMSNorm` + modulation
- fused `LigerModulatedRMSNorm`

Out of scope for the first PR:

- Fusing the modulation MLP that produces `scale` and `shift`
- Fusing residual gates
- Diffusers monkey-patching
- LayerNorm + modulation

## Alternatives

The current alternative is to use `LigerRMSNorm` and then apply scale/shift with regular PyTorch ops. That works, but still materializes the normalized output and launches extra elementwise kernels.

## Additional context

This is a common pattern in recent diffusion Transformer architectures:

- FiLM introduced feature-wise affine modulation:
https://ojs.aaai.org/index.php/AAAI/article/view/11671
- DiT uses adaptive layer norm / adaLN-Zero conditioning:
https://www.wpeebles.com/DiT
https://github.com/facebookresearch/DiT/blob/main/models.py
- FLUX uses per-block `shift`, `scale`, and `gate` modulation:
https://github.com/black-forest-labs/flux/blob/main/src/flux/modules/layers.py
- Diffusers supports adaptive normalization with RMSNorm variants:
https://huggingface.co/docs/diffusers/api/normalization
https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py

Contributor guide

Open the contributing guide

First steps

  1. Read the whole issue, then the project's contributing guide.
  2. Comment on the issue to say you are picking it up — it saves two people doing the same work.
  3. Fork the repository and make your change on a branch.
  4. Open a pull request that references the issue number.

Research direction

Start from the existing LigerRMSNorm kernel and module, then compare its behavior with the PyTorch reference operation described in the issue. Define the fused API and broadcasting, affine, forward, and backward cases before adding correctness tests; done includes those tests passing and a benchmark covering the three requested alternatives.

Written by the indexing model from the issue text.

Assessment

Tech stack
python, pytorch
Domain
machine-learning, performance
Issue type
Feature
Difficulty
5/5
Estimated time
Over a week
Activity status
Quiet
Clarity
Mostly clear
Newbie friendliness
45/100

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.