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
32 changes: 31 additions & 1 deletion tests/rl/test_producer.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@

import asyncio
import unittest
from unittest.mock import AsyncMock, MagicMock
from unittest.mock import AsyncMock, MagicMock, patch

from xtuner.v1.data_proto.rl_data import RolloutState, Status, discard_rollout_state
from xtuner.v1.rl.agent_loop_manager import (
Expand Down Expand Up @@ -383,6 +383,36 @@ def is_valid_sample_fn(samples):
self.assertEqual(result.failed_samples, 1)
self.assertEqual(result.filtered_samples, 1)

async def test_put_generated_group_releases_non_retryable_trace_sessions(self):
# FAILED / FILTERED 不进入 replay buffer,必须在丢弃 RolloutState 前释放对应 trace session。
cases = (
(Status.FAILED, True, 101),
(Status.COMPLETED, False, 102),
)
for status, is_valid, session_id in cases:
with self.subTest(status=status):
strategy = SyncProduceStrategyConfig(is_valid_sample_fn=lambda _samples: is_valid).build()
ctx = self._build_context(
strategy,
f"non_retryable_{status.name.lower()}",
self._build_agent_loop(),
self._build_sampler(),
batch_size=1,
)
item = make_rollout_state(session_id, status=status, reward_score=1.0)
item.session_id = session_id
item.routed_experts = MagicMock()

with patch(
"xtuner.v1.rl.rollout.trace_store.release_existing_sessions",
new=AsyncMock(return_value={str(session_id)}),
) as release_sessions:
self.assertFalse(await ctx.put_generated_group([item]))

release_sessions.assert_awaited_once_with([str(session_id)])
self.assertIsNone(item.session_id)
self.assertIsNone(item.routed_experts)

async def test_put_generated_group_records_raw_rewards_before_filtering(self):
# 验证 raw reward 在过滤前统计,filtered group 仍能贡献生成侧 reward 指标。
task_name = "test_raw_reward_before_filter"
Expand Down
132 changes: 100 additions & 32 deletions tests/rl/test_replay_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,9 +21,11 @@
# 11. save/resume 保留 Ray ObjectRef:直接 ObjectRef 和 dict(dict(ObjectRef)) 嵌套结构恢复后,
# 解引用得到的内容都应与保存前一致。

import asyncio
import tempfile
import unittest
from pathlib import Path
from unittest.mock import AsyncMock, patch

import numpy as np
import ray
Expand All @@ -41,6 +43,7 @@
def make_rollout_state(
uid: int,
*,
session_id: int | None = None,
status: Status = Status.COMPLETED,
seq_staleness: int = 0,
prompt_ids: list[int] | None = None,
Expand All @@ -66,6 +69,7 @@ def make_rollout_state(
group_id=uid,
message=[{"role": "user", "content": f"prompt {uid}"}],
prompt_ids=prompt_ids,
session_id=session_id,
tokens=list(tokens) if tokens is not None else list(prompt_ids),
response=response if response is not None else f"response {uid}",
response_ids=response_ids,
Expand Down Expand Up @@ -193,6 +197,7 @@ async def test_common_put_drops_expired_group_when_tail_batch_is_disabled(self):
replay_buffer = replay_buffer_config_cls().build()
stale = make_rollout_state(
1,
session_id=101,
prompt_ids=[101, 102],
tokens=[999],
response="stale response",
Expand All @@ -205,14 +210,19 @@ async def test_common_put_drops_expired_group_when_tail_batch_is_disabled(self):
extra_fields={"train_prompt_ids": [101, 102]},
)

await replay_buffer.put(
[stale],
"task",
current_train_step=5,
stale_threshold=3,
expired_groups_retryable=False,
)
with patch(
"xtuner.v1.rl.rollout.trace_store.release_existing_sessions",
new=AsyncMock(return_value={"101"}),
) as release_sessions:
await replay_buffer.put(
[stale],
"task",
current_train_step=5,
stale_threshold=3,
expired_groups_retryable=False,
)

release_sessions.assert_awaited_once_with(["101"])
assert stale.status == Status.EXPIRED
assert stale.prompt_ids is None
assert stale.tokens is None
Expand All @@ -230,6 +240,7 @@ async def test_common_put_defaults_to_retryable_expired_group(self):
pixel_values = np.ones((2, 3), dtype=np.float32)
stale = make_rollout_state(
1,
session_id=102,
prompt_ids=[101, 102],
tokens=[999],
response="stale response",
Expand All @@ -243,13 +254,18 @@ async def test_common_put_defaults_to_retryable_expired_group(self):
extra_fields={"train_prompt_ids": [101, 102]},
)

await replay_buffer.put(
[stale],
"task",
current_train_step=5,
stale_threshold=3,
)
with patch(
"xtuner.v1.rl.rollout.trace_store.release_existing_sessions",
new=AsyncMock(),
) as release_sessions:
await replay_buffer.put(
[stale],
"task",
current_train_step=5,
stale_threshold=3,
)

release_sessions.assert_not_awaited()
expired = await replay_buffer.get(1, "task", Status.EXPIRED)
reusable = expired[0][0]
assert reusable.status == Status.EXPIRED
Expand Down Expand Up @@ -412,46 +428,98 @@ async def test_common_refresh_token_expiry_moves_mixed_group_to_expired_pool(sel
self.assertEqual(group[1].response_ids, [12])
self.assertEqual(group[1].reward, {"score": 0.9})

async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self):
# 同一轮 refresh 仍统计两类过期;只删除 terminal EXPIRED,保留 tail batch 可重试项。
async def test_common_refresh_staleness_drops_only_non_retryable_expired_groups(self):
# 同一轮 refresh 仍统计两类过期;只删除 non-retryable EXPIRED,保留 tail batch 可重试项。
for config_name, replay_buffer_config_cls in REPLAY_BUFFER_CONFIGS:
with self.subTest(replay_buffer_config=config_name):
replay_buffer = replay_buffer_config_cls().build()
terminal_stale = make_rollout_state(
non_retryable_stale = make_rollout_state(
1,
session_id=201,
response_model_steps=[1],
mm_info={"pixel_values": np.ones((2, 3), dtype=np.float32)},
)
retryable_stale = make_rollout_state(
2,
session_id=202,
response_model_steps=[1],
mm_info={"pixel_values": np.ones((2, 3), dtype=np.float32)},
)
await replay_buffer.put([terminal_stale], "terminal_task")
await replay_buffer.put([non_retryable_stale], "non_retryable_task")
await replay_buffer.put([retryable_stale], "retryable_task")
assert len(replay_buffer) == 2

expired_counts = await replay_buffer.refresh_staleness(
task_stale_thresholds={"terminal_task": 2, "retryable_task": 2},
expired_groups_retryable_by_task={
"terminal_task": False,
"retryable_task": True,
},
current_train_step=4,
)
with patch(
"xtuner.v1.rl.rollout.trace_store.release_existing_sessions",
new=AsyncMock(return_value={"201"}),
) as release_sessions:
expired_counts = await replay_buffer.refresh_staleness(
task_stale_thresholds={"non_retryable_task": 2, "retryable_task": 2},
expired_groups_retryable_by_task={
"non_retryable_task": False,
"retryable_task": True,
},
current_train_step=4,
)

assert expired_counts == {"terminal_task": 1, "retryable_task": 1}
assert terminal_stale.status == Status.EXPIRED
assert terminal_stale.prompt_ids is None
assert terminal_stale.mm_info is None
release_sessions.assert_awaited_once_with(["201"])
assert expired_counts == {"non_retryable_task": 1, "retryable_task": 1}
assert non_retryable_stale.status == Status.EXPIRED
assert non_retryable_stale.prompt_ids is None
assert non_retryable_stale.mm_info is None
assert retryable_stale.status == Status.EXPIRED
assert retryable_stale.prompt_ids == [2, 1002]
assert retryable_stale.mm_info is not None
assert await replay_buffer.count("terminal_task", Status.COMPLETED) == 0
assert await replay_buffer.count("terminal_task", Status.EXPIRED) == 0
assert await replay_buffer.count("non_retryable_task", Status.COMPLETED) == 0
assert await replay_buffer.count("non_retryable_task", Status.EXPIRED) == 0
assert await replay_buffer.count("retryable_task", Status.EXPIRED) == 1
assert len(replay_buffer) == 1
assert await replay_buffer.get(1, "terminal_task", Status.EXPIRED) == []
assert await replay_buffer.get(1, "non_retryable_task", Status.EXPIRED) == []

async def test_refresh_staleness_batches_non_retryable_release_outside_lock(self):
replay_buffer = AsyncReplayBufferConfig().build()
first = make_rollout_state(1, session_id=301, response_model_steps=[1])
second = make_rollout_state(2, session_id=302, response_model_steps=[1])
await replay_buffer.put([first], "non_retryable_task")
await replay_buffer.put([second], "non_retryable_task")

release_started = asyncio.Event()
allow_release = asyncio.Event()

async def delayed_release(session_ids):
release_started.set()
await allow_release.wait()
return set(session_ids)

with patch(
"xtuner.v1.rl.rollout.trace_store.release_existing_sessions",
new=AsyncMock(side_effect=delayed_release),
) as release_sessions:
refresh_task = asyncio.create_task(
replay_buffer.refresh_staleness(
task_stale_thresholds={"non_retryable_task": 2},
expired_groups_retryable_by_task={"non_retryable_task": False},
current_train_step=4,
)
)
await asyncio.wait_for(release_started.wait(), timeout=1.0)
try:
# The non-retryable records are already removed and the buffer lock is
# available while the trace-store RPC is still blocked.
count = await asyncio.wait_for(
replay_buffer.count("non_retryable_task", Status.COMPLETED),
timeout=1.0,
)
assert count == 0
finally:
allow_release.set()
expired_counts = await refresh_task

release_sessions.assert_awaited_once_with(["301", "302"])
assert expired_counts == {"non_retryable_task": 2}
assert first.status == Status.EXPIRED
assert second.status == Status.EXPIRED
assert len(replay_buffer) == 0

async def test_common_refresh_staleness_contract(self):
# refresh_staleness 同时覆盖默认刷新 completed/aborted,以及 status filter 只刷新指定状态。
Expand Down
5 changes: 4 additions & 1 deletion tests/rl/test_rl_colocate_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,7 +142,8 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_
trainer._benchmark_training_samples = 0
trainer._benchmark_training_tokens = 0
trainer._save_trajectories = MagicMock()
trainer._release_trace_store = MagicMock()
trainer._release_trace_sessions = MagicMock(return_value=set())
trainer._release_all_trace_sessions = MagicMock()
trainer._sync_weights_and_save = MagicMock(
side_effect=lambda train_step, step_timer_dict: train_step % trainer._sync_weights_interval == 0
)
Expand Down Expand Up @@ -208,6 +209,8 @@ def test_fit_accepts_async_strategy_manager_on_colocate_path(self):
trainer.rollout_controller.offload.remote.assert_called_once_with()
trainer.train_controller.onload.assert_called_once_with(target="all")
trainer.train_controller.fit.assert_called_once()
trainer._release_all_trace_sessions.assert_called_once_with()
trainer._release_trace_sessions.assert_not_called()
self.assertEqual(trainer._cur_step, 1)

def test_fit_requires_non_empty_batch_from_manager(self):
Expand Down
Loading
Loading