You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
Commit 3d90567
Browse filesBrowse the repository at this point in the historyBrowse files
fix: chunk target logprobs without tensor parallelism
`_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>
0 commit comments