Skip to content

perf(megatron): autotune fused LM-head Triton kernels - #2373

Open
dyurk-lila wants to merge 3 commits into
dyurk/align-packing-bins-to-dpfrom
dyurk/triton-fused-lm-autotune
Open

dyurk-lila wants to merge 3 commits into
dyurk/align-packing-bins-to-dpfrom
dyurk/triton-fused-lm-autotune

Conversation

@dyurk-lila

@dyurk-lila dyurk-lila commented Oct 1, 2026 •

Copy link
Copy Markdown
Collaborator

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

  • Tune tile dimensions, pipeline stages, and warp counts instead of using one fixed launch configuration.
  • Bucket token counts so neighboring lengths reuse compiled kernels and tuning results.
  • Keep raw, reduced, and result buffers separate: Triton executes multiple candidates during tuning, so an epilogue must not overwrite inputs needed by later candidates.
  • Start log-sum-exp from the true row maximum, including rows where every logit is negative.
  • Remove unreachable vendored reductions and backward strategies, and release the final partial d-logits buffer before its projections.

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} with TP*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:

Metric across 104 local shapes Autotuned
Mean latency reduction 5.57%
Median latency reduction 3.63%
p90 latency reduction 14.47%
Maximum latency reduction 26.82%
Shapes faster 88 / 104
Worst regression 1.40%

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:

  • 20,000 tokens, new 32K bucket: 13.27 s cold compile+tune;
  • 20,001 tokens, same bucket: 6.13 ms, with no exact-length recompile or retune;
  • 40,000 tokens, new 64K bucket: 4.66 s retune.

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.

  1. perf(sft): use modified first-fit decreasing packing #2109 — use modified first-fit decreasing packing
  2. perf(train): make sequence packing alignment-aware #2108 — make sequence packing alignment-aware
  3. refactor(megatron): share packed segment layout derivation #2264 — share packed segment layout derivation
  4. perf(sft): align packed bin counts to data parallelism #2265 — align packed bin counts to data parallelism
  5. perf(megatron): autotune fused LM-head Triton kernels #2141 — autotune fused LM-head Triton kernels
  6. feat(megatron): fuse LM-head entropy for RL #2266 — fuse LM-head entropy for RL
  7. feat(megatron): add block-sparse fused LM-head forward #2267 — add block-sparse fused LM-head forward
  8. perf(megatron): add block-sparse fused LM-head backward #2268 — add block-sparse fused LM-head backward

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 with compute_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 -inf padding so all-negative logits stay finite. Backward is narrowed to the split-d_logits path (unused VERL reductions/strategies removed) and drops the staging d_logits buffer before the last matmul projections.

Benchmarks gain --autotune-sweep and 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.

@dyurk-lila
dyurk-lila added this pull request to stack #2377 October 1, 2026 14:10

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +1051 to +1053
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()

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

high

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

@greptile-apps

greptile-apps Bot commented Oct 1, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 4/5

[High risk] Rewrites Triton kernel autotuning and scheduling logic.

The PR appears safe to merge, with non-blocking benchmark-fidelity and regression-coverage improvements recommended.

Findings

  1. P2 Sweep measures different kernel ▶
  2. P2 TP path lacks regression coverage ▶
Diagram
%%{init: {'theme': 'neutral'}}%%
flowchart LR
  A[SkyRL fused LM-head adapter] -->|compute_entropy=False| B[Production mainloop]
  C[Autotune sweep] -->|COMPUTE_ENTROPY=True| D[Measured mainloop]
  B --> E[Log-probabilities]
  D --> F[Reported winning schedule]
Loading

Reviews (1) · Last reviewed commit: "fix(megatron): pass entropy mode in auto..."

config.num_warps,
)


Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

Comment on lines +361 to +363
logprobs, entropy, maximum, accumulate, entropy_b = fused_linear_logprob_triton.efficient_entropy_forward(
hidden, weight, labels, 1.0, None
)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 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

This branch was successfully deployed

1 active deployment
Preview — 0f50e566 Deployed Oct 1, 2026 by vercel[bot]
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants