Skip to content

Commit 3d90567

Browse 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>
1 parent 612d527 commit 3d90567

7 files changed

Lines changed: 700 additions & 19 deletions

File tree

‎docs/design-docs/automodel-context-parallel.md‎

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -236,6 +236,35 @@ flowchart TB
236236

237237
For `CP=1`, the worker keeps the direct fast path without constructing a sharder.
238238

239+
#### Memory of the "TP target logprob" step
240+
241+
The `TP target logprob` box above is `_tp_target_logprobs`. It has two branches, and only
242+
the vocabulary-parallel one is driven by `policy.logprob_chunk_size`:
243+
244+
| Layout | `policy.logprob_chunk_size` | Kernel |
245+
| --- | --- | --- |
246+
| `TP > 1` (DTensor logits) | set | `ChunkedDistributedLogprob` — chunked |
247+
| `TP > 1` (DTensor logits) | unset | `DistributedLogprob` — unchunked |
248+
| `TP = 1` (plain logits) | set or unset | `LocalChunkedLogprob` — always chunked, falling back to `DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE` (1024) |
249+
250+
The chunked kernels (`ChunkedDistributedLogprob`, `LocalChunkedLogprob`) save only the
251+
input logits in their own dtype and rematerialize the softmax per chunk in backward, so the
252+
largest float32 tensor alive is one chunk. The unchunked `DistributedLogprob` instead casts
253+
the whole tensor to float32 and saves a full `[B, S, V_local]` float32 softmax for backward.
254+
Backward additionally allocates a full-size gradient buffer, in the logits' dtype for the
255+
chunked kernels and in float32 for the unchunked one. At 20k tokens per GPU with a
256+
131k-entry vocabulary, the unchunked float32 activation and its float32 gradient are two
257+
~10 GiB tensors.
258+
259+
The `TP = 1` branch always chunks because running it unchunked has no upside: unlike the
260+
vocabulary-parallel kernels it needs no collectives, so chunking costs nothing but a Python
261+
loop.
262+
263+
For the same reason, `get_next_token_logprobs_from_logits` must not pre-cast the whole
264+
logits tensor to float32 before dispatching to a chunked kernel: those kernels cast per
265+
chunk, and the pre-cast would reintroduce exactly the tensor the chunking exists to avoid.
266+
It still pre-casts on the paths whose kernel materializes float32 itself.
267+
239268
### 3.5 X-Token Distillation Workflow Before and After
240269

241270
The outer algorithm is unchanged: tokenize and align fixed text, export teacher full-vocab

‎nemo_rl/algorithms/loss/utils.py‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -459,7 +459,11 @@ def prepare_loss_input(
459459
vocab_parallel_group=vocab_parallel_group,
460460
context_parallel_group=context_parallel_group,
461461
sampling_params=None, # no filtering
462-
# Only reachable with top-k/top-p sampling active that has its own kernel path so don't chunk here
462+
# Only reachable with top-k/top-p sampling active, whose own
463+
# kernel path owns chunking for the filtered call above.
464+
# Without filtering here the vocabulary-parallel kernels run
465+
# unchunked; the full-vocabulary one still chunks at its own
466+
# default (see nemo_rl.distributed.model_utils).
463467
chunk_size=None,
464468
cp_sharder=cp_sharder,
465469
)

‎nemo_rl/distributed/model_utils.py‎

Lines changed: 174 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,14 @@
3333
ContextParallelSharder,
3434
)
3535

36+
# Sequence chunk used by the full-local-vocabulary log-prob path when the caller
37+
# does not set ``policy.logprob_chunk_size``. Unlike the vocabulary-parallel
38+
# kernels this path needs no collectives, so chunking costs nothing but a Python
39+
# loop, while a single chunk makes every per-chunk float32 buffer full-size again
40+
# (~10 GiB at 20k tokens/GPU with a 131k-entry vocabulary). Chunking only
41+
# reassociates the reduction, so the result is numerically equivalent.
42+
DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE = 1024
43+
3644

