Conversation
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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 + 1rounds to1.0, while the FP32 sum is1.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
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:
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.
make testto ensure correctnessmake checkstyleto ensure code stylemake test-convergenceto ensure convergenceFull-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
95b01e9and 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" -qOriginal pytest summary lines, labelled by GPU and revision:
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:
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:
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 includesautograd.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):
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.