fix(loss): mask filtered tokens in force-on-policy reductions - #3778
Closed
tianyi-zhang-02 wants to merge 14 commits into
Closed
tianyi-zhang-02 wants to merge 14 commits into
tianyi-zhang-02 wants to merge 14 commits into
Conversation
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>
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
force-pushed
the
fix/neg-inf-narrowing-on-policy
branch
from
August 26, 2026 21:37
2b7165b to
1b090c3
Compare
…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
force-pushed
the
fix/neg-inf-narrowing-on-policy
branch
from
August 28, 2026 15:29
1b090c3 to
1ec5155
Compare
…_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
force-pushed
the
fix/neg-inf-narrowing-on-policy
branch
from
August 28, 2026 17:43
1ec5155 to
006907a
Compare
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>
Contributor
Author
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_ratiopath, whereprev_logprobsaliasescurr_logprobs.detach().When a current logprob is filtered, its finite substitute must not leak into quantities that read
prev_logprobs. This change:seq-mask-tis;prev_logprobspath 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.
force_on_policy_ratioStack
main:e831622435180c88d6adcc9fce46bfd88a85f3e5cbfd57d454fdb8e6f963b4d241bece9c2b6f2df769b9b424876a53900054d8e9ce8c2e616074339bThe Files changed view includes #3551 until that PR lands. This branch merges both the current #3551 head and the current
mainand 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:
test_loss_functions.py,test_utils.py,test_sequence_packing_fusion.py,single_controller/test_utils.pytest_grpo.py,test_ppo.pyThe 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"
main