Skip to content

fix(automodel): honor training logprob chunks for local vocabulary - #4114

Open
pstjohn wants to merge 2 commits into
NVIDIA-NeMo:mainfrom
pstjohn:fix/automodel-local-logprob-chunks
Open

pstjohn wants to merge 2 commits into
NVIDIA-NeMo:mainfrom
pstjohn:fix/automodel-local-logprob-chunks

Conversation

@pstjohn

@pstjohn pstjohn commented Sep 12, 2026 •

Copy link
Copy Markdown
Contributor

What does this PR do ?

Fix excessive memory use in AutoModel's local-vocabulary training loss by honoring policy.logprob_chunk_size.

With TP=1 and no context parallelism, the training loss still materializes full [batch, sequence, vocabulary] FP32 logits and log-softmax tensors even when chunking is configured. Reducing the chunk size therefore does not address this allocation during policy updates.

This PR passes the worker's TP mesh and configured chunk size through LossPostProcessor, allowing local logits to use the existing chunked autograd implementation with a singleton TP group. The kernel computes FP32 intermediates per chunk and recomputes them during backward.

Measured on a GB300, [1, 4096, 128256] bf16 logits with logprob_chunk_size=512, forward plus backward through LossPostProcessor, peak allocated over baseline:

no sampling top_k=5
before 5.87 GiB 6.36 GiB
after 2.45 GiB 6.36 GiB

The repository already has a chunked local-vocabulary loop in _tp_target_logprobs, but it is plain autograd and so keeps every chunk's FP32 log-softmax alive for backward; it reaches 3.55 GiB on the same case. ChunkedDistributedLogprob recomputes per chunk instead, which is what buys the remaining ~1.1 GiB.

This builds on #918, which introduced chunked logprob computation with deferred FP32 casting, and complements #3124, which reduced memory in AutoModel's separate logprob-evaluation path.

Issues

No linked issue. Related prior work: #918 and #3124.

Usage

For an AutoModel policy with TP=1 and CP=1, set the existing option:

policy:
  logprob_chunk_size: 256

The local-vocabulary training loss now honors this setting.

Top-k/top-p filtering is excluded from the chunked path and keeps its existing local apply_top_k_top_p plus log_softmax kernel. This is load-bearing rather than incidental: a non-None vocabulary-parallel group is what selects the vocab-parallel branch, so building one unconditionally would have moved sampling onto DistributedLogprobWithSampling — numerically identical, but an extra [B*S, V] all-to-all plus a saved softmax_output, which measured 10.28 GiB. The guard checks need_top_k_or_top_p_filtering explicitly.

Where the TP mesh is absent or not a singleton, the loss falls back to the unchunked path rather than raising. Non-DTensor logits carry the full vocabulary whatever the TP size — the DeciLM and FalconH1 lm_head plans use ColwiseParallel(output_layouts=Replicate()) — so a non-singleton mesh is a valid configuration, not an error.

Before your PR is "Ready for review"

Pre checks:

  • Make sure you read and followed Contributor guidelines
  • Did you write any new necessary tests?
  • Did you run the unit tests and functional tests locally? Visit our Testing Guide for how to run tests
  • Did you add or update any necessary documentation? Visit our Document Development Guide for how to write, build and test the docs.

Additional Information

Adds tests/unit/models/automodel/test_automodel_local_logprob_chunks.py (10 tests, automodel-marked and CUDA-gated):

  • selected logprobs, loss values, and gradients against the unchunked path in FP32 and BF16, over chunk sizes 1 (fully ragged), 7 (ragged tail), and 37 (exact fit)
  • the chunk loop is actually entered, asserted on the observed chunk widths
  • top-k/top-p never reaches from_parallel_logits_to_logprobs
  • non-positive logprob_chunk_size is rejected, since PolicyConfig does not validate it upstream

Mutation-checked: reverting the singleton group fails one test, and dropping the top-k guard clause fails another.

Local runs: tests/unit/models/automodel/ 249 passed; test_sequence_packing_gradients.py, test_sequence_packing_fusion.py, test_loss_functions.py 90 passed; pre-commit clean on both changed files.

Signed-off-by: Ubuntu <ubuntu@nvidia-lepton052.cm.cluster>
@pstjohn
pstjohn requested review from a team as code owners September 12, 2026 13:14
@copy-pr-bot

copy-pr-bot Bot commented Sep 12, 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.

@pstjohn pstjohn left a comment

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Reviewed with a team of six agents (RL-codebase, bug-finder, test, design, devil's-advocate, plus leader verification). Every finding below was independently re-verified before staging; three were dropped in the adversarial pass and two had their severity reduced.

The core change is good and I reproduced the win. On a GB300 with [1, 4096, 128256] bf16 logits and chunk_size=512 (forward+backward, peak allocated over baseline), this PR takes the training loss from 5.87 GiB to 2.45 GiB. Routing through ChunkedDistributedLogprob rather than one of the two existing hand-rolled loops is the right call — only the custom autograd function recomputes in backward, which is worth ~1.1 GiB over the plain-autograd alternative.

Two things stand between it and merge, both mechanical:

  1. The new test file reds CI — collection error in all four L0_Unit_Tests_Models_* shards, and it never runs in the Automodel shard. It also fails ruff-format. See the comment on the test file header.
  2. A silent memory regression on the sampling path — logprob_chunk_size together with top-k/top-p measures 6.36 → 10.28 GiB, the opposite of the PR's purpose. One condition fixes it.

Worth knowing about the test: it currently passes unchanged with the feature disabled, so it doesn't yet protect the behavior it was written for.

Two items I am not asking you to change here, noted only so they don't get lost:

  • chunk_size is still dropped on the DTensor (TP>1) training path — get_next_token_logprobs_from_logits doesn't forward it to get_logprobs_from_vocab_parallel_logits (model_utils.py#L1873-L1879) although LogprobsPostProcessor does on the inference side. Pre-existing, in a file this PR doesn't touch — a follow-up issue.
  • There are now three chunked-logprob implementations in the tree. They aren't interchangeable (two bound only the forward), so this isn't a straight deduplication, but it's worth a tracking issue.

Generated by Claude Code

Comment on lines +649 to +654
if (
chunk_size is not None
and token_layout is None
and self.loss_fn.input_type == LossInputType.LOGPROB
and not isinstance(logits, torch.distributed.tensor.DTensor)
):

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

train.py:649-654

2 action items.

TL;DR — with logprob_chunk_size and top-k/top-p both set, this guard sends logprobs through DistributedLogprobWithSampling and peak memory goes 6.36 → 10.28 GiB: the opposite of what the PR is for.

PR-introduced. use_chunking already excludes sampling (model_utils.py#L1838-L1842), so no chunking happens — but a non-None vocab_parallel_group is what selects the branch, and this guard sets it regardless of sampling_params. So the call leaves the local apply_top_k_top_p+log_softmax path (#L1881-L1897) and enters the vocab-parallel one (#L1846-L1860).

Measured on a GB300, [1, 4096, 128256] bf16, chunk_size=512, forward+backward, peak allocated over baseline:

no sampling top_k=5
pre-PR 5.87 GiB 6.36 GiB
this PR 2.45 GiB 10.28 GiB

The extra cost is an all_to_all copy of [B*S, V] plus softmax_output = log_probs.exp() saved for backward (#L516-L517). Values are unaffected — a direct comparison gave identical logprobs and max |dgrad| = 0.0. Latent rather than live: every in-tree automodel recipe inherits top_p: 1.0, top_k: null, so only user configs hit it.

AI-1

Don't build the group when filtering is active. need_top_k_or_top_p_filtering is already imported here. (Same edit also switches L653 to the bare DTensor imported at the top and used unqualified at L223/280/768/968/1070 — L653 is the only fully-qualified use in the file.)

Suggested change
if (
chunk_size is not None
and token_layout is None
and self.loss_fn.input_type == LossInputType.LOGPROB
and not isinstance(logits, torch.distributed.tensor.DTensor)
):
if (
chunk_size is not None
and token_layout is None
and self.loss_fn.input_type == LossInputType.LOGPROB
and not need_top_k_or_top_p_filtering(self.sampling_params)
and not isinstance(logits, DTensor)
):

One alternative worth a thought, of several: let use_chunking cover the sampling case so ChunkedDistributedLogprobWithSampling is used instead — that measured 5.45 GiB, better than pre-PR. Bigger change; the guard above is the safe fix for this PR.

AI-2

The PR description says "Top-k/top-p filtering continues to use the existing unchunked sampling path" — it's unchunked, but it's a different kernel now, so that sentence needs a correction. While editing the body, could you also paste the before/after peak-allocated number with the logits shape you measured? The PR is entirely a memory argument and currently carries no number.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989, taking AI-1 as suggested.

The guard now carries and not need_top_k_or_top_p_filtering(self.sampling_params) (train.py#L654-L661), and L657 uses the bare DTensor.

Re-measured on the same GB300 with [1, 4096, 128256] bf16 / chunk_size=512, forward+backward, peak allocated over baseline, driven through LossPostProcessor rather than the kernel directly:

no sampling top_k=5
pre-PR 5.87 GiB 6.36 GiB
this PR, before the fix 2.45 GiB 10.28 GiB
this PR, after the fix 2.45 GiB 6.36 GiB

So the sampling path is back to exactly the pre-PR number and the target path keeps its win. test_top_k_filtering_keeps_the_local_unchunked_path (test#L144-L161) pins it: it spies from_parallel_logits_to_logprobs and asserts it is never reached with top_k=5. Deleting the new guard clause makes that test fail.

On the ChunkedDistributedLogprobWithSampling alternative -- agreed it's the better end state at 5.45 GiB, but it changes which kernel the sampling path uses, and this PR has no coverage for that. Leaving it out.

AI-2: PR body updated with the corrected sentence and the table above.

Comment thread nemo_rl/models/automodel/train.py Outdated
Comment on lines +655 to +659
if chunk_size <= 0:
raise ValueError("logprob_chunk_size must be positive")
if self.tp_mesh is None or self.tp_mesh.size() != 1:
raise ValueError("chunked local logits require a singleton TP mesh")
local_vocab_group = self.tp_mesh.get_group()

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

train.py:655-659

1 action item.

TL;DR — this raises on a property that doesn't imply what the message says; for replicated full-vocab logits at TP>1, chunking is exactly what you want, and a ValueError is the wrong answer.

PR-introduced. The branch is already gated on not isinstance(logits, DTensor) — i.e. the vocabulary is not sharded. At that point tp_mesh.size() says nothing about the logits. Two Automodel lm_head plans produce exactly this shape at TP>1, because ColwiseParallel defaults to use_local_output=True, so a Replicate() output layout returns a plain full-vocab torch.Tensor:

(Both at the pinned submodule SHA 1814c6c9.) The repo already encodes this invariant — see the # Non-DTensor path (no TP sharding) comment at train.py#L777-L781.

To be clear about severity: I could not find an in-tree recipe that sets automodel-v2 dtensor_cfg.tensor_parallel_size > 1 together with logprob_chunk_size — the TP>1 values in the automodel recipes are under generation.vllm_cfg. So this is a latent trap rather than a live regression, but it turns a config that trains fine today into a hard failure the moment someone sets the chunk size.

AI-1

Fall back to the unchunked path instead of raising.

Suggested change
if chunk_size <= 0:
raise ValueError("logprob_chunk_size must be positive")
if self.tp_mesh is None or self.tp_mesh.size() != 1:
raise ValueError("chunked local logits require a singleton TP mesh")
local_vocab_group = self.tp_mesh.get_group()
if chunk_size <= 0:
raise ValueError("logprob_chunk_size must be positive")
# Non-DTensor logits are full-vocabulary regardless of TP size (some
# lm_head plans replicate their output), so a world-size-1 group is
# the correct reduction. Where the mesh cannot supply one, fall back
# to the unchunked path rather than failing a valid config.
if self.tp_mesh is not None and self.tp_mesh.size() == 1:
local_vocab_group = self.tp_mesh.get_group()

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989, taking AI-1 as written -- fallback instead of raising.

train.py#L662-L670:

            if chunk_size <= 0:
                raise ValueError("logprob_chunk_size must be positive")
            # Non-DTensor logits carry the full vocabulary whatever the TP size
            # (some lm_head plans replicate their output), so a world-size-1
            # group is the correct reduction. Where the mesh cannot supply one,
            # fall back to the unchunked path rather than failing a valid config.
            if self.tp_mesh is not None and self.tp_mesh.size() == 1:
                local_vocab_group = self.tp_mesh.get_group()

Confirmed the ColwiseParallel(output_layouts=Replicate()) reading at the pinned SHA for both plans. Agreed on severity -- latent, but it would have turned a working DeciLM/FalconH1 config into a hard failure the moment someone set the chunk size.

Side effect worth noting: this also removes the only crash the tp_mesh=None default could cause, which is what makes the revised typing-only ask on L554-L556 sufficient.

Comment on lines +666 to +667
vocab_parallel_group=local_vocab_group,
vocab_parallel_rank=0 if local_vocab_group is not None else None,

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

train.py:666-667

1 action item.

TL;DR — with sequence_packing.enabled: true these two keywords are silently discarded and the fix becomes a no-op, with no error.

PR-introduced. SequencePackingLossWrapper.__init__ defaults vocab_parallel_rank / vocab_parallel_group / context_parallel_group to None, and __call__ re-passes all three explicitly on every packed sequence. Explicit call-site keywords override functools.partial keywords, so:

  1. local_vocab_group = self.tp_mesh.get_group() (L659)
  2. the wrapper is constructed without it (L675-L680) and resets it to None per sequence
  3. use_chunking → False (model_utils.py#L1838-L1842)
  4. next_token_logits.to(torch.float32) runs on the whole tensor (#L1844) — the exact [batch, seq, vocab] FP32 materialization this PR removes

Instrumenting prepare_loss_input with enable_seq_packing=True and cfg={"logprob_chunk_size": 4} shows what actually arrives per packed sequence:

{'chunk_size': 4, 'vocab_parallel_group': None, 'vocab_parallel_rank': None, ...}

Note chunk_size is not in the wrapper's re-passed set, which is precisely why this fails silently rather than loudly. No in-tree automodel recipe currently combines packing with logprob_chunk_size, so nothing regresses today — but the fix disappears the moment someone enables packing.

AI-1

Pass the groups to the wrapper's constructor, not only into the partial. The edit is at L675-L680, outside this hunk, so as prose:

            loss_fn = SequencePackingLossWrapper(
                loss_fn=self.loss_fn,
                prepare_fn=prepare_loss_input_wrapped,
                cu_seqlens_q=processed_inputs.flash_attn_kwargs.cu_seqlens_q,
                cu_seqlens_q_padded=processed_inputs.flash_attn_kwargs.cu_seqlens_q,
                vocab_parallel_group=local_vocab_group,
                vocab_parallel_rank=0 if local_vocab_group is not None else None,
                context_parallel_group=(
                    self.cp_mesh.get_group() if self.cp_size > 1 else None
                ),
            )

The context_parallel_group line fixes the same clobbering for CP, which is pre-existing rather than something this PR introduced.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989, taking AI-1 including the context_parallel_group line.

train.py#L692-L703:

                cu_seqlens_q_padded=processed_inputs.flash_attn_kwargs.cu_seqlens_q,
                # __call__ re-passes these three explicitly per packed sequence,
                # and a call-site keyword beats a functools.partial keyword, so
                # setting them only on prepare_fn would silently drop them.
                vocab_parallel_group=local_vocab_group,
                vocab_parallel_rank=0 if local_vocab_group is not None else None,
                context_parallel_group=(
                    self.cp_mesh.get_group() if self.cp_size > 1 else None
                ),

The chunk_size-survives-but-the-group-doesn't asymmetry was the part that made this worth catching -- it fails as a silent no-op rather than an error. tests/unit/algorithms/test_sequence_packing_gradients.py and test_sequence_packing_fusion.py still pass (90 tests).

Comment thread nemo_rl/models/automodel/train.py Outdated
prepare_loss_input_wrapped = partial(
prepare_loss_input,
sampling_params=self.sampling_params,
chunk_size=chunk_size,

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

train.py:665

1 action item.

PR-introduced, low severity. chunk_size is bound here unconditionally, outside the LossInputType.LOGPROB / token_layout is None guard above, so it now reaches two paths that previously always got None on this worker:

(LOGIT, OPD_FULL, DISTILLATION and DRAFT are unaffected — I traced each; OPD_FULL accepts chunk_size but never reads it.) Probably beneficial in both cases, but it's an untested behavior change outside the PR's stated TP=1/no-CP scope.

AI-1

Either scope it to the case the guard enables:

Suggested change
chunk_size=chunk_size,
chunk_size=chunk_size if local_vocab_group is not None else None,

or keep it and say in the PR body that CP-sharded training logprobs now chunk too, with coverage for that path.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989, taking the first option -- scoped rather than widened.

train.py#L672-L673:

            chunk_size=chunk_size if local_vocab_group is not None else None,

This restores exactly the pre-PR None on the CP-sharded LOGPROB and DISTILLATION_CROSS_TOKENIZER paths. Chunking those is probably a win, but it's a separate change that deserves its own coverage rather than riding along here.

Thanks for tracing each input type -- OPD_FULL accepting chunk_size without reading it is the kind of thing that would have made this look already-covered.

Comment thread nemo_rl/models/automodel/train.py Outdated
Comment on lines +554 to +556
enable_seq_packing: bool = False,
sampling_params: Optional[TrainingSamplingParams] = None,
tp_mesh: Any = None,

@pstjohn pstjohn Sep 14, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

train.py:554-556

1 action item (revised — my original ask here was wrong, see below).

AI-1

Tighten the annotation only; DeviceMesh is already imported at L38.

Suggested change
enable_seq_packing: bool = False,
sampling_params: Optional[TrainingSamplingParams] = None,
tp_mesh: Any = None,
tp_mesh: Optional[DeviceMesh] = None,

Context — no action. I originally asked for tp_mesh to be made a required parameter, on the reasoning that the None default is dead and only serves to turn a missing wire-up into ValueError("chunked local logits require a singleton TP mesh") deep in a training step. That was wrong on both counts:

  • It isn't dead. 15 existing call sites construct this class without tp_mesh — 12 in test_automodel_train.py and 3 in test_automodel_context_parallel.py — none of which this PR touches. Making it required would mean adding tp_mesh=None to all 15, which respells the same None rather than removing it, and grows a 78-line PR by ~15 unrelated files. A post-processor not doing vocab-parallel work legitimately has no TP mesh.
  • The crash it guards against disappears anyway if you take the fallback suggested on L655-L659: once a missing or non-singleton mesh degrades to the unchunked path instead of raising, tp_mesh=None is harmless.

The TopkLogitsPostProcessor comparison I drew is also weaker than I implied — it takes tp_mesh required because vocab-parallel top-k is its entire purpose, whereas here it feeds one optional branch.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Taking the revised AI-1. train.py#L556 is now tp_mesh: Optional[DeviceMesh] = None.

Appreciate you re-checking the 15 call sites and walking this back -- I'd started on the required-parameter version and it was turning into tp_mesh=None sprayed across two test files this PR has no other business touching.

Also extended the docstring to say what the None case does, since with the L662-L670 fallback it is now a real supported state rather than an error:

            tp_mesh: Tensor-parallel mesh; its singleton group permits the
                existing chunked vocabulary loss to handle unsharded logits.
                Without one, the local-vocabulary loss stays unchunked.

Comment on lines +1 to +13
"""Qualify local-vocabulary loss chunking through the AutoModel loss boundary."""

from collections.abc import Iterator
from typing import Any

import pytest
import torch
import torch.distributed as dist
from torch.distributed.device_mesh import DeviceMesh

from nemo_rl.algorithms.loss.interfaces import LossInputType
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
from nemo_rl.models.automodel.train import LossPostProcessor

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

test_automodel_local_logprob_chunks.py:1-13

3 action items.

TL;DR — as committed this file errors at collection in all four L0_Unit_Tests_Models_* shards, never executes in the Automodel shard, and fails the ruff-format pre-commit hook.

All PR-introduced. Reproduced each:

(a) Collection error in the base shards. train.py:33 imports nemo_automodel at module scope, and nemo-automodel is an optional extra (pyproject.toml#L131). L0_Unit_Tests_Models_{1..4}.sh collect unit/models/ under uv run --no-sync with no extras (Models_1#L23), and this is the only test file sitting directly in unit/models/ — everything else is in a subdirectory.

ERROR collecting tests/unit/models/test_automodel_local_logprob_chunks.py
E   ModuleNotFoundError: No module named 'nemo_automodel'

(b) Deselected in the Automodel shard. Marker filtering happens in pytest_collection_modifyitems, and --automodel-only keeps only marked items. With no @pytest.mark.automodel:

$ pytest ... --hf-gated --automodel-only
collected 6 items
Running 0 items in this shard

All 7 existing files that import nemo_rl.models.automodel.* carry both the guard and the marker — e.g. test_automodel_train.py#L22-L25.

AI-1

Move the file to tests/unit/models/automodel/, next to test_automodel_train.py which already holds the LossPostProcessor tests.

AI-2

Add the import guard, and mark each test @pytest.mark.automodel. Adding the NVIDIA header at the same time matches every neighbour in that directory (the copyright skill does exclude tests/, so this one is consistency, not a rule).

Suggested change
"""Qualify local-vocabulary loss chunking through the AutoModel loss boundary."""
from collections.abc import Iterator
from typing import Any
import pytest
import torch
import torch.distributed as dist
from torch.distributed.device_mesh import DeviceMesh
from nemo_rl.algorithms.loss.interfaces import LossInputType
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
from nemo_rl.models.automodel.train import LossPostProcessor
# Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Qualify local-vocabulary loss chunking through the AutoModel loss boundary."""
from collections.abc import Iterator
from typing import Any
import pytest
import torch
import torch.distributed as dist
from torch.distributed.device_mesh import DeviceMesh
try:
import nemo_automodel # noqa: F401
except ImportError:
pytest.skip("nemo_automodel not available", allow_module_level=True)
from nemo_rl.algorithms.loss.interfaces import LossInputType
from nemo_rl.distributed.batched_data_dict import BatchedDataDict
from nemo_rl.models.automodel.train import LossPostProcessor

AI-3

Run uv run --group dev pre-commit run --files <path>. ruff-format (pinned v0.9.9, .pre-commit-config.yaml#L18) rewrites 46 of 53 lines here — single→double quotes, .01→0.01, and five lines over 200 chars — so the pre-commit CI job is red as-is.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

All three fixed in 617d989.

AI-1. Moved to tests/unit/models/automodel/test_automodel_local_logprob_chunks.py, beside test_automodel_train.py.

AI-2. Import guard and header applied as suggested; every test is now @pytest.mark.automodel and @pytest.mark.skipif(not torch.cuda.is_available(), ...).

Verified the guard behaves identically to its neighbour by blocking nemo_automodel through a sys.meta_path hook that raises ImportError on that name:

=== test_automodel_local_logprob_chunks ===
SKIPPED [1] .../test_automodel_local_logprob_chunks.py:28: nemo_automodel not available
=== test_automodel_train ===
SKIPPED [1] .../test_automodel_train.py:26: nemo_automodel not available

Module-level skip, no collection error. And in the Automodel shard it now selects rather than reporting zero:

$ pytest tests/unit/models/automodel/test_automodel_local_logprob_chunks.py --automodel-only
Running 10 items in this shard
10 passed, 20 warnings in 26.44s

AI-3. pre-commit run --files is clean on both changed files -- ruff, ruff-format, pyrefly, and the rest all pass. (The taplo-format hook can't build in my environment; no TOML changed.) The full tests/unit/models/automodel/ suite is 249 passed.


@pytest.mark.parametrize('dtype', [torch.float32, torch.bfloat16])
@pytest.mark.parametrize('chunk_size', [1, 7, 64])
def test_chunked_local_loss_preserves_selected_values_and_gradients(singleton_tp_mesh: DeviceMesh, dtype: torch.dtype, chunk_size: int) -> None:

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

test_automodel_local_logprob_chunks.py:39

1 action item.

TL;DR — this test passes unchanged when the PR's feature is turned off, so it cannot catch a regression in the thing it was written for.

PR-introduced. It compares chunk_size=None against chunk_size=N through the same code path and never asserts that chunking happened. Mutating train.py:647 to force chunk_size = None — i.e. reverting this PR's behavior — gives:

6 passed, 20 warnings in 27.11s

with the chunked kernel firing zero times (instrumented _compute_distributed_selected_logprobs; with the fix live it fires 37/6/1 times for chunk_size 1/7/64). The nearest analog does assert this — test_logprob_memory.py#L155-L177 spies log_softmax and checks the observed shapes.

Separately, the two new raise sites (L656, L658) have no coverage at all — grep over tests/ finds neither message. logprob_chunk_size is NotRequired[int | None] in PolicyConfig with no positivity validation upstream, so 0 is reachable straight from YAML.

AI-1

Keep the parity test, and add a test that observes the chunking plus one per raise. These are additions rather than a replacement, so as a block — verified passing, and the first one fails under the mutation above:

@pytest.mark.automodel
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
def test_loss_chunking_splits_the_sequence_into_chunks(
    singleton_tp_mesh: DeviceMesh, monkeypatch: pytest.MonkeyPatch
) -> None:
    """The training loss must actually chunk, not silently fall back to one pass."""
    import nemo_rl.distributed.model_utils as model_utils

    seen: list[int] = []
    original = model_utils._compute_distributed_selected_logprobs

    def spy(logits: torch.Tensor, *args: Any, **kwargs: Any) -> torch.Tensor:
        seen.append(int(logits.shape[1]))
        return original(logits, *args, **kwargs)

    monkeypatch.setattr(model_utils, "_compute_distributed_selected_logprobs", spy)
    _run(singleton_tp_mesh, chunk_size=7)
    assert seen == [7, 7, 7, 7, 7, 2]


@pytest.mark.automodel
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available")
@pytest.mark.parametrize("chunk_size", [0, -1])
def test_nonpositive_chunk_size_is_rejected(
    singleton_tp_mesh: DeviceMesh, chunk_size: int
) -> None:
    with pytest.raises(ValueError, match="logprob_chunk_size must be positive"):
        _run(singleton_tp_mesh, chunk_size=chunk_size)

(_run = the body of the existing test factored into a helper taking tp_mesh / chunk_size / dtype.) If you take the fallback in my comment on train.py:655-659, the second test becomes the only raise left to cover.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989. This was the right call to make -- the test was measuring nothing.

Both new tests added, plus a third for the top-k routing. The parity test is factored onto a _run(tp_mesh, chunk_size, dtype, sampling_params) helper (test#L65-L101).

One correction to the suggested block: the observed split is [7, 7, 7, 7, 7, 2], not [..., 1] -- the chunk loop runs over all 37 positions, not SEQ - 1. Your original instrumentation had it right; I'd mis-transcribed it and the test caught me.

Mutation-checked both directions, which is the part the old test failed:

mutation result
never build the singleton group (reverts the PR) 1 failed, 6 passed
drop the need_top_k_or_top_p_filtering clause 1 failed (test_top_k_filtering_keeps_the_local_unchunked_path), 7 passed
unmutated 10 passed

Previously the first of those was 6 passed. On the raise sites: the singleton-mesh one is gone entirely now that it degrades to a fallback, so test_nonpositive_chunk_size_is_rejected[0, -1] covers the one that remains.



@pytest.mark.parametrize('dtype', [torch.float32, torch.bfloat16])
@pytest.mark.parametrize('chunk_size', [1, 7, 64])

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

test_automodel_local_logprob_chunks.py:38

1 action item.

PR-introduced, minor. chunk_size=64 against S=37 is a single chunk, not an oversized one — instrumenting the chunk loop gives seen == [37] for 64, versus [7, 7, 7, 7, 7, 2] for 7 and 37×[1] for 1. So that parametrization compares the unchunked kernel against a one-pass run of the chunked kernel, which the None leg already covers. (Ragged-tail coverage via chunk_size=7 is real and worth keeping.)

AI-1

Use the exact-fit boundary instead, which is the case actually worth pinning:

Suggested change
@pytest.mark.parametrize('chunk_size', [1, 7, 64])
@pytest.mark.parametrize('chunk_size', [1, 7, 37])

and drop "oversized chunks" from the docstring on the next line.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989 -- now @pytest.mark.parametrize("chunk_size", [1, 7, 37]) (test#L107), and the docstring reads "ragged final chunks, the exact-fit boundary, signed weights, and both dtypes".

37 is the exact fit against the 37 positions the chunk loop actually walks (same off-by-one as in the thread above), so it's a single full-width chunk with no ragged tail -- the boundary worth pinning, and distinct from the None leg.

results.append((loss.detach(), metrics['selected'], logits.grad))
torch.testing.assert_close(results[1][0], results[0][0], rtol=1e-5, atol=1e-6)
torch.testing.assert_close(results[1][1], results[0][1], rtol=1e-5, atol=1e-6)
torch.testing.assert_close(results[1][2], results[0][2], rtol=.01 if dtype == torch.bfloat16 else 1e-5, atol=5e-5 if dtype == torch.bfloat16 else 1e-7)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

test_automodel_local_logprob_chunks.py:53

1 action item.

PR-introduced, minor. The gradient assertion is close to vacuous in bf16. On the exact tensors this test builds:

torch.bfloat16  max|g|=3.589e-02  median|g|=2.241e-05  frac(|g| < 5e-5): 0.797
torch.float32   max|g|=2.908e-02  median|g|=1.458e-05  frac(|g| < 5e-5): 0.797

~80% of gradient entries are below atol=5e-5, so assert_close accepts arbitrary error on most of the tensor. That's inherent to the setup rather than a typo — the gradient w.r.t. non-target logits is weight * softmax, and weights is divided by 72 across a 257-wide vocabulary.

AI-1

Scale the signal up rather than loosening further — dropping the /72 on weights at L43 raises the gradient magnitudes into a range where the existing tolerances actually bite. Tightening atol to ~1e-6 works too, but is more likely to need per-dtype tuning.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989 -- dropped the /72, so weights is now plain torch.randn(BATCH, SEQ - 1) (test#L77).

Took the scale-up rather than the tolerance tightening, for the reason you gave: atol at ~1e-6 would need per-dtype tuning and would drift the moment the fixture shape changes. The existing tolerances now bite on a gradient distribution that isn't ~80% below them. Both dtypes still pass at all three chunk sizes.

Comment on lines +25 to +34
@pytest.fixture(scope='module')
def singleton_tp_mesh(tmp_path_factory: Any) -> Iterator[DeviceMesh]:
"""Use a real singleton NCCL group, matching each FSDP rank's TP=1 mesh."""
assert torch.cuda.is_available(), 'this qualification requires a GPU'
store = tmp_path_factory.mktemp('local-logprob') / 'store'
dist.init_process_group('nccl', init_method=f'file://{store}', rank=0, world_size=1)
try:
yield DeviceMesh('cuda', [0], mesh_dim_names=('tp',))
finally:
dist.destroy_process_group()

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

test_automodel_local_logprob_chunks.py:25-34

1 action item.

PR-introduced, minor — an idiom mismatch rather than a live failure. This fixture initializes and destroys the global default process group unconditionally. The two existing in-process PG fixtures in the directory this file should move to — test_automodel_checkpoint.py#L148-L156 and test_automodel_setup.py#L2814-L2822 — both guard with if not torch.distributed.is_initialized(): and deliberately never destroy.

Co-running with either of them today gives ValueError: trying to initialize the default process group twice!. That can't bite in CI right now only because those files are automodel-marked and this one isn't — which stops being true once the marker is added.

AI-1

Match the existing idiom: guard the init, drop the destroy_process_group(), and use skipif rather than a bare assert for the CUDA requirement (the repo's 8 GPU-gated unit tests all use @pytest.mark.skipif(not torch.cuda.is_available(), ...)).

Suggested change
@pytest.fixture(scope='module')
def singleton_tp_mesh(tmp_path_factory: Any) -> Iterator[DeviceMesh]:
"""Use a real singleton NCCL group, matching each FSDP rank's TP=1 mesh."""
assert torch.cuda.is_available(), 'this qualification requires a GPU'
store = tmp_path_factory.mktemp('local-logprob') / 'store'
dist.init_process_group('nccl', init_method=f'file://{store}', rank=0, world_size=1)
try:
yield DeviceMesh('cuda', [0], mesh_dim_names=('tp',))
finally:
dist.destroy_process_group()
@pytest.fixture(scope='module')
def singleton_tp_mesh() -> DeviceMesh:
"""A real singleton TP mesh, matching each FSDP rank's TP=1 mesh."""
if not dist.is_initialized():
os.environ.setdefault('MASTER_ADDR', 'localhost')
os.environ.setdefault('MASTER_PORT', '29519')
dist.init_process_group(backend='nccl', rank=0, world_size=1)
return DeviceMesh('cuda', [0], mesh_dim_names=('tp',))

(Needs import os; Iterator then becomes unused.)

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Fixed in 617d989 (test#L53-L62) -- guarded init, no destroy_process_group(), and @pytest.mark.skipif(not torch.cuda.is_available(), ...) on each test in place of the bare assert. Also set RANK/WORLD_SIZE to match _init_gloo_pg in test_automodel_setup.py.

Agreed it's an idiom point rather than a live failure, but it stops being hypothetical the moment the automodel marker lands and this file can co-run with the other two PG fixtures -- which is this same commit. Confirmed the whole directory runs clean together: 249 passed.

Follow-up to the review on NVIDIA-NeMo#4114.

- Do not build the singleton vocabulary group when top-k/top-p filtering is
  active. `use_chunking` already excludes sampling, but a non-None
  `vocab_parallel_group` is what selects the branch, so the sampling path was
  being routed through `DistributedLogprobWithSampling`. Measured on
  [1, 4096, 128256] bf16 with chunk_size=512: top_k=5 peak allocated goes
  10.28 -> 6.36 GiB, matching the pre-change path. The target path is
  unchanged at 5.87 -> 2.45 GiB.
- Fall back to the unchunked path instead of raising when the TP mesh is
  absent or non-singleton. Non-DTensor logits carry the full vocabulary
  whatever the TP size (DeciLM and FalconH1 lm_head plans replicate their
  output), so the old ValueError failed configs that train fine today.
- Pass the vocabulary and context-parallel groups to
  `SequencePackingLossWrapper.__init__`. Its `__call__` re-passes all three
  explicitly per packed sequence, and a call-site keyword beats a
  functools.partial keyword, so the fix was silently dropped under sequence
  packing. Also fixes the same pre-existing clobbering for CP.
- Scope `chunk_size` to the case the guard enables, keeping the CP-sharded and
  cross-tokenizer distillation paths on their previous behavior.
- Annotate `tp_mesh` as `Optional[DeviceMesh]` rather than `Any`.

Move the test beside its siblings in tests/unit/models/automodel/ and give it
the `nemo_automodel` import guard plus `@pytest.mark.automodel` that all seven
neighbours carry; without them it errored at collection in the four
L0_Unit_Tests_Models shards and was deselected in the Automodel shard. The
suite now observes the chunk boundaries, pins the top-k routing, and covers
the non-positive chunk size — reverting either behavior fails it, which the
previous version did not.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Peter St. John <pstjohn@nvidia.com>
@pstjohn

pstjohn commented Sep 14, 2026

Copy link
Copy Markdown
Contributor Author

Pushed 617d9898 addressing all 11 comments. Replies are inline on each thread; summary:

Two things that blocked merge.

The top-k/top-p regression is fixed — the guard now checks need_top_k_or_top_p_filtering explicitly, so sampling stays on its local kernel. Re-measured end to end through LossPostProcessor rather than the kernel:

no sampling top_k=5
pre-PR 5.87 GiB 6.36 GiB
this PR, before the fix 2.45 GiB 10.28 GiB
now 2.45 GiB 6.36 GiB

The test file moved to tests/unit/models/automodel/ with the nemo_automodel guard and @pytest.mark.automodel, so it no longer errors at collection in the four L0_Unit_Tests_Models shards and actually runs in the Automodel shard (10 items, previously 0). pre-commit is clean.

The test now has teeth. It previously passed with the feature disabled. It now asserts the observed chunk widths and the top-k routing, and mutation-checks confirm it: reverting the singleton group fails one test, dropping the top-k guard clause fails another.

Also fixed: the SequencePackingLossWrapper clobbering (groups now go to the constructor, which also fixes the pre-existing CP case), chunk_size scoped back to the guarded path so CP-sharded and cross-tokenizer distillation keep their previous behavior, the ValueError on a non-singleton TP mesh replaced with a fallback, and tp_mesh typed as Optional[DeviceMesh].

Local: tests/unit/models/automodel/ 249 passed, sequence-packing and loss-function suites 90 passed.

One correction to a suggested snippet — the chunk loop walks all 37 positions, not SEQ - 1, so the expected split is [7, 7, 7, 7, 7, 2]. Noted on that thread.

PR description updated with the measurements and the corrected sampling-path sentence.

🤖 Generated with Claude Code

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

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant