Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions deepspeed/runtime/rollout/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
"""
Expand All @@ -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",
]


Expand Down
133 changes: 133 additions & 0 deletions deepspeed/runtime/rollout/opsd.py
Original file line number Diff line number Diff line change
@@ -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)
29 changes: 29 additions & 0 deletions docs/code-docs/source/inference-engine.rst
Original file line number Diff line number Diff line change
Expand Up @@ -142,3 +142,32 @@ capacity-exhaustion fallback reclaims a dead prefix.
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.
94 changes: 94 additions & 0 deletions tests/unit/runtime/rollout/test_rollout_interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ---------------------------------------------------
Expand Down Expand Up @@ -70,6 +73,97 @@ 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))


@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=beta)

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


Expand Down
Loading