perf(megatron): autotune fused LM-head Triton kernels - #2373
dyurk-lila wants to merge 3 commits into
Conversation
Signed-off-by: David Yurk <dyurk@users.noreply.github.com>
There was a problem hiding this comment.
Code Review
This pull request refactors and optimizes the Triton fused linear log-prob kernel, introducing cached and bucketed autotuning to prevent recompilation churn, compiling out entropy-only operations when unused, and ensuring numerical stability for strongly negative logits by using the true row maximum for epilogue shifts. It also updates the benchmark suite and adds comprehensive tests. The review feedback correctly identifies an optimization opportunity in the backward pass: when vocab_size is smaller than 9504, hardcoding vocab_per_split to 9504 causes unnecessary memory allocation and redundant copy overhead. Implementing the suggested min(vocab_size, 9504) limit will improve efficiency.
| vocab_per_split = 9504 | ||
| num_splits = (vocab_size + vocab_per_split - 1) // vocab_per_split | ||
| d_logits = torch.empty((num_tokens, vocab_per_split), device=hidden.device, dtype=hidden.dtype).contiguous() |
There was a problem hiding this comment.
When vocab_size is smaller than 9504 (which is very common per rank when using Tensor Parallelism), hardcoding vocab_per_split = 9504 leads to allocating a significantly larger d_logits staging buffer than necessary. Furthermore, because split_width will be less than vocab_per_split, the code will fall into the else block and perform an expensive .contiguous() copy to slice the tensor, before freeing the original.\n\nBy setting vocab_per_split = min(vocab_size, 9504), we can avoid allocating excess columns, eliminate the redundant copy overhead, and significantly reduce peak GPU memory usage.
vocab_per_split = min(vocab_size, 9504)\n num_splits = (vocab_size + vocab_per_split - 1) // vocab_per_split\n d_logits = torch.empty((num_tokens, vocab_per_split), device=hidden.device, dtype=hidden.dtype).contiguous()
|
| config.num_warps, | ||
| ) | ||
|
|
||
|
|
There was a problem hiding this comment.
Sweep measures different kernel The sweep launches candidates with
COMPUTE_ENTROPY=True, but the SkyRL adapter uses compute_entropy=False. These kernels do different work, so the reported winning schedule and latency gains may not represent the production path. Measuring the log-probability-only specialization would make the results useful for that path.
| logprobs, entropy, maximum, accumulate, entropy_b = fused_linear_logprob_triton.efficient_entropy_forward( | ||
| hidden, weight, labels, 1.0, None | ||
| ) |
There was a problem hiding this comment.
TP path lacks regression coverage This negative-logit test uses the local, entropy-enabled forward path. Tensor parallelism uses a separate epilogue, and the production adapter disables entropy, so the test would not catch a regression in that combination. A TP log-probability-only case would cover the changed production path.
Knowledge Base Used: Backend workers and distributed strategies
What does this PR do?
Autotune the fused LM-head Triton kernels across bounded shape buckets and skip entropy work for log-probability-only calls.
TLDR: select launch geometry for each workload class while keeping autotuning inputs immutable and log-sum-exp stable for strongly negative logits.
How it works
H100 benchmark
The standalone sweep used one 8x NVIDIA H100 80GB node with 180 CPU requested and limited. It covered four public LM-head configurations, total token counts from 16K through 256K, and every
(TP, CP)combination in{1,2,4,8}withTP*CP <= 8. The 200 logical cases collapse to 104 unique rank-local(M,H,V)shapes. This is a synthetic kernel-shape sweep; it does not consume dataset examples.Compared with the prior fixed five-stage mainloop:
The largest observed win was the Nemotron Ultra 550B shape at local
(M=8192,H=8192,V=131072), from 34.91 ms to 25.54 ms. These are isolated fused-mainloop results on one GPU generation, not end-to-end training claims; the candidate set intentionally retains schedule diversity for other models and accelerators.Compile and tuning behavior
A specialization probe measured:
Crossing a power-of-two bucket intentionally retunes because the winning geometry can change. Broad cold tuning is nontrivial (14.9 s median first call in this sweep), so long-running jobs should warm expected buckets before timing.
Validation
Before the short-tail metrics follow-up, focused CPU suites passed in both mirrors: 415 tests across collation, DP sampling, worker batching, layout, active spans, and Megatron correctness (13 Megatron-only skips); 40 bin-packer tests; and 3 unpacked active-span guard tests with megatron-core installed. Ruff and Black passed. Seven targeted Triton value/gradient tests passed on one H100; the full GPU suite and benchmarks were not rerun.
After the metrics follow-up, 39 focused SFT CPU tests passed in each mirror, including synchronous and asynchronous short-tail logging; Ruff, Black, and secret detection passed on the changed files.
Review order
Part 5 of 8 in a linked series. All eight PRs target
main; later PRs have cumulative diffs that include their prerequisites. Review and merge from bottom to top.Depends on #2265. Review only this layer.
Note
Medium Risk
Touches core Megatron LM-head forward/backward numerics and TP reductions; first-seen autotune buckets can add multi-second cold-start latency before training is timed.
Overview
Autotunes the vendored fused LM-head Triton stack (forward mainloop, epilogues, backward split) over multiple tile/stage/warp candidates, with cached tuning keyed by power-of-two token buckets so nearby sequence lengths reuse kernels without per-exact-length recompilation.
The log-prob-only Megatron adapter (
FusedLinearLogprobTriton) now drives kernels withcompute_entropy=False, skipping entropy buffers and math while keeping TP behavior. Epilogues write final log-probs to separate buffers so autotune replays cannot corrupt reduction inputs; log-sum-exp uses true row max and-infpadding so all-negative logits stay finite. Backward is narrowed to the split-d_logitspath (unused VERL reductions/strategies removed) and drops the stagingd_logitsbuffer before the last matmul projections.Benchmarks gain
--autotune-sweepand an entropy-vs-logprob-only probe; tests add schedule/bucket invariants and a strongly negative logits regression case.Reviewed by Cursor Bugbot for commit 0f50e56. Bugbot is set up for automated code reviews on this repo. Configure here.