fix: chunk target logprobs without tensor parallelism (long-context OOM) - #4222
Draft
zhiyuanma-agents wants to merge 1 commit into
Draft
zhiyuanma-agents wants to merge 1 commit into
zhiyuanma-agents wants to merge 1 commit into
Conversation
`_tp_target_logprobs` takes a full-local-vocabulary branch when `tp_group is None`. In NeMo RL that branch is reached from the Automodel context-parallel path when this rank owns the full vocabulary, i.e. context parallel > 1 with tensor parallel 1. It ran a plain float32 `log_softmax` + gather under ordinary autograd, so its chunk loop bounded nothing: every chunk's float32 log-softmax output was retained for backward, and with `chunk_size=None` it was a single float32 [B, S, V] tensor whose float32 gradient is allocated again in backward (two ~10 GiB tensors at 20k tokens/GPU with a 131k-entry vocabulary). An allocator snapshot of an out-of-memory long-context step showed exactly that pair: one live 9.94 GiB block and a 9.94 GiB request inside `backward()`. - Add `LocalChunkedLogprob`, the `tp_group is None` counterpart of `ChunkedDistributedLogprob`: forward computes a per-chunk float32 logsumexp + gather and saves only the input logits (in their own dtype) and the targets; backward recomputes the per-chunk softmax and writes `(onehot - softmax) * grad_output` into a gradient buffer of the logits' dtype. Peak memory is one chunk of float32 logits instead of the whole tensor. Non-positive chunk sizes are rejected rather than silently returning an uninitialized buffer. The top-k/top-p filtering path is unchanged apart from also handling a zero-length sequence. - Default this branch's chunk to `DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE` (1024) when the caller passes no usable `chunk_size`. Unlike the vocabulary-parallel kernels this path needs no collectives, so chunking costs nothing but a Python loop. - `get_next_token_logprobs_from_logits` computed `use_chunking` only for the vocabulary-parallel path, so on the Automodel context-parallel path (`cp_sharder` set, `vocab_parallel_group` unset) it pre-cast the whole logits tensor to float32 before dispatch and defeated the chunking below it. Decide the pre-cast per dispatch target instead. The DTensor (tensor-parallel, context-parallel-size 1) branch was in the same position and additionally never forwarded `chunk_size` to `get_logprobs_from_vocab_parallel_logits` even though it accepts one; both are fixed. - Thread `policy.logprob_chunk_size` into the Automodel training loss path. `LossPostProcessor`'s `prepare_loss_input` partial never forwarded it, so `chunk_size` arrived as `None` during training and chunking was active only in `get_logprobs`. The Megatron `LossPostProcessor` already forwards it. Tests: CPU-only numerics for the new function (forward against `log_softmax` + gather, backward against the autograd reference, float32 and bfloat16 inputs, chunk sizes that do not divide the sequence, the default chunk size, `inference_only`, and the `no_grad` path), that it saves only its inputs, that non-positive and empty inputs are handled, and that the dispatcher casts per dispatch target -- including a gloo-gated single-rank DTensor case. Plus a test that `policy.logprob_chunk_size` reaches `prepare_loss_input`. Signed-off-by: Zhiyuan Ma <zhiyuan.ma@scale.com> Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
zhiyuanma-agents
force-pushed
the
fix/chunked-logprobs-without-tensor-parallel
branch
from
September 21, 2026 06:05
38ccfbe to
3d90567
Compare
Contributor
|
likely related to (but not a direct overlap of) #4114 |
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.
What
_tp_target_logprobstakes a full-local-vocabulary branch whentp_group is None. In NeMo RL that branch is reached from the Automodel context-parallel path (get_cp_sharded_next_token_logprobs) when this rank owns the full vocabulary — context parallel > 1 with tensor parallel 1;cp_sharderisNonebelow that, soCP=1never gets here. With TP > 1 the same function routes throughChunkedDistributedLogprob, a custom autograd Function that saves only the input logits and recomputes per chunk in backward. The TP = 1 branch instead fell through to a plain Python loop oflog_softmax(logits[:, s:e].to(float32))+ gather under ordinary autograd.That loop bounds nothing. Every chunk's float32 log-softmax output is retained for backward, and when
chunk_sizeisNoneit is one float32[B, S, V]tensor; backward then allocates a second one for its gradient. At 20k tokens/GPU with a 131k-entry vocabulary that is two ~10 GiB tensors. A CUDA allocator snapshot of an out-of-memory long-context step shows exactly that pair: one live 9.94 GiB block and a 9.94 GiB request insidebackward(), while weights plus checkpointed activations were ~54 GiB.Two adjacent problems kept the fix from taking effect on its own, and are fixed here too:
get_next_token_logprobs_from_logitscomputeduse_chunkingonly for the vocabulary-parallel path. On the Automodel context-parallel path (cp_sharderset,vocab_parallel_groupunset) it therefore pre-cast the whole logits tensor to float32 before dispatch — materializing the very tensor the chunked kernels below it exist to avoid. The DTensor branch (tensor parallel,CP=1) was in the same position, and additionally never forwardedchunk_sizetoget_logprobs_from_vocab_parallel_logitseven though that function accepts one and plumbs it to the chunked kernel.policy.logprob_chunk_size:LossPostProcessor'sprepare_loss_inputpartial innemo_rl/models/automodel/train.pyomitschunk_size, so it arrived asNoneand chunking was active only inget_logprobs, never in training. The MegatronLossPostProcessoralready forwards it. Becausechunk_size=Noneis the realistic training case, the local kernel needs a sane default chunk and the dispatcher must not pre-cast in that case — both are part of this change.Changes
LocalChunkedLogprob(nemo_rl/distributed/model_utils.py), thetp_group is Nonecounterpart ofChunkedDistributedLogprob. Forward computes a per-chunk float32 logsumexp + gather and saves only the input logits in their own dtype plus the targets; backward recomputes the per-chunk softmax and writes(onehot - softmax) * grad_outputinto a gradient buffer of the logits' dtype. Peak memory becomes one chunk of float32 logits instead of the whole tensor, in both passes. A non-positivechunk_sizeraises instead of skipping the loop and returning an uninitialized buffer. The existing top-k/top-p filtering path stays on plain autograd, with the same zero-length-sequence handling as the new branch.DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE = 1024is used on this branch when the caller passes no usablechunk_size. Unlike the vocabulary-parallel kernels this path needs no collectives, so chunking costs nothing but a Python loop, andNonecurrently means "one chunk". Documented on thechunk_sizeargument, inPolicyConfig, and in the Automodel context-parallel design note.get_next_token_logprobs_from_logits, pluschunk_sizeforwarded on the DTensor branch. The pre-cast is skipped exactly when the kernel that will run casts per chunk. It is kept on the context-parallel + vocabulary-shard +chunk_size=Nonecombination because the unchunkedDistributedLogprobcasts the whole tensor to float32 itself and saves a full float32 softmax for backward, so skipping it would bound nothing. (To be explicit, since an earlier draft of this description said otherwise: this is not a dtype-correctness requirement — autograd silently casts a mismatched floating gradient back to the input dtype.) Behavior on the vocabulary-parallel and non-parallel branches is unchanged.policy.logprob_chunk_sizethreaded into the Automodel training loss path viaLossPostProcessor, matching the Megatron post-processor.Behavior change and trade-off
Numerics are invariant — chunking only reassociates the reduction — but two things do change for existing users:
chunk_sizeturns on chunking in the training forward for recipes that already setpolicy.logprob_chunk_size; 11 automodel recipes inexamples/configs/recipes/do, at 1024–4096. Previously onlyget_logprobschunked.DistributedLogprobis replaced byChunkedDistributedLogprob. That trades the saved full float32 softmax for recompute, and 2 large all-reduces for 2 ×num_chunkssmall ones per pass (plus the same per-chunk collectives again in backward). This is the same trade the vocabulary-parallel path has always made whenlogprob_chunk_sizeis set; it is a deliberate memory-for-collectives exchange, and worth a look from someone with a latency-sensitive TP recipe.Verification
Numerics (CPU). New
tests/unit/distributed/test_local_chunked_logprob.py, 44 cases:log_softmax+ gather, and backward against autograd through that reference, for float32 and bfloat16 logits and for chunk sizes that do not divide the sequence (1, 2, 3, 5, 7, 8, 64);1024 + 7tokens, so the fallback path is exercised across an odd boundary;[B, S, V]is materialized for backward;inference_only, theno_gradpath used byget_logprobs, non-positive chunk sizes, and zero-length sequences on both branches;chunk_size), still does when top-k/top-p filtering is requested, and — in a gloo-gated single-ranktpmesh test — followschunk_sizewhen the logits are a DTensor, which stays a DTensor either way;_tp_target_logprobsstill matches its own reference.Measured deviation from the unchunked reference at those small shapes, over 40 seeds x the chunk sizes above: max 9.5e-07 forward and 2.4e-07 backward for float32 logits. bfloat16 gradients happen to round identically at that shape, but that is shape-dependent — larger vocabularies drift to ~1.5e-05. Test tolerances are 2e-06 forward and 2e-05 backward.
tests/unit/models/automodel/test_automodel_train.pygains a test thatpolicy.logprob_chunk_sizereachesprepare_loss_input.End to end. A 160k-token context-parallel-8 run with dynamic batching that ran out of memory at step 1 in six consecutive launches trains with this change: step 1 completed in 281 s of productive time at 99.6% step efficiency. Peak memory at 20k tokens/GPU drops by roughly 15 GiB.
Notes for reviewers
MixedPrecisionPolicy(output_dtype=torch.float32)inautomodel/setup.pyanddtensor/parallelize.pystores every FSDP unit's output in float32, which is worth its own knob. Ruling that out (and an unconditional.float()logits upcast in one Transformers modeling file) is how the allocations above were attributed to this path.🤖 Generated with Claude Code