From 3f21d78ef796bae7d31a07bdf7e8c25c8fb1a197 Mon Sep 17 00:00:00 2001 From: nathon-lee Date: Sat, 26 Sep 2026 14:22:22 +0000 Subject: [PATCH 1/2] feat(rollout): add OPSD response token loss primitives Derive response-only causal targets from RolloutBatch and provide generalized JSD loss statistics for distributed training. Signed-off-by: nathon-lee --- deepspeed/runtime/rollout/__init__.py | 5 + deepspeed/runtime/rollout/opsd.py | 133 ++++++++++++++++++ docs/code-docs/source/inference-engine.rst | 29 ++++ .../runtime/rollout/test_rollout_interface.py | 72 ++++++++++ 4 files changed, 239 insertions(+) create mode 100644 deepspeed/runtime/rollout/opsd.py diff --git a/deepspeed/runtime/rollout/__init__.py b/deepspeed/runtime/rollout/__init__.py index 16f6fc595da6..a7af574ff419 100644 --- a/deepspeed/runtime/rollout/__init__.py +++ b/deepspeed/runtime/rollout/__init__.py @@ -6,6 +6,7 @@ Provides: - :class:`RolloutEngine` — abstract base class - :class:`RolloutRequest`, :class:`RolloutBatch`, :class:`SamplingConfig` — dataclasses + - :class:`ResponseTokenBatch` — response-only causal training targets - :class:`HybridEngineRollout` — concrete implementation using DeepSpeed hybrid engine - :func:`build_rollout` — factory that selects the engine from config """ @@ -18,15 +19,19 @@ SamplingConfig, ) from deepspeed.runtime.rollout.hybrid_engine_rollout import HybridEngineRollout +from deepspeed.runtime.rollout.opsd import JSDLossOutput, ResponseTokenBatch, generalized_jsd_loss __all__ = [ "HybridEngineRollout", + "JSDLossOutput", "RolloutBatch", "RolloutConfig", "RolloutEngine", "RolloutRequest", + "ResponseTokenBatch", "SamplingConfig", "build_rollout", + "generalized_jsd_loss", ] diff --git a/deepspeed/runtime/rollout/opsd.py b/deepspeed/runtime/rollout/opsd.py new file mode 100644 index 000000000000..fa58b0a3041f --- /dev/null +++ b/deepspeed/runtime/rollout/opsd.py @@ -0,0 +1,133 @@ +# SPDX-License-Identifier: Apache-2.0 + +# DeepSpeed Team +"""Response-token batching and distillation losses for OPSD-style training.""" + +from dataclasses import dataclass +from typing import Optional + +import torch +import torch.nn.functional as F + +from deepspeed.runtime.rollout.base import RolloutBatch + + +@dataclass(frozen=True) +class ResponseTokenBatch: + """Causal training targets derived from an on-policy rollout. + + ``target_ids`` and ``response_mask`` are shifted by one position relative + to ``input_ids``: logits at position ``i`` predict ``target_ids[:, i]``. + """ + + input_ids: torch.Tensor + attention_mask: torch.Tensor + target_ids: torch.Tensor + response_mask: torch.Tensor + + @classmethod + def from_rollout(cls, rollout: RolloutBatch) -> "ResponseTokenBatch": + """Derive response-only causal targets from ``rollout``.""" + if rollout.input_ids.device != rollout.attention_mask.device: + raise ValueError("rollout input_ids and attention_mask must be on the same device") + if rollout.input_ids.device != rollout.response_start_idx.device: + raise ValueError("rollout input_ids and response_start_idx must be on the same device") + if rollout.response_start_idx.dtype not in (torch.int32, torch.int64): + raise ValueError("rollout response_start_idx must use an integer dtype") + + sequence_length = rollout.input_ids.shape[1] + response_start_idx = rollout.response_start_idx + if (response_start_idx < 1).any() or (response_start_idx > sequence_length).any(): + raise ValueError("rollout response_start_idx must be in [1, sequence_length]") + + target_positions = torch.arange(1, sequence_length, device=rollout.input_ids.device) + response_positions = target_positions.unsqueeze(0) >= response_start_idx.unsqueeze(1) + response_mask = rollout.attention_mask[:, 1:].bool() & response_positions + return cls( + input_ids=rollout.input_ids, + attention_mask=rollout.attention_mask, + target_ids=rollout.input_ids[:, 1:], + response_mask=response_mask, + ) + + def select_causal_logits(self, logits: torch.Tensor) -> torch.Tensor: + """Align model logits with this batch's shifted response targets.""" + if logits.dim() != 3: + raise ValueError("logits must be 3-D [batch, sequence_length, vocabulary_size]") + if logits.shape[:2] != self.input_ids.shape: + raise ValueError("logits batch and sequence dimensions must match input_ids") + return logits[:, :-1, :] + + +@dataclass(frozen=True) +class JSDLossOutput: + """JSD loss and its unreduced statistics for distributed aggregation.""" + + loss: torch.Tensor + loss_sum: torch.Tensor + valid_token_count: torch.Tensor + + +def generalized_jsd_loss(student_logits: torch.Tensor, + teacher_logits: torch.Tensor, + response_mask: torch.Tensor, + beta: float = 0.5, + temperature: float = 1.0, + teacher_top_k: Optional[int] = None, + pointwise_clip: Optional[float] = None) -> JSDLossOutput: + """Compute generalized JSD on valid response tokens. + + ``beta=0`` yields forward KL from teacher to student; ``beta=1`` yields + reverse KL from student to teacher. A positive ``teacher_top_k`` restricts + both distributions to the teacher's most likely vocabulary entries before + they are normalized. ``pointwise_clip`` caps individual vocabulary-level + divergence contributions before their per-token reduction. + """ + if student_logits.shape != teacher_logits.shape or student_logits.dim() != 3: + raise ValueError("student_logits and teacher_logits must have matching 3-D shapes") + if response_mask.shape != student_logits.shape[:2]: + raise ValueError("response_mask shape must match logits batch and sequence dimensions") + if not 0.0 <= beta <= 1.0: + raise ValueError("beta must be in [0, 1]") + if temperature <= 0.0: + raise ValueError("temperature must be positive") + if teacher_top_k is not None and teacher_top_k <= 0: + raise ValueError("teacher_top_k must be positive when specified") + if teacher_top_k is not None and teacher_top_k > student_logits.shape[-1]: + raise ValueError("teacher_top_k must not exceed the vocabulary size") + if pointwise_clip is not None and pointwise_clip < 0.0: + raise ValueError("pointwise_clip must be non-negative when specified") + + student_logits = student_logits / temperature + teacher_logits = teacher_logits / temperature + if teacher_top_k is not None: + teacher_top_indices = torch.topk(teacher_logits, k=teacher_top_k, dim=-1).indices + student_logits = torch.gather(student_logits, dim=-1, index=teacher_top_indices) + teacher_logits = torch.gather(teacher_logits, dim=-1, index=teacher_top_indices) + + student_log_probs = F.log_softmax(student_logits, dim=-1) + teacher_log_probs = F.log_softmax(teacher_logits, dim=-1) + student_probs = student_log_probs.exp() + teacher_probs = teacher_log_probs.exp() + + if beta == 0.0: + pointwise_jsd = teacher_probs * (teacher_log_probs - student_log_probs) + elif beta == 1.0: + pointwise_jsd = student_probs * (student_log_probs - teacher_log_probs) + else: + mixture_log_probs = torch.logaddexp( + student_log_probs + + torch.log1p(torch.tensor(-beta, dtype=student_logits.dtype, device=student_logits.device)), + teacher_log_probs + + torch.log(torch.tensor(beta, dtype=student_logits.dtype, device=student_logits.device))) + pointwise_jsd = beta * teacher_probs * (teacher_log_probs - mixture_log_probs) + pointwise_jsd += (1.0 - beta) * student_probs * (student_log_probs - mixture_log_probs) + + if pointwise_clip is not None: + pointwise_jsd = pointwise_jsd.clamp(max=pointwise_clip) + + response_mask = response_mask.to(dtype=pointwise_jsd.dtype) + loss_sum = (pointwise_jsd.sum(dim=-1) * response_mask).sum() + valid_token_count = response_mask.sum().to(dtype=torch.long) + loss = loss_sum / valid_token_count.clamp_min(1) + return JSDLossOutput(loss=loss, loss_sum=loss_sum, valid_token_count=valid_token_count) diff --git a/docs/code-docs/source/inference-engine.rst b/docs/code-docs/source/inference-engine.rst index 7ae7b89bebc6..b269b3415e50 100644 --- a/docs/code-docs/source/inference-engine.rst +++ b/docs/code-docs/source/inference-engine.rst @@ -110,3 +110,32 @@ Transformers. active rows while preserving its static tensor addresses. This mirrors the scheduler/cache separation used by systems such as vLLM and SGLang without copying their backend-specific kernels. + +OPSD response-token utilities +----------------------------- + +``ResponseTokenBatch.from_rollout`` derives causal training targets and a +response-only loss mask from a ``RolloutBatch``. The mask excludes prompt and +padding tokens while retaining an attended EOS token. Use +``generalized_jsd_loss`` to compare aligned student and teacher logits on that +mask: + +.. code-block:: python + + from deepspeed.runtime.rollout import ResponseTokenBatch, generalized_jsd_loss + + response_batch = ResponseTokenBatch.from_rollout(rollout_batch) + student_logits = response_batch.select_causal_logits(student_model(**model_inputs).logits) + teacher_logits = aligned_teacher_logits + loss_output = generalized_jsd_loss( + student_logits, + teacher_logits, + response_batch.response_mask, + ) + loss = loss_output.loss + +The rollout supplies only the student on-policy trajectory. Construct the +teacher inputs, including any privileged context, separately and align the +teacher logits with the student causal-logit shape before calling the loss. +``loss_sum`` and ``valid_token_count`` are available for callers that need +explicit distributed loss normalization. diff --git a/tests/unit/runtime/rollout/test_rollout_interface.py b/tests/unit/runtime/rollout/test_rollout_interface.py index 7a94c01cc40b..a70d7c0a6e11 100644 --- a/tests/unit/runtime/rollout/test_rollout_interface.py +++ b/tests/unit/runtime/rollout/test_rollout_interface.py @@ -13,11 +13,14 @@ import deepspeed.runtime.rollout as rollout_module from deepspeed.runtime.rollout import ( + JSDLossOutput, RolloutBatch, RolloutEngine, RolloutRequest, + ResponseTokenBatch, SamplingConfig, build_rollout, + generalized_jsd_loss, ) # --- dataclass invariants --------------------------------------------------- @@ -70,6 +73,75 @@ def test_continuous_batching_details_are_not_public_exports(): assert "ContinuousBatchUpdate" not in rollout_module.__all__ +def test_response_token_batch_masks_prompts_and_padding(): + rollout = RolloutBatch( + input_ids=torch.tensor([[0, 11, 12, 21, 2, 0], [13, 14, 15, 22, 23, 2]]), + attention_mask=torch.tensor([[0, 1, 1, 1, 1, 0], [1, 1, 1, 1, 1, 1]]), + response_start_idx=torch.tensor([3, 3]), + ) + + response_batch = ResponseTokenBatch.from_rollout(rollout) + + assert torch.equal(response_batch.input_ids, rollout.input_ids) + assert response_batch.target_ids.tolist() == [[11, 12, 21, 2, 0], [14, 15, 22, 23, 2]] + assert response_batch.response_mask.tolist() == [[False, False, True, True, False], + [False, False, True, True, True]] + + +@pytest.mark.parametrize("response_start_idx", [torch.tensor([0]), torch.tensor([4]), torch.tensor([2.0])]) +def test_response_token_batch_rejects_invalid_response_boundaries(response_start_idx): + rollout = RolloutBatch( + input_ids=torch.tensor([[1, 2, 3]]), + attention_mask=torch.ones((1, 3), dtype=torch.long), + response_start_idx=response_start_idx, + ) + + with pytest.raises(ValueError): + ResponseTokenBatch.from_rollout(rollout) + + +def test_response_token_batch_aligns_causal_logits(): + rollout = RolloutBatch( + input_ids=torch.tensor([[1, 2, 3]]), + attention_mask=torch.ones((1, 3), dtype=torch.long), + response_start_idx=torch.tensor([2]), + ) + logits = torch.randn(1, 3, 5) + + response_batch = ResponseTokenBatch.from_rollout(rollout) + + assert torch.equal(response_batch.select_causal_logits(logits), logits[:, :-1]) + with pytest.raises(ValueError, match="sequence"): + response_batch.select_causal_logits(torch.randn(1, 2, 5)) + + +def test_generalized_jsd_loss_matches_forward_kl_on_response_tokens(): + student_logits = torch.log(torch.tensor([[[0.8, 0.2], [0.5, 0.5]]])) + teacher_logits = torch.log(torch.tensor([[[0.5, 0.5], [0.9, 0.1]]])) + response_mask = torch.tensor([[True, False]]) + + output = generalized_jsd_loss(student_logits, teacher_logits, response_mask, beta=0.0) + + expected = 0.5 * torch.log(torch.tensor(0.5 / 0.8)) + 0.5 * torch.log(torch.tensor(0.5 / 0.2)) + assert isinstance(output, JSDLossOutput) + assert output.valid_token_count.item() == 1 + assert torch.allclose(output.loss_sum, expected) + assert torch.allclose(output.loss, expected) + + +def test_generalized_jsd_loss_supports_top_k_and_empty_response_masks(): + student_logits = torch.tensor([[[20.0, 0.0, 0.0]]]) + teacher_logits = torch.tensor([[[0.0, 10.0, 9.0]]]) + response_mask = torch.tensor([[True]]) + + top_one = generalized_jsd_loss(student_logits, teacher_logits, response_mask, teacher_top_k=1) + empty = generalized_jsd_loss(student_logits, teacher_logits, torch.tensor([[False]])) + + assert top_one.loss.item() == pytest.approx(0.0) + assert empty.valid_token_count.item() == 0 + assert empty.loss.item() == pytest.approx(0.0) + + # --- interface conformance via FakeRollout --------------------------------- From 9da5134fa36f09620eb81514e349ec8d8bf20810 Mon Sep 17 00:00:00 2001 From: nathon-lee Date: Tue, 6 Oct 2026 23:48:40 +0800 Subject: [PATCH 2/2] tests: expand generalized_jsd_loss beta coverage Signed-off-by: nathon-lee --- .../runtime/rollout/test_rollout_interface.py | 28 +++++++++++++++++-- 1 file changed, 25 insertions(+), 3 deletions(-) diff --git a/tests/unit/runtime/rollout/test_rollout_interface.py b/tests/unit/runtime/rollout/test_rollout_interface.py index a70d7c0a6e11..e760a75e55d2 100644 --- a/tests/unit/runtime/rollout/test_rollout_interface.py +++ b/tests/unit/runtime/rollout/test_rollout_interface.py @@ -115,14 +115,36 @@ def test_response_token_batch_aligns_causal_logits(): response_batch.select_causal_logits(torch.randn(1, 2, 5)) -def test_generalized_jsd_loss_matches_forward_kl_on_response_tokens(): +@pytest.mark.parametrize( + "beta, expected", + [ + ( + 0.0, + 0.5 * torch.log(torch.tensor(0.5 / 0.8)) + 0.5 * torch.log(torch.tensor(0.5 / 0.2)), + ), + ( + 1.0, + 0.8 * torch.log(torch.tensor(0.8 / 0.5)) + 0.2 * torch.log(torch.tensor(0.2 / 0.5)), + ), + ( + 0.5, + 0.5 * ( + 0.5 * torch.log(torch.tensor(0.5 / 0.65)) + + 0.5 * torch.log(torch.tensor(0.5 / 0.35)) + ) + 0.5 * ( + 0.8 * torch.log(torch.tensor(0.8 / 0.65)) + + 0.2 * torch.log(torch.tensor(0.2 / 0.35)) + ), + ), + ], +) +def test_generalized_jsd_loss_covers_beta_variants(beta, expected): student_logits = torch.log(torch.tensor([[[0.8, 0.2], [0.5, 0.5]]])) teacher_logits = torch.log(torch.tensor([[[0.5, 0.5], [0.9, 0.1]]])) response_mask = torch.tensor([[True, False]]) - output = generalized_jsd_loss(student_logits, teacher_logits, response_mask, beta=0.0) + output = generalized_jsd_loss(student_logits, teacher_logits, response_mask, beta=beta) - expected = 0.5 * torch.log(torch.tensor(0.5 / 0.8)) + 0.5 * torch.log(torch.tensor(0.5 / 0.2)) assert isinstance(output, JSDLossOutput) assert output.valid_token_count.item() == 1 assert torch.allclose(output.loss_sum, expected)