Skip to content

feat(rollout): add OPSD response token loss primitives - #8677

Open
nathon-lee wants to merge 3 commits into
deepspeedai:masterfrom
nathon-lee:feat/opsd-response-token-loss
Open

nathon-lee wants to merge 3 commits into
deepspeedai:masterfrom
nathon-lee:feat/opsd-response-token-loss

Conversation

@nathon-lee

@nathon-lee nathon-lee commented Sep 26, 2026 •

Copy link
Copy Markdown
Contributor

Summary

This PR adds phase-one OPSD training primitives for deriving response-only causal targets from RolloutBatch and computing generalized JSD on valid response tokens.

Why this is needed

RolloutBatch already 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

  • Adds ResponseTokenBatch.from_rollout() to derive shifted target IDs and a response-only mask from RolloutBatch.
  • Adds generalized_jsd_loss() with configurable JSD mixing, temperature, teacher top-k vocabulary restriction, and pointwise clipping.
  • Returns loss_sum and valid_token_count in addition to the normalized loss so distributed callers can choose correct global normalization.
  • Documents the public API and validates left padding, response boundaries, EOS/padding masking, causal-logit alignment, empty responses, and top-k loss restriction.

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 passed
  • Changed-file pre-commit checks: passed

Thank you for your review.

Derive response-only causal targets from RolloutBatch and provide generalized JSD loss statistics for distributed training.

Signed-off-by: nathon-lee <leejianwoo@gmail.com>
Comment thread tests/unit/runtime/rollout/test_rollout_interface.py Outdated
@delock

delock commented Oct 5, 2026

Copy link
Copy Markdown
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!

@nathon-lee

Copy link
Copy Markdown
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.
For the DeepSpeedExamples question, this PR mainly provides the reusable rollout/loss primitives, so example simplification can likely be explored in a follow-up.

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.

2 participants