Skip to content

Support filtered sampling in score centering - #3641

Closed
Shi-Dong wants to merge 6 commits into
shi/score-centering-docsfrom
shi/score-centering-filtered-sampling
Closed

Shi-Dong wants to merge 6 commits into
shi/score-centering-docsfrom
shi/score-centering-filtered-sampling

Conversation

@Shi-Dong

@Shi-Dong Shi-Dong commented Sep 23, 2026 •

Copy link
Copy Markdown
Collaborator

Summary

Add top-p and top-k filtered sampling support to score centering, building on the recorded sampling support from #3354 and SGLang support log probabilities from #40932 (cherry-picked into sglang-miles by #41047).

  • Request sampling_logprobs_mode="support" and store the sampler's post-filter probabilities for every surviving token. Select the sampled token's probability from the same response.
  • Normalize trainer probabilities over the recorded support for the sampled token, centering correction, and entropy diagnostic.
  • Require positive sampling top-k no greater than the saved candidate count; reject missing, misaligned, or incomplete support probabilities.
  • Keep the existing top-logprobs path for unfiltered sampling. Update native and session rollout handling, tests, and documentation.

Validation

99 focused tests passed; 6 GPU-dependent cases skipped on the local Mac. Ruff, Black, and diff checks passed. The session integration test was updated but could not be collected locally because SGLang is not installed in this environment. Live GPU training remains untested.

Stacked on #3626.

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Claude Code Review

This repository is configured for manual code reviews. Comment @claude review for a one-time review, or @claude review always to subscribe this PR to a review on every future push.

Tip: disable this comment in your organization's Code Review settings.

@Shi-Dong

Copy link
Copy Markdown
Collaborator Author

@claude review always

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Code review found no issues

No high-confidence issues detected in this change.

@Shi-Dong
Shi-Dong force-pushed the shi/score-centering-docs branch from 53c51bf to 170f508 Compare September 24, 2026 02:54
@Shi-Dong
Shi-Dong force-pushed the shi/score-centering-filtered-sampling branch from bb34e90 to abd8bae Compare September 24, 2026 02:54

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

I reviewed this PR (including the commit pushed since my last review) and didn't find any bugs. Because it touches training-critical gradient/renormalization logic for score-centering under sampling-support replay, a human look would still be worthwhile.

What was reviewed: the replay-branch renormalization in score_centering_loss.py (head-logsumexp applied to both the sampled-token column and candidate columns) against the new dense-gradient reference test in test_score_centering_filtered.py, which matches an independent analytic gradient (including the entropy term) to 2e-6 — this rules out a biased-gradient concern I looked at closely. Also checked the vectorized _candidate_support_membership sorted-key matching for support coverage/duplicate detection and the widened validate_score_centering_sampling bounds in miles/utils/score_centering.py.

Extended reasoning...

The diff (score_centering_loss.py, generate_utils/score_centering.py, utils/score_centering.py plus docs/tests) extends score centering to support bounded top-p/top-k filtered sampling by restricting recorded candidates to the realized sampler support and renormalizing both training-side logprobs and the entropy diagnostic over that support; it touches no auth/injection/data-exposure surface but is numerically subtle RL-training logic where a silent bug would bias gradients rather than crash. The PR author (Shi-Dong) is a CODEOWNER of all changed miles/ paths, and a new commit landed after my prior "no issues" review on this PR, so this is a re-check of updated code; the new test_filtered_loss_matches_dense_support_gradient test gave concrete, verifiable evidence ruling out the specific bias concern investigated this run, but the overall change remains complex enough to warrant a human pass before merge.

@Shi-Dong
Shi-Dong force-pushed the shi/score-centering-filtered-sampling branch from abd8bae to e859915 Compare September 24, 2026 07:48
@Shi-Dong
Shi-Dong force-pushed the shi/score-centering-docs branch from 170f508 to 9ccf9df Compare September 24, 2026 08:07
@Shi-Dong Shi-Dong closed this Sep 24, 2026

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Beyond the inline findings, I also checked whether a custom rollout producer that fills only rollout_topk_token_ids/rollout_topk_log_probs/rollout_log_probs (the documented contract) but leaves rollout_sampling_mask unset could silently corrupt use_sampling_support_replay's renormalization in score_centering_loss.py — it can't, since that renormalization only reads the topk candidate arrays and never touches rollout_sampling_mask (that field feeds the separate policy_loss sampling-mask path).

Extended reasoning...

This run's inline findings (a TypeError crash on zero-token aborts in filtered mode, and a spurious startup validation error for filtered sampling with session-server + large top-k) already signal that human review is warranted, so no new top-level verdict is needed. I additionally verified one candidate concern beyond those findings — that omitting rollout_sampling_mask for a custom producer could silently break the new use_sampling_support_replay renormalization — and confirmed that path never reads rollout_sampling_mask, so it's a non-issue worth recording briefly since this is new information not covered by my prior "no issues" review on an earlier push of this PR.

Additional findings (outside the current diff — GitHub can't attach inline comments there):

  • 🔴 miles/utils/score_centering.py — Operators using filtered score-centering (top_p<1 or top_k>0) with --score-centering-top-k>128 (default) and --use-session-server now get a spurious startup ValueError forcing --use-miles-router, even though filtered mode never needs it. validate_score_centering_args (score_centering.py:76-84) still applies the '>20 candidates requires --use-miles-router' guard unconditionally, but configure_score_centering_request (generate_utils/score_centering.py:25-29) already made filtered requests skip top_logprobs entirely and use sampling_logprobs_mode='support' instead, so the router's OpenAI top_logprobs cap this guard exists for cannot fire in that mode. …

    Why this was flagged

    …This directly contradicts the PR's own new docs ('Filtered rollouts request support probabilities instead of top_logprobs and do not use that cap'). Fix: skip this check when args.rollout_top_p<1.0 or args.rollout_top_k>0 (i.e. filtered/support mode), keeping it only for the unfiltered top_logprobs path.

    Trigger: launch with --loss-type score_centering, --score-centering-top-k 128 (the documented default), --rollout-top-p 0.9 --rollout-top-k 64 (filtered, now permitted by this diff), --use-session-server, and no --use-miles-router. validate_score_centering_args at miles/utils/score_centering.py:76-84 raises 'Score-centering session rollouts with more than 20 candidates require --use-miles-router: the SGLang Rust router caps OpenAI top_logprobs at 20', aborting the run at startup. But configure_score_centering_request (miles/rollout/generate_utils/score_centering.py:25-29) never sends top_logprobs for this config; it sets sampling_logprobs_mode='support' and requests return_sampling_mask instead, so the cited router cap does not apply. Before this diff, score centering only supported…

    Verification: normal (contradicts the docs shipped in this same PR). The >20 candidates require --use-miles-router guard in validate_score_centering_args (miles/utils/score_centering.py:76-84) fires purely on args.score_centering_top_k > 20 and use_session_server and not use_miles_router — it never inspects args.rollout_top_p/args.rollout_top_k, both of which are available in this same function…

Comment on lines +52 to +60
support_ids = meta.get("output_token_sampling_mask") if support_mode else None
rows = meta.get("output_token_sampling_logprobs" if support_mode else "output_top_logprobs")
if support_mode and n and (sample.rollout_sampling_mask is None or support_ids is None):
raise ValueError("Score centering requires SGLang sampling support IDs")
if rows is None and n:
field = "output_token_sampling_logprobs" if support_mode else "output_top_logprobs"
raise ValueError(f"Score centering requires SGLang {field} from generation")
rows = rows if rows is not None else []
if len(rows) != n or (support_mode and len(support_ids) != n):

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

🔴 Filtered score-centering (sampling_logprobs_mode="support") crashes with an unhandled TypeError, instead of a clean error, when a turn aborts with zero generated tokens. append_score_centering_topk only guards the rollout_sampling_mask/support_ids None-check at line 54 with support_mode and n, so when n==0 and meta_info lacks output_token_sampling_mask, support_ids stays None and line 60's len(support_ids) != n raises TypeError: object of type 'NoneType' has no len(). append_sampling_metadata already special-cases this abort/empty case; this sibling function was not. …

Why this was flagged

…Fix: when n==0, skip or explicitly tolerate support_ids being None (mirror append_sampling_metadata's abort handling) so an abort with 0 tokens under filtered sampling fails cleanly or no-ops instead of raising an unhandled TypeError.

Trigger: sampling_logprobs_mode="support" (filtered top-p/top-k score centering) plus a partial-rollout retract/abort where SGLang's meta_info has no output_token_logprobs and no output_token_sampling_mask, so n=0 and support_ids=None. score_centering.py:54's None-check is skipped because it is gated by n, so line 60's support_mode and len(support_ids) != n evaluates len(None) and raises TypeError instead of the intended ValueError. This function is called unconditionally from sglang_rollout.py's generate() (if not evaluation: append_score_centering_topk(...)) and gated only on sampling_logprobs_mode == "support" in generate_endpoint_utils.py and merge.py, so any abort under filtered sampling crashes the rollout task/worker instead of surfacing a controlled error, unlike the base branch where no support_mode path existed.

Verification: normal. In append_score_centering_topk (miles/rollout/generate_utils/score_centering.py), line 54 gates the None-check on n: if support_mode and n and (sample.rollout_sampling_mask is None or support_ids is None): raise ValueError(...). When a rollout aborts with zero generated tokens, n=0, so this check is skipped, and support_ids = meta.get("output_token_sampling_mask") stays None (abort…

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