Skip to content

fix: chunk target logprobs without tensor parallelism (long-context OOM) - #4222

Draft
zhiyuanma-agents wants to merge 1 commit into
NVIDIA-NeMo:mainfrom
zhiyuanma-agents:fix/chunked-logprobs-without-tensor-parallel
Draft

zhiyuanma-agents wants to merge 1 commit into
NVIDIA-NeMo:mainfrom
zhiyuanma-agents:fix/chunked-logprobs-without-tensor-parallel

Conversation

@zhiyuanma-agents

@zhiyuanma-agents zhiyuanma-agents commented Sep 21, 2026 •

Copy link
Copy Markdown

What

_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 (get_cp_sharded_next_token_logprobs) when this rank owns the full vocabulary — context parallel > 1 with tensor parallel 1; cp_sharder is None below that, so CP=1 never gets here. With TP > 1 the same function routes through ChunkedDistributedLogprob, 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 of log_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_size is None it 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 inside backward(), 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:

  1. get_next_token_logprobs_from_logits computed use_chunking only for the vocabulary-parallel path. On the Automodel context-parallel path (cp_sharder set, vocab_parallel_group unset) 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 forwarded chunk_size to get_logprobs_from_vocab_parallel_logits even though that function accepts one and plumbs it to the chunked kernel.
  2. The training-side caller never forwarded policy.logprob_chunk_size: LossPostProcessor's prepare_loss_input partial in nemo_rl/models/automodel/train.py omits chunk_size, so it arrived as None and chunking was active only in get_logprobs, never in training. The Megatron LossPostProcessor already forwards it. Because chunk_size=None is 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), 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 plus the targets; backward recomputes the per-chunk softmax and writes (onehot - softmax) * grad_output into 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-positive chunk_size raises 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 = 1024 is used on this branch 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, and None currently means "one chunk". Documented on the chunk_size argument, in PolicyConfig, and in the Automodel context-parallel design note.
  • Pre-cast decided per dispatch target in get_next_token_logprobs_from_logits, plus chunk_size forwarded 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=None combination because the unchunked DistributedLogprob casts 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_size threaded into the Automodel training loss path via LossPostProcessor, 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:

  • Threading chunk_size turns on chunking in the training forward for recipes that already set policy.logprob_chunk_size; 11 automodel recipes in examples/configs/recipes/ do, at 1024–4096. Previously only get_logprobs chunked.
  • On TP > 1 context-parallel paths that now chunk, DistributedLogprob is replaced by ChunkedDistributedLogprob. That trades the saved full float32 softmax for recompute, and 2 large all-reduces for 2 × num_chunks small ones per pass (plus the same per-chunk collectives again in backward). This is the same trade the vocabulary-parallel path has always made when logprob_chunk_size is 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:

  • forward against an unchunked 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);
  • the default chunk size on a sequence of 1024 + 7 tokens, so the fallback path is exercised across an odd boundary;
  • the gradient buffer keeps the logits' dtype, and the saved tensors are only the inputs — nothing of size [B, S, V] is materialized for backward;
  • inference_only, the no_grad path used by get_logprobs, non-positive chunk sizes, and zero-length sequences on both branches;
  • the dispatcher does not upcast on the context-parallel path with full-vocabulary logits (with and without chunk_size), still does when top-k/top-p filtering is requested, and — in a gloo-gated single-rank tp mesh test — follows chunk_size when the logits are a DTensor, which stays a DTensor either way;
  • the top-k/top-p branch of _tp_target_logprobs still 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.py gains a test that policy.logprob_chunk_size reaches prepare_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

  • Related but deliberately out of scope: FSDP2's MixedPrecisionPolicy(output_dtype=torch.float32) in automodel/setup.py and dtensor/parallelize.py stores 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

@copy-pr-bot

copy-pr-bot Bot commented Sep 21, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@github-actions github-actions Bot added the Documentation Improvements or additions to documentation label Sep 21, 2026
`_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
zhiyuanma-agents force-pushed the fix/chunked-logprobs-without-tensor-parallel branch from 38ccfbe to 3d90567 Compare September 21, 2026 06:05
@pstjohn

pstjohn commented Sep 23, 2026

Copy link
Copy Markdown
Contributor

likely related to (but not a direct overlap of) #4114

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

community-request Documentation Improvements or additions to documentation

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants