Skip to content

fix(loss): mask filtered tokens in force-on-policy reductions - #3778

Closed
tianyi-zhang-02 wants to merge 14 commits into
NVIDIA-NeMo:mainfrom
tianyi-zhang-02:fix/neg-inf-narrowing-on-policy
Closed

tianyi-zhang-02 wants to merge 14 commits into
NVIDIA-NeMo:mainfrom
tianyi-zhang-02:fix/neg-inf-narrowing-on-policy

Conversation

@tianyi-zhang-02

@tianyi-zhang-02 tianyi-zhang-02 commented Aug 23, 2026 •

Copy link
Copy Markdown
Contributor

What does this PR do?

This follows up on #3551. That PR now owns the filtered actor mask and the actor-side retained-token aggregation; this PR extends the same protocol to the force_on_policy_ratio path, where prev_logprobs aliases curr_logprobs.detach().

When a current logprob is filtered, its finite substitute must not leak into quantities that read prev_logprobs. This change:

  • narrows the four train/generation mismatch diagnostics to the retained tokens;
  • narrows token- and sequence-level importance-sampling inputs, including TIS/ICEPOP and seq-mask-tis;
  • emits raw numerator fragments with their matching retained-token counts, then divides once after microbatch and data-parallel aggregation;
  • leaves the normal, separately computed prev_logprobs path unchanged.

The actor loss still uses the original global-token denominator, so filtering does not increase the effective step size. Reference-policy KL also keeps the full token mask.

