Repository navigation
feat(rollout): add OPSD response token loss primitives - #8677
Open
nathon-lee wants to merge 3 commits into
Open
nathon-lee wants to merge 3 commits into
nathon-lee wants to merge 3 commits into
Conversation
Derive response-only causal targets from RolloutBatch and provide generalized JSD loss statistics for distributed training. Signed-off-by: nathon-lee <leejianwoo@gmail.com>
nathon-lee
requested review from
loadams,
tjruwase and
tohtana
as code owners
September 26, 2026 14:26
delock
self-requested a review
September 28, 2026 14:40
delock
reviewed
Oct 5, 2026
Collaborator
|
Hi @nathon-lee , this PR overall looks good to me except one test coverage. One question that might be related, this PR provide additional support for OPD/OPSD support, does it mean the examples in DeepSpeedExamples can be simplified when this PR is merged? Thanks! |
Contributor
Author
|
Hi @delock, Thanks for the review — good point. I’ll add test coverage for beta=1.0 and beta=0.5 as well. |
Signed-off-by: nathon-lee <leejianwoo@gmail.com>
This branch has not been deployed
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.
Summary
This PR adds phase-one OPSD training primitives for deriving response-only causal targets from
RolloutBatchand computing generalized JSD on valid response tokens.Why this is needed
RolloutBatchalready preserves the student on-policy token trajectory, attention mask, EOS handling, and response boundary. Before this change, training code had no shared way to turn that contract into causal targets and a response-only loss mask. Each consumer would need to reimplement the causal shift and boundary logic, risking accidental inclusion of prompt or padding tokens and inconsistent distributed loss normalization.OPSD needs this exact boundary: it compares student and teacher distributions only on the student's generated response tokens. These utilities establish that reusable, model-agnostic layer without coupling rollout code to a trainer or a particular teacher implementation.
Changes
ResponseTokenBatch.from_rollout()to derive shifted target IDs and a response-only mask fromRolloutBatch.generalized_jsd_loss()with configurable JSD mixing, temperature, teacher top-k vocabulary restriction, and pointwise clipping.loss_sumandvalid_token_countin addition to the normalized loss so distributed callers can choose correct global normalization.Scope
The rollout remains the source of the student on-policy trajectory. Teacher-context construction, privileged solution inputs, and cross-prefix teacher-logit alignment are intentionally deferred to a follow-up. This keeps the first phase small and independently useful for future OPSD or response-token distillation integrations.
Validation
tests/unit/runtime/rollout/test_rollout_interface.py: 18 passedThank you for your review.