3745
def _compute_distributed_log_softmax_with_grad(
3846
vocab_parallel_logits: torch.Tensor, group: torch.distributed.ProcessGroup
@@ -888,6 +896,113 @@ def backward(
888896
return grad_input, None, None, None, None, None, None
889897

890898

899+
class LocalChunkedLogprob(torch.autograd.Function):
900+
"""Target log probabilities over a full (non-vocabulary-parallel) vocabulary.
901+
902+
The ``tp_group is None`` counterpart of :class:`ChunkedDistributedLogprob`:
903+
the sequence dimension is chunked in *both* passes, and the float32 cast
904+
happens inside the chunk loop. Forward keeps only the per-chunk logsumexp and
905+
gather, saving nothing but the input logits (in their own dtype) and the
906+
targets; backward recomputes the per-chunk softmax and writes
907+
``(onehot - softmax) * grad_output`` into a gradient buffer of the logits'
908+
dtype.
909+
910+
A plain ``log_softmax(...).gather(...)`` under ordinary autograd cannot do
911+
this: every chunk's float32 log-softmax output stays alive until backward, so
912+
looping over chunks bounds nothing and backward allocates a second float32
913+
[B, S, V] tensor for the gradient.
914+
915+
In NeMo RL this branch is reached from the Automodel context-parallel path
916+
(``get_cp_sharded_next_token_logprobs``) when this rank owns the full
917+
vocabulary, i.e. context parallel > 1 with tensor parallel 1.
918+
"""
919+
920+
@staticmethod
921+
def forward( # pyrefly: ignore[bad-override] Always ignore torch.autograd.Function.forward's type since it's always more specific than the base class
922+
ctx: Any,
923+
logits: torch.Tensor,
924+
target: torch.Tensor,
925+
chunk_size: int,
926+
inference_only: bool = False,
927+
) -> torch.Tensor:
928+
if chunk_size <= 0:
929+
# Guard the output buffer: a non-positive chunk makes the loop below
930+
# run zero times and return uninitialized memory.
931+
raise ValueError(f"chunk_size must be positive, got {chunk_size}")
932+
933+
seq_size = int(logits.shape[1])
934+
num_chunks = (seq_size + chunk_size - 1) // chunk_size
935+
936+
log_probs = torch.empty(
937+
logits.shape[:2], dtype=torch.float32, device=logits.device
938+
)
939+
for chunk_idx in range(num_chunks):
940+
chunk_start = chunk_idx * chunk_size
941+
chunk_end = min(seq_size, (chunk_idx + 1) * chunk_size)
942+
943+
logits_chunk = logits[:, chunk_start:chunk_end, :].to(dtype=torch.float32)
944+
selected = logits_chunk.gather(
945+
dim=-1, index=target[:, chunk_start:chunk_end].unsqueeze(-1)
946+
).squeeze(-1)
947+
log_probs[:, chunk_start:chunk_end] = selected - torch.logsumexp(
948+
logits_chunk, dim=-1
949+
)
950+
951+
# Explicitly free before the next iteration allocates
952+
del logits_chunk, selected
953+
954+
if not inference_only:
955+
# Only the inputs are saved: backward rematerializes the softmax.
956+
ctx.save_for_backward(logits, target)
957+
ctx.chunk_size = chunk_size
958+
959+
return log_probs
960+
961+
@staticmethod
962+
def backward(
963+
ctx: Any,
964+
*grad_outputs: torch.Tensor,
965+
) -> tuple[torch.Tensor, None, None, None]:
966+
grad_output = grad_outputs[0]
967+
logits, target = ctx.saved_tensors
968+
chunk_size = ctx.chunk_size
969+
970+
seq_size = int(logits.shape[1])
971+
num_chunks = (seq_size + chunk_size - 1) // chunk_size
972+
973+
# Every element is overwritten below, so this does not need zeroing.
974+
grad_input: torch.Tensor = torch.empty_like(logits)
975+
# The scatter_add_ source is all ones; allocate it once and slice it for
976+
# the (possibly shorter) tail chunk.
977+
ones = torch.ones(
978+
(int(logits.shape[0]), min(chunk_size, seq_size), 1),
979+
dtype=torch.float32,
980+
device=logits.device,
981+
)
982+
983+
for chunk_idx in range(num_chunks):
984+
chunk_start = chunk_idx * chunk_size
985+
chunk_end = min(seq_size, (chunk_idx + 1) * chunk_size)
986+
987+
logits_chunk = logits[:, chunk_start:chunk_end, :].to(dtype=torch.float32)
988+
989+
# d(log p_t)/d(z_v) = onehot(t)_v - softmax(z)_v, built in place so
990+
# the chunk never holds more than one [B, chunk, V] float32 tensor.
991+
chunk_grad_fp32 = torch.softmax(logits_chunk, dim=-1).neg_()
992+
del logits_chunk
993+
994+
chosen = target[:, chunk_start:chunk_end].unsqueeze(-1)
995+
chunk_grad_fp32.scatter_add_(-1, chosen, ones[:, : chunk_end - chunk_start])
996+
chunk_grad_fp32.mul_(grad_output[:, chunk_start:chunk_end].unsqueeze(-1))
997+
grad_input[:, chunk_start:chunk_end, :].copy_(chunk_grad_fp32)
998+
999+
# Explicitly free before the next iteration allocates
1000+
del chosen, chunk_grad_fp32
1001+
1002+
# if you add an argument to the forward method, then you must add a corresponding None here
1003+
return grad_input, None, None, None
1004+
1005+
8911006
def _tp_target_logprobs(
8921007
vocab_parallel_logits: torch.Tensor,
8931008
target: torch.Tensor,
@@ -913,7 +1028,11 @@ def _tp_target_logprobs(
9131028
vocab_end_index: Exclusive global ID after the local vocabulary.
9141029
tp_group: Vocabulary-parallel process group, or ``None`` when this rank
9151030
owns the full vocabulary.
916-
chunk_size: Optional sequence chunk size to bound peak memory.
1031+
chunk_size: Optional sequence chunk size to bound peak memory. The
1032+
vocabulary-parallel kernels run unchunked when this is ``None``; the
1033+
full-vocabulary path instead falls back to
1034+
``DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE`` because running it unchunked has
1035+
no upside (see :class:`LocalChunkedLogprob`).
9171036
sampling_params: Optional top-k/top-p filtering configuration.
9181037
inference_only: Skip saving tensors for backward.
9191038
@@ -959,20 +1078,40 @@ def _tp_target_logprobs(
9591078
inference_only,
9601079
).contiguous()
9611080

962-
# Full local vocabulary: plain log-softmax + gather, chunked when requested.
1081+
# Full local vocabulary: chunk in both passes so the float32 log-softmax is
1082+
# never retained for backward. Always chunked -- see the ``chunk_size``
1083+
# argument docs above.
1084+
if not need_top_k_or_top_p_filtering(sampling_params):
1085+
local_chunk_size = (
1086+
chunk_size
1087+
if chunk_size and chunk_size > 0
1088+
else DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE
1089+
)
1090+
return LocalChunkedLogprob.apply( # type: ignore[no-any-return]
1091+
vocab_parallel_logits,
1092+
target,
1093+
local_chunk_size,
1094+
inference_only,
1095+
).contiguous()
1096+
1097+
# Top-k/top-p filtering stays on plain autograd: the filtered logits are what
1098+
# log-softmax must see, and the mask is produced inside the chunk loop.
1099+
assert sampling_params is not None
9631100
seq_len = int(target.shape[1])
1101+
if seq_len == 0:
1102+
return torch.empty(
1103+
target.shape, dtype=torch.float32, device=vocab_parallel_logits.device
1104+
)
9641105
effective_chunk_size = chunk_size or seq_len
9651106
out_chunks: list[torch.Tensor] = []
9661107
for start in range(0, seq_len, effective_chunk_size):
9671108
end = min(seq_len, start + effective_chunk_size)
9681109
logits_chunk = vocab_parallel_logits[:, start:end, :].to(torch.float32)
969-
if need_top_k_or_top_p_filtering(sampling_params):
970-
assert sampling_params is not None
971-
logits_chunk, _ = apply_top_k_top_p(
972-
logits_chunk,
973-
top_k=sampling_params.top_k,
974-
top_p=sampling_params.top_p,
975-
)
1110+
logits_chunk, _ = apply_top_k_top_p(
1111+
logits_chunk,
1112+
top_k=sampling_params.top_k,
1113+
top_p=sampling_params.top_p,
1114+
)
9761115
log_probs = torch.nn.functional.log_softmax(logits_chunk, dim=-1)
9771116
out_chunks.append(
9781117
log_probs.gather(dim=-1, index=target[:, start:end].unsqueeze(-1)).squeeze(
@@ -1824,8 +1963,10 @@ def get_next_token_logprobs_from_logits(
18241963
vocab_parallel_group: Process group for vocab parallelism
18251964
context_parallel_group: Process group for context parallelism
18261965
sampling_params: Sampling parameters for top-k/top-p filtering
1827-
chunk_size: Sequence-dim chunk size for the vocab-parallel path; only
1828-
applied without top-k/top-p sampling.
1966+
chunk_size: Sequence-dim chunk size for the vocabulary-parallel and
1967+
DTensor paths; only applied without top-k/top-p sampling. The
1968+
full-vocabulary (tensor-parallel-size 1) kernel chunks at
1969+
``DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE`` when this is unset.
18291970
cp_sharder: Automodel ``ContextParallelSharder`` that sharded this
18301971
forward's model batch (V2 automodel worker with cp_size > 1). When
18311972
set, ``next_token_logits`` is this rank's CP-local shard and the
@@ -1834,12 +1975,27 @@ def get_next_token_logprobs_from_logits(
18341975
Returns:
18351976
Token log-probabilities of shape [batch_size, seq_len - 1]
18361977
"""
1837-
# ChunkedDistributedLogprob casts each chunk to float32 internally.
1838-
use_chunking = (
1839-
vocab_parallel_group is not None
1840-
and chunk_size is not None
1841-
and not need_top_k_or_top_p_filtering(sampling_params)
1842-
)
1978+
# The chunked kernels cast each chunk to float32 internally, so pre-casting
1979+
# the whole tensor here would defeat them: it materializes a float32
1980+
# [B, S, V] activation, and backward then allocates a same-size gradient.
1981+
# Only skip the pre-cast on paths that really do cast per chunk.
1982+
needs_filtering = need_top_k_or_top_p_filtering(sampling_params)
1983+
if vocab_parallel_group is not None:
1984+
use_chunking = not needs_filtering and chunk_size is not None
1985+
elif cp_sharder is not None:
1986+
# _tp_target_logprobs always chunks its full-vocabulary branch, so the
1987+
# context-parallel path casts per chunk even when chunk_size is None.
1988+
# With a vocabulary shard (TP > 1, DTensor logits) it only chunks on
1989+
# request; its unchunked kernel casts the whole tensor to float32 and
1990+
# saves a full float32 softmax anyway, so pre-casting costs nothing
1991+
# there and keeps the dtype flow explicit.
1992+
use_chunking = not needs_filtering and (
1993+
chunk_size is not None or not isinstance(next_token_logits, DTensor)
1994+
)
1995+
elif isinstance(next_token_logits, DTensor):
1996+
use_chunking = not needs_filtering and chunk_size is not None
1997+
else:
1998+
use_chunking = False
18431999
if not use_chunking:
18442000
next_token_logits = next_token_logits.to(torch.float32)
18452001

@@ -1875,6 +2031,7 @@ def get_next_token_logprobs_from_logits(
18752031
next_token_logits,
18762032
input_ids,
18772033
seq_index=seq_index,
2034+
chunk_size=chunk_size if use_chunking else None,
18782035
sampling_params=sampling_params,
18792036
)
18802037

‎nemo_rl/models/automodel/train.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -574,6 +574,9 @@ def __init__(
574574
self.dp_size = dp_size
575575
self.enable_seq_packing = enable_seq_packing
576576
self.sampling_params = sampling_params
577+
# Same knob the logprob path reads; without it the training forward runs
578+
# the log-prob computation unchunked (see LogprobsPostProcessor).
579+
self.logprob_chunk_size = cfg.get("logprob_chunk_size", None)
577580
self._cp_gradient_fanout = (
578581
cp_size
579582
if cp_size > 1
@@ -645,6 +648,7 @@ def __call__(
645648
self.cp_mesh.get_group() if self.cp_size > 1 else None
646649
),
647650
cp_sharder=token_layout,
651+
chunk_size=self.logprob_chunk_size,
648652
)
649653
# Wrap loss function for sequence packing if needed
650654
if self.enable_seq_packing:

‎nemo_rl/models/policy/__init__.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -648,7 +648,9 @@ class PolicyConfig(TypedDict):
648648
logprob_batch_size: NotRequired[int]
649649
# If set, log probability computation is chunked along the sequence dimension to avoid GPU OOM (especially during backward pass).
650650
# Within each chunk loop, logits casting (from float16/bfloat16 to float32) is done to prevent holding the entire float32 logits tensor in memory.
651-
# If None, chunking is disabled and the full sequence is processed at once.
651+
# If None, chunking is disabled and the full sequence is processed at once, except on the
652+
# full-vocabulary (tensor-parallel-size 1) log-prob path, which always chunks and falls back to
653+
# nemo_rl.distributed.model_utils.DEFAULT_LOCAL_LOGPROB_CHUNK_SIZE.
652654
logprob_chunk_size: NotRequired[int | None]
653655
generation: NotRequired[GenerationConfig]
654656
generation_batch_size: NotRequired[

0 commit comments

Comments
 (0)