Quantity Mask / denominator
actor diagnostics retained actor tokens / retained actor-token count (#3551)
prev-derived diagnostics with force_on_policy_ratio retained prev tokens / retained prev-token count
actor gradient retained actor tokens / original global-token count
reference-policy KL full token mask / original global-token count

Stack

The Files changed view includes #3551 until that PR lands. This branch merges both the current #3551 head and the current main and is conflict-free at the SHA above.

Validation

Final-head CPU validation on macOS arm64, Python 3.12.2, PyTorch 2.8.0, Ray 2.51.1:

Suite Result
test_loss_functions.py, test_utils.py, test_sequence_packing_fusion.py, single_controller/test_utils.py 117 passed, 45 CUDA-only skipped
test_grpo.py, test_ppo.py 259 passed, 7 CUDA-only skipped
Ruff check + format check on all 11 touched Python files passed

The tests cover uneven microbatches/ranks, zero retained tokens, all affected diagnostics, token- and sequence-level IS, TIS/ICEPOP normalization, and classic/async/sync/Single Controller caller aggregation.

The final head was not rerun on CUDA because the private 4090 host became unavailable. No functional training run is claimed here.

Before your PR is "Ready for review"

  • Read and followed the contributor guidelines
  • Added the necessary unit and caller-level tests
  • Aligned the branch with current main
  • Ran lint and formatting checks
  • Final-head CUDA / functional run

mask_out_neg_inf_logprobs prints "Masking out these positions", computes a
narrowed mask, and then returns only the logprobs -- the narrowed mask never
leaves the function. All five call sites keep the un-narrowed mask, and
ClippedPGLossFn re-derives its own reduction mask from the untouched batch.

The substituted 0.0 is not a neutral filler. It means log p = 0, i.e. p = 1,
while generation_logprobs at the same position is finite by construction
(vLLM sampled the token from its own filtered distribution). Everything that
compares the two therefore reads a large fabricated difference:
token_mult_prob_error, gen_kl_error and js_divergence_error all evaluate
exp(|0 - log pi_gen|), which is ~224 for a typical -5.4. Under
sequence_level_importance_ratios the same difference is summed over the
sequence and exponentiated, so it inflates the importance weight multiplying
the whole sequence's clipped loss -- a gradient effect, not just a metric.

Return the keep mask and fold it into token_mask at the two curr-logprob call
sites, which is where the loss reads it back. The three prev-logprob sites
cannot pass it on today because LogprobOutputSpec carries logprobs only; they
discard it explicitly with a comment rather than silently.

Reachable in shipped recipes: grpo-llama3.2-1b-instruct-1n8g-fsdp2tp2-
temp0.8-topp0.9-topk50.yaml and its megatron twin set top_k=50, top_p=0.9,
and the defect fires whenever the existing warning prints.

Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
…nostics

Folding the -inf keep mask into ``token_mask`` narrowed every reduction in
ClippedPGLossFn, not just the actor term. Two of those should not narrow:

- The reference-policy KL reduces with ``token_mask`` but is computed from
  ``curr_logprobs_unfiltered``, which is finite at the filtered positions.
  Dropping them removes real KL terms, and removes them selectively: a token
  is filtered precisely where the training policy assigns it near-zero mass
  while the reference checkpoint may not, i.e. where the two disagree most.
- token_mult_prob_error / gen_kl_error / policy_kl_error / JS divergence read
  only prev and generation logprobs, so they are unaffected by the
  substitution and narrowing them hides mismatch the metrics exist to report.

Publish the mask as ``curr_logprobs_keep_mask`` instead and apply it only to
quantities derived from ``curr_logprobs``: the actor loss (both reduction
modes), the GSPO sequence-level ratio, the entropy approximation, the VAPO
NLL term, and the probs_ratio metrics -- where a substituted logprob of 0.0
otherwise fabricates a ratio of exp(-prev_logprob).

The new test pins both halves and fails under either mutation: ignoring the
keep mask, and folding it back into ``token_mask``.

Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
The note still described the old mechanism, where the curr side folded its
narrowing into token_mask. It now publishes curr_logprobs_keep_mask and
narrows only the actor reduction.

Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
@copy-pr-bot

copy-pr-bot Bot commented Aug 23, 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.

@svcnvidia-nemo-ci svcnvidia-nemo-ci added the waiting-on-maintainers Waiting on maintainers to respond label Aug 26, 2026
The keep mask is applied at six reductions; one test pinned one of them.
Four of the rest are now covered, and two of those move the gradient rather
than a dashboard number:

  - the GSPO sequence ratio. masked_mean with no global normalization factor,
    so the filtered position corrupts numerator and divisor, and the result is
    exponentiated.
  - the SEQUENCE_LEVEL actor loss. Its inner masked_mean normalizes by the
    mask sum, so the position both injects a fabricated clip_loss and inflates
    its own divisor.
  - VAPO's positive-example NLL, which normalizes by its own mask sum: the
    filtered position adds -0.0 to the numerator and +1 to the denominator, so
    the term is deflated rather than merely noisy.
  - probs_ratio and probs_ratio_clamped. The existing test pins only
    probs_ratio_max, which comes from a separate reduction.

The two sequence-level tests widen ratio_clip deliberately: at the default
0.2 both ratios land above the ceiling and clamp to the same value, so the
loss is identical either way and the assertion would be vacuous -- while the
metrics still differ, which makes the degenerate case easy to miss.

Mutation-tested: reverting each of the four sites to the un-narrowed mask
turns one of these red.

Signed-off-by: Tianyi Zhang <zhangtianyi975@gmail.com>
Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
@tianyi-zhang-02
tianyi-zhang-02 force-pushed the fix/neg-inf-narrowing-on-policy branch from 2b7165b to 1b090c3 Compare August 26, 2026 21:37
…owed

The comment claimed they "read only prev/generation logprobs, so narrowing
them would drop real terms". That is wrong, and it contradicts the comment
this same PR adds at the worker call sites, which says the opposite.

prev_logprobs carries its own 0.0 substitutions: mask_out_neg_inf_logprobs
runs on the logprob-inference pass too, at the three worker sites this PR
edits. So those diagnostics are not reading real terms at the filtered
positions either.

The actual reason they stay on the full mask is different and weaker: the keep
mask published this step identifies what THIS forward filtered, not what the
earlier one did, so narrowing with it would drop the wrong positions.
force_on_policy_ratio is the single case where the two sets coincide, because
prev is then an alias of curr -- which is what the follow-up handles.

Comment only. No behaviour change.

Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
@tianyi-zhang-02
tianyi-zhang-02 force-pushed the fix/neg-inf-narrowing-on-policy branch from 1b090c3 to 1ec5155 Compare August 28, 2026 15:29
…_logprobs

force_on_policy_ratio sets `prev_logprobs = curr_logprobs.detach()` before the
mismatch diagnostics run. The substituted 0.0 that curr_logprobs carries under
top-k/top-p filtering is therefore in prev too, so the reductions that read
prev must narrow with the keep mask -- the comment justifying the full mask
("the mismatch diagnostics read only prev/generation logprobs") does not hold
once prev is an alias for curr.

0.0 means log p = 0, i.e. p = 1, so each filtered position contributes
exp(|log pi_gen|). Measured on a 3-token microbatch with pi_gen = exp(-5) at
the filtered position: token_mult_prob_error 50.14 instead of 1.0, and under
sequence_level_importance_ratios the whole sequence is weighted exp(5) = 148x.
That last one scales the gradient, since the per-sequence scalar multiplies
every token and the actor_mask reduction cannot undo it.

Adds prev_mask/prev_token_mask, equal to the existing masks unless
force_on_policy_ratio is set, and routes the four mismatch diagnostics, the
sequence-level IS weight, the TIS/icepop out-of-bounds fractions, the
seq-mask-tis gate and sampling_importance_ratio through them. With
force_on_policy_ratio unset the reductions are unchanged.

Six tests assert the filtered position no longer steers the result, by varying
generation_logprobs there and requiring the metric to hold still.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
@tianyi-zhang-02
tianyi-zhang-02 force-pushed the fix/neg-inf-narrowing-on-policy branch from 1ec5155 to 006907a Compare August 28, 2026 17:43
Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
Emit actor-mask diagnostic numerators with their retained-token count, then finalize the global ratio after microbatch and rank aggregation. This keeps KL on the full valid-token mask while making filtered ratio and entropy metrics invariant to uneven shards.

Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
Resolve current-main conflicts and normalize narrowed diagnostic metrics with actual surviving-token counts across every trainer caller.

Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
Reuse the parent retained-token aggregation protocol and extend it to forced-on-policy prev reductions.

Signed-off-by: Tianyi Zhang <123608656+tianyi-zhang-02@users.noreply.github.com>
@tianyi-zhang-02 tianyi-zhang-02 changed the title fix(loss): narrow the prev-side reductions when prev_logprobs is curr_logprobs fix(loss): mask filtered tokens in force-on-policy reductions Sep 7, 2026
@tianyi-zhang-02

Copy link
Copy Markdown
Contributor Author

Closing to avoid duplicating the active top-k/top-p masking work in #4154 and the masked-value safety fix in #4120. This branch also still needs a denominator redesign, so keeping a parallel implementation open would make the intended loss semantics less clear.

@svcnvidia-nemo-ci svcnvidia-nemo-ci removed the waiting-on-maintainers Waiting on maintainers to respond label Sep 25, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants