Skip to content

fix(cutedsl): fuse Gemma RMSNorm weight offset in FP32 - #1462

Open
luca-888 wants to merge 1 commit into
linkedin:mainfrom
luca-888:fix/cutedsl-gemma-weight-offset
Open

luca-888 wants to merge 1 commit into
linkedin:mainfrom
luca-888:fix/cutedsl-gemma-weight-offset

Conversation

@luca-888

Copy link
Copy Markdown
Contributor

Summary

The CuTe DSL Gemma paths in RMSNorm and FusedAddRMSNorm add the offset to BF16/FP16 weights before promoting them to FP32. This introduces extra rounding in the forward output and input gradient. For example, BF16 0.00390625 + 1 rounds to 1.0, while the FP32 sum is 1.00390625.

Pass the weight and offset separately, then add the offset in FP32 inside the forward and backward kernels. This also removes the separate weight-add kernel and temporary vector.

The fused offset handling follows Quack.

Testing Done

  • Hardware Type: NVIDIA H100 and NVIDIA B200.
  • Existing selected suites: 233 passed, 0 failed, 0 skipped on each GPU, before and after; 125 deselected. Covers RMSNorm, LayerNorm, FusedAddRMSNorm and CUDA Graph replay.
  • No new pytest cases or changes to existing test tolerances.

The existing BF16 Gemma RMSNorm test at M=128, N=4096 passes on both revisions. Its maximum absolute error against Triton improves identically on both GPUs:

Output Before After
Y 0.0625 0
dX 0.0625 0.0078125
dW 0 0

Cache-cleared benchmarks show lower forward latency across the seven measured BF16 RMSNorm shapes: 3.8–31.4% on H100 and 13.9–42.1% on B200. Removing the BF16 weight vector saves 8/16/36 KiB of forward/full peak allocation at N=4096/8192/18432. Full-call timings have more variation; details are below.

  • run make test to ensure correctness
  • run make checkstyle to ensure code style
  • run make test-convergence to ensure convergence

Full-repository and model-convergence suites were not run. No end-to-end training speed, memory or convergence claim is made.

Environment, test command and original pytest output

Python 3.12.10, PyTorch 2.9.1, CUDA 12.8, Triton 3.5.1, CuTe DSL 4.6.0, TVM FFI 0.1.13.post3. CUDA-visible memory: H100 79.2 GiB; B200 178.4 GiB.

Run on the base commit 95b01e9 and this branch:

LIGER_KERNEL_IMPL=cutedsl python -m pytest \
  test/cutedsl/test_rms_norm.py test/cutedsl/test_rms_norm_stream.py \
  test/ops/test_rms_norm.py test/ops/test_layer_norm.py \
  test/ops/test_fused_add_rms_norm.py \
  -k "nvidia-cutedsl or test_rms_norm_parity or test_rms_norm_mixed_weight_dtype or graph" -q

Original pytest summary lines, labelled by GPU and revision:

H100 — before
========= 233 passed, 125 deselected, 2 warnings in 116.55s (0:01:56) ==========

H100 — after
========== 233 passed, 125 deselected, 2 warnings in 90.10s (0:01:30) ==========

B200 — before
========= 233 passed, 125 deselected, 2 warnings in 153.33s (0:02:33) ==========

B200 — after
========= 233 passed, 125 deselected, 2 warnings in 102.62s (0:01:42) ==========

The error table uses test_rms_norm_parity[False-True-1.0-gemma-dtype1-2-64-4096]. Some irregular parity shapes exercise the existing inline fallback.

Supplemental diagnostics on each GPU passed 122 numerical cases against a CuTe FP32-pre-add control and 26 controls for non-Gemma/zero-offset behavior, offset cache keys and CUDA Graph replay with changed X, W and dY. Coverage includes BF16/FP16/FP32, mixed FP32 weights, strided inputs, in-place/out-of-place backward, residual-output gradients and hidden sizes 512–18432, including boundaries around 4096, 8192 and 16384. These diagnostics are separate from the committed pytest suite.

PyTorch accuracy comparison and minimal reproduction

Relative L2 error against PyTorch/autograd, using BF16 X/W, M=128, N=4096, offset=1, seed=42 and weight scale=0.01:

GPU Op Y before → after dX before → after
H100 RMSNorm 0.00292 → 2.3e-05 0.00288 → 1.22e-05
H100 FusedAddRMSNorm 0.00289 → 2.53e-05 0.00191 → 0.000253
B200 RMSNorm 0.00292 → 2.3e-05 0.00288 → 1.21e-05
B200 FusedAddRMSNorm 0.00289 → 2.53e-05 0.00191 → 0.000256

