diff --git a/tests/models/test_value_head.py b/tests/models/test_value_head.py new file mode 100644 index 000000000..64be35e33 --- /dev/null +++ b/tests/models/test_value_head.py @@ -0,0 +1,15 @@ +"""CPU tests for ValueHead.""" + +from __future__ import annotations + +import torch + +from unirl.models.types.value_head import ValueHead + + +def test_value_head_output_shape() -> None: + head = ValueHead(hidden_size=8) + hidden = torch.randn(5, 8) + values = head(hidden) + assert values.shape == (5,) + assert values.dtype == torch.float32 diff --git a/tests/types/test_advantages_gae.py b/tests/types/test_advantages_gae.py index 2c647a4b4..5fdc1e6ad 100644 --- a/tests/types/test_advantages_gae.py +++ b/tests/types/test_advantages_gae.py @@ -7,7 +7,7 @@ import pytest import torch -from unirl.types.advantages import compute_gae_advantages +from unirl.types.advantages import compute_gae_advantages, scatter_terminal_rewards def test_gae_hand_computed_lambda_one() -> None: @@ -107,3 +107,11 @@ def test_gae_single_step() -> None: advantages, returns = compute_gae_advantages(rewards, values, gamma=1.0, gae_lambda=0.95) assert math.isclose(float(advantages.item()), 0.75, rel_tol=0, abs_tol=1e-6) assert math.isclose(float(returns.item()), 1.0, rel_tol=0, abs_tol=1e-6) + + +def test_scatter_terminal_rewards_packed() -> None: + lengths = torch.tensor([2, 1]) + cu = torch.tensor([0, 2, 3]) + rewards = torch.tensor([1.0, 0.0]) + out = scatter_terminal_rewards(rewards, lengths=lengths, cu_seqlens=cu) + assert out.tolist() == [0.0, 1.0, 0.0] diff --git a/tests/types/test_rollout_track_gae.py b/tests/types/test_rollout_track_gae.py new file mode 100644 index 000000000..5586ceabe --- /dev/null +++ b/tests/types/test_rollout_track_gae.py @@ -0,0 +1,88 @@ +"""CPU tests for scatter_terminal_rewards and track-level GAE wiring.""" + +from __future__ import annotations + +import torch + +from unirl.types.advantages import scatter_terminal_rewards +from unirl.types.rollout_resp import RolloutTrack +from unirl.types.segments.text import TextSegment + + +def test_scatter_terminal_rewards_places_reward_on_last_token() -> None: + segment = TextSegment.pack( + tokens=[torch.tensor([10, 11]), torch.tensor([20])], + values=[torch.tensor([0.2, 0.5]), torch.tensor([0.8])], + ) + assert segment.cu_seqlens is not None + assert segment.lengths is not None + rewards = torch.tensor([1.0, 0.5]) + token_rewards = scatter_terminal_rewards( + rewards, lengths=segment.lengths, cu_seqlens=segment.cu_seqlens + ) + assert token_rewards.shape == (3,) + assert token_rewards.tolist() == [0.0, 1.0, 0.5] + + +def test_rollout_track_compute_gae_advantages() -> None: + segment = TextSegment.pack( + tokens=[torch.tensor([10, 11, 12])], + values=[torch.tensor([0.2, 0.5, 0.8])], + ) + track = RolloutTrack( + sample_ids=["s0"], + rewards=torch.tensor([1.0]), + segment=segment, + ) + updated = track.compute_gae_advantages(gamma=1.0, gae_lambda=1.0) + assert updated.segment is not None + assert updated.segment.token_advantages is not None + assert updated.segment.returns is not None + assert updated.segment.token_advantages.shape == (3,) + assert updated.advantages is not None + assert updated.advantages.shape == (1,) + # Hand-check from PR1 test: sparse terminal reward on 3 tokens. + expected = torch.tensor([0.8, 0.5, 0.2]) + assert torch.allclose(updated.segment.token_advantages, expected, atol=1e-6) + + +def test_rollout_track_compute_gae_advantages_multi_sample_no_leak() -> None: + """Packed batch: GAE must not bootstrap across trajectory boundaries.""" + segment = TextSegment.pack( + tokens=[torch.tensor([10, 11]), torch.tensor([20])], + values=[torch.tensor([0.2, 0.5]), torch.tensor([0.8])], + ) + track = RolloutTrack( + sample_ids=["s0", "s1"], + rewards=torch.tensor([1.0, 0.5]), + segment=segment, + ) + updated = track.compute_gae_advantages(gamma=1.0, gae_lambda=1.0) + assert updated.segment is not None + assert updated.segment.token_advantages is not None + + track0 = RolloutTrack( + sample_ids=["s0"], + rewards=torch.tensor([1.0]), + segment=TextSegment.pack( + tokens=[torch.tensor([10, 11])], + values=[torch.tensor([0.2, 0.5])], + ), + ) + expected0 = track0.compute_gae_advantages(gamma=1.0, gae_lambda=1.0).segment + assert expected0 is not None and expected0.token_advantages is not None + + track1 = RolloutTrack( + sample_ids=["s1"], + rewards=torch.tensor([0.5]), + segment=TextSegment.pack( + tokens=[torch.tensor([20])], + values=[torch.tensor([0.8])], + ), + ) + expected1 = track1.compute_gae_advantages(gamma=1.0, gae_lambda=1.0).segment + assert expected1 is not None and expected1.token_advantages is not None + + packed_adv = updated.segment.token_advantages + assert torch.allclose(packed_adv[:2], expected0.token_advantages, atol=1e-6) + assert torch.allclose(packed_adv[2:], expected1.token_advantages, atol=1e-6) diff --git a/unirl/models/qwen3/ar.py b/unirl/models/qwen3/ar.py index ce0fb7a7d..7e392343c 100644 --- a/unirl/models/qwen3/ar.py +++ b/unirl/models/qwen3/ar.py @@ -23,7 +23,7 @@ from dataclasses import dataclass from dataclasses import field as dc_field from types import MethodType -from typing import Any, List, Optional, Tuple +from typing import Any, List, Optional, Tuple, Union import torch import torch.distributed as dist @@ -31,6 +31,7 @@ from torch.utils.checkpoint import checkpoint from unirl.models.types.ar import ARSamplingParams, ARStage, ARStep, left_pad_prompt +from unirl.models.types.replay_result import ReplayResult from unirl.types.segments import TextSegment from unirl.utils.dtypes import parse_torch_dtype @@ -87,6 +88,7 @@ def _replay_aware_forward( temperature: float = 1.0, autocast_dtype: Optional[torch.dtype] = None, packed_predict_index: Optional[torch.Tensor] = None, + return_values: bool = False, **kw: Any, ) -> Any: """Dual-mode ``forward`` installed on the Qwen3 CausalLM instance. @@ -108,6 +110,8 @@ def _replay_aware_forward( return f(self, **kw) raise RuntimeError("_replay_aware_forward: no class-level forward found in the MRO") + _require_value_head_for_replay(self, return_values) + # cuDNN's fused SDPA backward (ScaledDotProductCudnnAttentionBackward0) returns # NaN grads on some bf16 sequences while the forward stays finite (confirmed via # torch.autograd.detect_anomaly): it floods every parameter grad and forces the @@ -129,6 +133,7 @@ def _replay_aware_forward( # [B, chunk, vocab] FP32 transient stays ~1.2 GiB, and each chunk is # gradient-checkpointed (recomputed in backward rather than held). T = float(temperature) if float(temperature) > 0.0 else 1.0 + value_head = getattr(self, "value_head", None) if return_values else None if packed_predict_index is not None: # Packed varlen replay: ``hidden`` is one packed row [1, L_total, H] @@ -155,8 +160,18 @@ def _flat_logp_chunk(h: torch.Tensor, tok: torch.Tensor) -> torch.Tensor: else: flat_parts.append(_flat_logp_chunk(h, tok)) if not flat_parts: - return hidden.new_zeros((0,), dtype=torch.float32) - return torch.cat(flat_parts, dim=0) + empty = hidden.new_zeros((0,), dtype=torch.float32) + if value_head is None: + return empty + return ReplayResult(log_probs=empty, values=empty) + log_probs = torch.cat(flat_parts, dim=0) + if value_head is None: + return log_probs + value_parts: List[torch.Tensor] = [] + for s in range(0, int(h_pred.size(0)), flat_chunk): + value_parts.append(value_head(h_pred[s : s + flat_chunk])) + values = torch.cat(value_parts, dim=0) if value_parts else log_probs.new_zeros(0) + return ReplayResult(log_probs=log_probs, values=values) T_max = int(response_tokens.size(1)) resp_hidden = hidden[:, prompt_len - 1 : prompt_len - 1 + T_max, :] @@ -176,8 +191,63 @@ def _logp_chunk(h: torch.Tensor, tok: torch.Tensor) -> torch.Tensor: else: parts.append(_logp_chunk(h, tok)) if not parts: - return resp_hidden.new_zeros((bsz, 0), dtype=torch.float32) # T_max == 0 - return torch.cat(parts, dim=1) + empty = resp_hidden.new_zeros((bsz, 0), dtype=torch.float32) + if value_head is None: + return empty + return ReplayResult(log_probs=empty, values=empty) + log_probs = torch.cat(parts, dim=1) + if value_head is None: + return log_probs + value_parts = [] + for s in range(0, T_max, chunk): + value_parts.append(value_head(resp_hidden[:, s : s + chunk, :])) + values = torch.cat(value_parts, dim=1) if value_parts else log_probs.new_zeros((bsz, 0)) + return ReplayResult(log_probs=log_probs, values=values) + + +def _require_value_head_for_replay(model: Any, return_values: bool) -> None: + if return_values and getattr(model, "value_head", None) is None: + raise ValueError( + "Qwen3 replay: return_values=True requires a value head " + "(set use_value_head=True in the pipeline config)" + ) + + +def _finalize_replay_output( + out: Union[torch.Tensor, ReplayResult], + *, + segment: TextSegment, + return_values: bool, + logprob_dtype: torch.dtype, + device: torch.device, +) -> Union[torch.Tensor, ReplayResult]: + """Cast log-probs and flatten packed values to match ``segment`` layout.""" + if isinstance(out, ReplayResult): + log_probs = out.log_probs.to(dtype=logprob_dtype) + if not return_values: + return log_probs + values = out.values + if values is None: + raise ValueError("Qwen3ARStage.replay: return_values=True but critic returned no values") + if log_probs.ndim == 1: + return ReplayResult(log_probs=log_probs, values=values.to(device=device)) + if segment.cu_seqlens is None or segment.lengths is None: + raise ValueError("Qwen3ARStage.replay: segment requires cu_seqlens to flatten values") + lengths = [int(n) for n in segment.lengths.tolist()] + cu = [int(c) for c in segment.cu_seqlens.tolist()] + flat: List[torch.Tensor] = [] + for b, n in enumerate(lengths): + if n <= 0: + continue + flat.append(values[b, :n]) + packed_values = torch.cat(flat, dim=0) if flat else values.new_zeros(0, device=device) + return ReplayResult(log_probs=log_probs, values=packed_values.to(device=device)) + if return_values: + raise ValueError( + "Qwen3ARStage.replay: return_values=True but critic returned no values " + "(set use_value_head=True in the pipeline config)" + ) + return out.to(dtype=logprob_dtype) # Attention backends with a sparse packed kernel (skip cross-sequence blocks): @@ -418,22 +488,28 @@ def replay( *, segment: TextSegment, temperature: float = 1.0, - ) -> torch.Tensor: + return_values: bool = False, + ) -> Union[torch.Tensor, ReplayResult]: """Per-token log-prob replay over a stored rollout segment. Branch: prefer :meth:`packed_replay` (packed-varlen, zero padding, B > 1) and fall back to :meth:`padding_replay` (the dense ``[B, P_max + T_max]`` padded path) when packing does not apply. Returns packed varlen - ``[total_tokens]`` aligned with ``segment.log_probs``; caller controls - grad / ``.train()`` scope. ``temperature`` divides logits before - ``log_softmax`` to match SGLang's sampler (``1.0`` is a no-op). + ``[total_tokens]`` aligned with ``segment.log_probs`` unless + ``return_values=True``, in which case an :class:`ReplayResult` with + packed ``values`` is returned alongside log-probs. """ + _require_value_head_for_replay(self.model.transformer, return_values) attn_impl = getattr(getattr(self.model.transformer, "config", None), "_attn_implementation", None) if _packed_replay_supported(attn_impl): - packed = self.packed_replay(conditions, segment=segment, temperature=temperature) + packed = self.packed_replay( + conditions, segment=segment, temperature=temperature, return_values=return_values + ) if packed is not None: return packed - return self.padding_replay(conditions, segment=segment, temperature=temperature) + return self.padding_replay( + conditions, segment=segment, temperature=temperature, return_values=return_values + ) def packed_replay( self, @@ -441,7 +517,8 @@ def packed_replay( *, segment: TextSegment, temperature: float = 1.0, - ) -> Optional[torch.Tensor]: + return_values: bool = False, + ) -> Optional[Union[torch.Tensor, ReplayResult]]: """Packed-varlen replay (B > 1): zero padding anywhere. Concatenate every sample's REAL prompt tokens + its flat response tokens @@ -517,9 +594,16 @@ def packed_replay( packed_predict_index=predict_index, prompt_len=0, temperature=temperature, + return_values=return_values, autocast_dtype=(self.autocast_dtype if device.type == "cuda" else None), ) - return per_token_flat.to(dtype=self.logprob_dtype) + return _finalize_replay_output( + per_token_flat, + segment=segment, + return_values=return_values, + logprob_dtype=self.logprob_dtype, + device=device, + ) def padding_replay( self, @@ -527,7 +611,8 @@ def padding_replay( *, segment: TextSegment, temperature: float = 1.0, - ) -> torch.Tensor: + return_values: bool = False, + ) -> Union[torch.Tensor, ReplayResult]: """Dense ``[B, P_max + T_max]`` padded replay — the default / fallback path. One teacher-forced forward over padded ``prompt + response``; gather @@ -626,28 +711,54 @@ def padding_replay( # root-wrapped or plain) and never materializes [B, L, vocab] logits. # The cuda-vs-cpu autocast decision lives here; dtype validity and the # autocast scope live in the patched forward. - per_token = self.model.transformer( + out = self.model.transformer( input_ids=full_ids, attention_mask=full_mask, position_ids=position_ids, response_tokens=response_tokens, prompt_len=prompt_len, temperature=temperature, + return_values=return_values, autocast_dtype=(self.autocast_dtype if device.type == "cuda" else None), - ) # [B, T_max] FP32 + ) # [B, T_max] FP32 or ReplayResult if T_max == 0: - return torch.zeros(0, dtype=self.logprob_dtype, device=device) + empty = torch.zeros(0, dtype=self.logprob_dtype, device=device) + if return_values: + return ReplayResult(log_probs=empty, values=empty) + return empty + + if isinstance(out, ReplayResult): + per_token = out.log_probs + per_value = out.values + else: + per_token = out + per_value = None + if return_values and per_value is None: + raise ValueError( + "Qwen3ARStage.replay: return_values=True but critic returned no values " + "(set use_value_head=True in the pipeline config)" + ) - flat: List[torch.Tensor] = [] + flat_logp: List[torch.Tensor] = [] + flat_val: List[torch.Tensor] = [] for b in range(batch_size): n = lengths[b] if n == 0: continue - flat.append(per_token[b, :n]) - if not flat: - return torch.zeros(0, dtype=self.logprob_dtype, device=device) - return torch.cat(flat, dim=0).to(dtype=self.logprob_dtype) + flat_logp.append(per_token[b, :n]) + if per_value is not None: + flat_val.append(per_value[b, :n]) + if not flat_logp: + empty = torch.zeros(0, dtype=self.logprob_dtype, device=device) + if return_values: + return ReplayResult(log_probs=empty, values=empty) + return empty + log_probs = torch.cat(flat_logp, dim=0).to(dtype=self.logprob_dtype) + if not return_values: + return log_probs + values = torch.cat(flat_val, dim=0) + return ReplayResult(log_probs=log_probs, values=values) def _resolve_stop_ids( self, diff --git a/unirl/models/qwen3/bundle.py b/unirl/models/qwen3/bundle.py index 3c13308a8..4b2182d39 100644 --- a/unirl/models/qwen3/bundle.py +++ b/unirl/models/qwen3/bundle.py @@ -28,6 +28,7 @@ from unirl.models.types.bundle import Bundle from unirl.models.types.meta_init import build_meta_init_transformer +from unirl.models.types.value_head import ValueHead from unirl.utils.dtypes import parse_torch_dtype from .config import Qwen3PipelineConfig @@ -115,6 +116,10 @@ def from_config(cls, config: Qwen3PipelineConfig) -> "Qwen3Bundle": if tokenizer.pad_token is None and tokenizer.eos_token is not None: tokenizer.pad_token = tokenizer.eos_token + if config.use_value_head: + hidden_size = int(getattr(transformer.config, "hidden_size")) + transformer.value_head = ValueHead(hidden_size).to(device) + bundle = cls( transformer=transformer, tokenizer=tokenizer, diff --git a/unirl/models/qwen3/config.py b/unirl/models/qwen3/config.py index c6c8a1c73..129c327b2 100644 --- a/unirl/models/qwen3/config.py +++ b/unirl/models/qwen3/config.py @@ -69,6 +69,9 @@ class Qwen3PipelineConfig: use_lora: bool = False lora_target_modules: Optional[List[str]] = None + # Attach a scalar value head on the transformer for PPO / GAE training. + use_value_head: bool = False + system_instruction: Optional[str] = None # Chat-template thinking switch; MUST agree with the rollout engine's # chat_template_kwargs.enable_thinking or train/rollout prompts diverge. diff --git a/unirl/models/types/replay_result.py b/unirl/models/types/replay_result.py index 3bf7a7222..f7e262105 100644 --- a/unirl/models/types/replay_result.py +++ b/unirl/models/types/replay_result.py @@ -8,10 +8,9 @@ Diffusion stages populate ``log_probs`` and ``prev_sample_means`` (the mean of the SDE Gaussian — μ_θ — used as the second moment in the KL -penalty). AR stages currently return a plain ``Tensor`` (signature -divergence with diffusion is intentional for now); when AR replay grows -``logits``-based KL support, the ``logits`` field on this result will be -the canonical home. +penalty). AR stages return a plain ``Tensor`` for GRPO-style replay, or a +:class:`ReplayResult` when optional critic ``values`` (or future ``logits``) +are requested. """ from __future__ import annotations @@ -28,8 +27,9 @@ class ReplayResult: others are stage-specific and may be ``None``.""" log_probs: torch.Tensor - """Aligned with ``segment.sde_logp`` (or its slice when ``step_indices`` - subsets). Shape ``[B, S']`` for diffusion replay.""" + """Aligned with ``segment.sde_logp`` / ``segment.log_probs`` (or a slice + when ``step_indices`` subsets). Shape ``[B, S']`` for diffusion replay; + packed ``[total_tokens]`` for AR varlen replay.""" prev_sample_means: Optional[torch.Tensor] = None """The SDE transition's mean μ_θ at each replayed step. Shape @@ -42,5 +42,10 @@ class ReplayResult: or entropy penalty support; not needed for Binary KL (which uses only per-token log-probs). Currently not populated.""" + values: Optional[torch.Tensor] = None + """Per-token critic predictions ``V_t`` from replay. Shape ``[B, T]`` or + packed ``[total_tokens]`` for AR. Used by PPO / GAE training paths. + ``None`` when the stage does not attach a value head.""" + __all__ = ["ReplayResult"] diff --git a/unirl/models/types/value_head.py b/unirl/models/types/value_head.py new file mode 100644 index 000000000..5a2d1969e --- /dev/null +++ b/unirl/models/types/value_head.py @@ -0,0 +1,24 @@ +"""Scalar value head for PPO-style critic training on AR hidden states.""" + +from __future__ import annotations + +import torch +import torch.nn as nn + + +class ValueHead(nn.Module): + """Linear critic ``V(h)`` on last hidden states. + + Kept in FP32 for stable value loss math (mirrors replay log-prob FP32 policy). + """ + + def __init__(self, hidden_size: int) -> None: + super().__init__() + self.proj = nn.Linear(hidden_size, 1, bias=True, dtype=torch.float32) + + def forward(self, hidden: torch.Tensor) -> torch.Tensor: + """Map ``[..., H]`` hidden states to ``[...,]`` scalar values.""" + return self.proj(hidden.float()).squeeze(-1) + + +__all__ = ["ValueHead"] diff --git a/unirl/types/advantages.py b/unirl/types/advantages.py index 7df4c68c4..37e0a500b 100644 --- a/unirl/types/advantages.py +++ b/unirl/types/advantages.py @@ -75,6 +75,48 @@ def compute_gae_advantages( return advantages, returns +def scatter_terminal_rewards( + rewards_per_sample: torch.Tensor, + *, + lengths: torch.Tensor, + cu_seqlens: torch.Tensor, +) -> torch.Tensor: + """Place each sample's scalar reward on its last response token in packed layout. + + Args: + rewards_per_sample: Per-trajectory rewards ``[B]``. + lengths: Response token counts per sample ``[B]``. + cu_seqlens: Packed cumulative offsets ``[B + 1]`` (``TextSegment.cu_seqlens``). + + Returns: + Packed per-token rewards ``[total_tokens]`` (zero except terminal positions). + """ + if rewards_per_sample.ndim != 1: + raise ValueError( + f"scatter_terminal_rewards: rewards_per_sample must be 1D, got shape {tuple(rewards_per_sample.shape)}" + ) + batch_size = int(lengths.shape[0]) + if int(rewards_per_sample.shape[0]) != batch_size: + raise ValueError( + f"scatter_terminal_rewards: rewards batch ({int(rewards_per_sample.shape[0])}) " + f"!= lengths batch ({batch_size})" + ) + if int(cu_seqlens.shape[0]) != batch_size + 1: + raise ValueError( + f"scatter_terminal_rewards: cu_seqlens length ({int(cu_seqlens.shape[0])}) " + f"!= batch_size + 1 ({batch_size + 1})" + ) + total = int(cu_seqlens[-1].item()) + out = rewards_per_sample.new_zeros(total) + cu = [int(c) for c in cu_seqlens.tolist()] + for b in range(batch_size): + n = int(lengths[b].item()) + if n <= 0: + continue + out[cu[b] + n - 1] = rewards_per_sample[b] + return out + + def _gae_1d( rewards: torch.Tensor, values: torch.Tensor, diff --git a/unirl/types/rollout_resp.py b/unirl/types/rollout_resp.py index 18ff7b65d..609805d14 100644 --- a/unirl/types/rollout_resp.py +++ b/unirl/types/rollout_resp.py @@ -43,8 +43,9 @@ from __future__ import annotations +import copy import logging -from dataclasses import dataclass +from dataclasses import dataclass, replace from dataclasses import fields as dc_fields from typing import Any, Callable, Dict, Iterable, List, Literal, Optional, Tuple, Type, TypeVar, Union @@ -59,10 +60,12 @@ shared_field, ) from unirl.distributed.tensor.ref import hydrate +from unirl.types.advantages import compute_gae_advantages as _compute_gae +from unirl.types.advantages import scatter_terminal_rewards from unirl.types.conditions import Condition from unirl.types.media_preview import MediaPreview from unirl.types.primitives import Audios, Images, Texts, Videos -from unirl.types.segments import Segment +from unirl.types.segments import Segment, TextSegment from unirl.utils.shard_balance import lpt_shard_permutation, shard_token_spread logger = logging.getLogger(__name__) @@ -132,8 +135,6 @@ def metadata_only(self) -> "RolloutTrack": (``sample_ids``, ``parent_ids``, ``parent_track``), rewards, advantages, and status. """ - import copy - light = copy.copy(self) light.conditions = {} light.segment = None @@ -417,6 +418,99 @@ def compute_advantages( adv = reshaped - mean return _track_with_field(self, "advantages", adv.flatten()) + def compute_gae_advantages( + self, + *, + gamma: float = 1.0, + gae_lambda: float = 0.95, + use_loss_mask: bool = True, + ) -> "RolloutTrack": + """GAE advantages from per-token ``segment.values`` and scalar ``rewards``. + + Writes packed ``segment.token_advantages`` and ``segment.returns``. + ``track.advantages`` is set to the per-sample mean of token advantages + (for logging / compatibility with existing wandb panels). + + Requires a :class:`~unirl.types.segments.text.TextSegment` with + ``values``, ``lengths``, and ``cu_seqlens`` populated (typically by + ``ARStage.replay(..., return_values=True)``). + + Args: + gamma: Discount factor passed to :func:`compute_gae_advantages`. + gae_lambda: GAE smoothing λ. + use_loss_mask: When ``segment.loss_mask`` is set, pass it as the + GAE validity mask (e.g. response-only tokens). + + Returns: + A new :class:`RolloutTrack` with GAE fields attached. + """ + if self.rewards is None: + raise ValueError("RolloutTrack.compute_gae_advantages: track has no rewards") + if self.segment is None or not isinstance(self.segment, TextSegment): + raise ValueError( + "RolloutTrack.compute_gae_advantages: requires a TextSegment with values" + ) + segment = self.segment + if segment.values is None: + raise ValueError("RolloutTrack.compute_gae_advantages: segment.values is None") + if segment.lengths is None or segment.cu_seqlens is None: + raise ValueError( + "RolloutTrack.compute_gae_advantages: segment requires framework-managed " + "cu_seqlens (construct via TextSegment.pack)" + ) + + values = hydrate(segment.values).to(torch.float32) + lengths = segment.lengths.to(device=values.device) + cu_seqlens = segment.cu_seqlens.to(device=values.device) + rewards_local = hydrate(self.rewards).to(device=values.device, dtype=torch.float32) + token_rewards = scatter_terminal_rewards( + rewards_local, lengths=lengths, cu_seqlens=cu_seqlens + ) + mask = None + if use_loss_mask and segment.loss_mask is not None: + mask = hydrate(segment.loss_mask).to(device=values.device, dtype=values.dtype) + + # Run GAE per packed trajectory so λ-carry and bootstrap reset at each + # sample boundary (a single 1D pass would leak across cu_seqlens gaps). + cu = [int(c) for c in cu_seqlens.tolist()] + token_adv = values.new_zeros(values.shape) + token_returns = values.new_zeros(values.shape) + for b, n in enumerate(lengths.tolist()): + n = int(n) + if n <= 0: + continue + start = cu[b] + sl_rewards = token_rewards[start : start + n] + sl_values = values[start : start + n] + sl_mask = mask[start : start + n] if mask is not None else None + adv, ret = _compute_gae( + sl_rewards, + sl_values, + gamma=gamma, + gae_lambda=gae_lambda, + mask=sl_mask, + ) + token_adv[start : start + n] = adv + token_returns[start : start + n] = ret + + updated_segment = replace( + segment, + token_advantages=token_adv, + returns=token_returns, + ) + # Per-sample mean for existing track-level advantage metrics. + sample_adv: List[torch.Tensor] = [] + for b, n in enumerate(lengths.tolist()): + n = int(n) + if n <= 0: + sample_adv.append(token_adv.new_zeros(())) + continue + sample_adv.append(token_adv[cu[b] : cu[b] + n].mean()) + track_adv = torch.stack(sample_adv) if sample_adv else token_adv.new_zeros((0,)) + + updated = _track_with_field(self, "segment", updated_segment) + return _track_with_field(updated, "advantages", track_adv) + def _root_group_per_sample(resp: "RolloutResp", track_name: str) -> List[str]: """Return the root-track group_id corresponding to each sample of ``track_name``. diff --git a/unirl/types/segments/text.py b/unirl/types/segments/text.py index 45f57b6ef..8e78a54c6 100644 --- a/unirl/types/segments/text.py +++ b/unirl/types/segments/text.py @@ -37,6 +37,10 @@ class TextSegment(Segment): tokens: Optional[torch.Tensor] = packed_field(default=None) log_probs: Optional[torch.Tensor] = packed_field(default=None) loss_mask: Optional[torch.Tensor] = packed_field(default=None) + # PPO / GAE path (optional): per-token critic and advantage plumbing. + values: Optional[torch.Tensor] = packed_field(default=None) + returns: Optional[torch.Tensor] = packed_field(default=None) + token_advantages: Optional[torch.Tensor] = packed_field(default=None) def as_condition_with(self, encoder: Callable[..., Any]) -> Condition: """Re-embed packed tokens via the supplied encoder into a TextEmbedCondition.