Skip to content
Closed
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
15 changes: 15 additions & 0 deletions tests/models/test_value_head.py
Original file line number Diff line number Diff line change
@@ -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
10 changes: 9 additions & 1 deletion tests/types/test_advantages_gae.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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]
88 changes: 88 additions & 0 deletions tests/types/test_rollout_track_gae.py
Original file line number Diff line number Diff line change
@@ -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)
155 changes: 133 additions & 22 deletions unirl/models/qwen3/ar.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,14 +23,15 @@
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
import torch.nn.functional as F
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

Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -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]
Expand All @@ -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, :]

Expand All @@ -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):
Expand Down Expand Up @@ -418,30 +488,37 @@ 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,
conditions: Qwen3ARConditions,
*,
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
Expand Down Expand Up @@ -517,17 +594,25 @@ 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,
conditions: Qwen3ARConditions,
*,
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
Expand Down Expand Up @@ -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,
Expand Down
Loading