FusedAddRMSNorm includes a residual-output gradient. Its backward uses a rounded saved residual, which contributes to the remaining difference from PyTorch. dW is unchanged for fixed inputs and upstream gradients.

Run this unchanged on each revision using the existing Gemma PyTorch reference:

import torch
from liger_kernel.ops.backends._cutedsl.rms_norm import _LigerRMSNormCuTeDSLFunction
from liger_kernel.testing._op_helpers import pytorch_reference_rms_norm

torch.manual_seed(42)
x = torch.randn(128, 4096, device="cuda", dtype=torch.bfloat16, requires_grad=True)
w = (torch.randn(4096, device="cuda") * 0.01).bfloat16().requires_grad_()
dy = torch.randn_like(x)
xr = x.detach().clone().requires_grad_()
wr = w.detach().clone().requires_grad_()
y = _LigerRMSNormCuTeDSLFunction.apply(x, w, 1e-6, 1.0, "gemma", False, None)
yr = pytorch_reference_rms_norm(xr, wr, 1e-6, 1.0, "gemma")
actual = (y, *torch.autograd.grad(y, (x, w), dy))
expected = (yr, *torch.autograd.grad(yr, (xr, wr), dy.clone()))
for name, a, b in zip(("Y", "dX", "dW"), actual, expected):
    delta = a.detach().double() - b.detach().double()
    rel_l2 = delta.norm() / b.detach().double().norm().clamp_min(1e-30)
    print(name, "max_abs=", delta.abs().max().item(), "relative_l2=", rel_l2.item())
Benchmark results and measurement details

Cache-cleared triton.testing.do_bench, warmup=5 ms, rep=30 ms, median of nine shuffled-provider rounds. Compilation and input/dY generation are excluded. in_place=False; full includes autograd.grad. An identical baseline is interleaved as an A/A control. GPU clocks were not locked; full timings include host enqueue gaps.

Measured 16 configurations × forward/backward/full = 48 combinations per GPU, including seven BF16 RMSNorm shapes, three FusedAddRMSNorm shapes, other dtypes and zero-offset controls. Profiling confirms removal of one weight-add kernel from each affected forward.

BF16 results, before → after (µs):

GPU Op M × N Forward Forward + backward
H100 RMSNorm 1 × 4096 8.864 → 6.080 102.640 → 95.184
H100 RMSNorm 128 × 4096 9.728 → 6.880 99.856 → 87.680
H100 RMSNorm 2048 × 4096 20.864 → 18.944 103.072 → 98.160
H100 RMSNorm 8192 × 4096 54.048 → 51.968 141.472 → 139.488
H100 RMSNorm 128 × 8192 10.592 → 8.000 103.744 → 92.160
H100 RMSNorm 2048 × 8192 31.808 → 29.504 107.216 → 94.672
H100 RMSNorm 128 × 18432 12.448 → 10.016 102.352 → 92.320
H100 FusedAddRMSNorm 128 × 4096 10.528 → 7.872 130.656 → 120.576
H100 FusedAddRMSNorm 2048 × 4096 30.800 → 28.224 130.672 → 119.888
H100 FusedAddRMSNorm 2048 × 8192 53.440 → 50.720 128.032 → 116.304
B200 RMSNorm 1 × 4096 11.104 → 6.432 61.120 → 55.680
B200 RMSNorm 128 × 4096 11.296 → 8.192 63.648 → 58.624
B200 RMSNorm 2048 × 4096 17.280 → 14.336 71.104 → 64.448
B200 RMSNorm 8192 × 4096 33.664 → 28.992 97.440 → 95.232
B200 RMSNorm 128 × 8192 12.288 → 8.192 71.008 → 62.528
B200 RMSNorm 2048 × 8192 23.360 → 18.496 66.256 → 64.528
B200 RMSNorm 128 × 18432 13.216 → 10.048 59.968 → 61.008
B200 FusedAddRMSNorm 128 × 4096 13.152 → 8.192 71.760 → 64.432
B200 FusedAddRMSNorm 2048 × 4096 22.496 → 18.464 85.216 → 77.312
B200 FusedAddRMSNorm 2048 × 8192 33.568 → 28.704 92.096 → 79.008

H100 RMSNorm backward-only changes range from −0.8% to +4.7%; B200 backward-only medians are unchanged. Backward incremental allocation and full-call live allocation are unchanged.

The B200 128×18432 full result has substantial A/A variation, and its apparent slowdown changed sign on repeat. A separate 21-round zero-offset repeat measured 53.760 → 56.096 µs (+4.3%), while the identical-baseline A/A control differed by +5.2%; forward was unchanged and backward differed by −0.15%. These measurements do not resolve a stable kernel regression or establish zero overhead.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

1 participant