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: 3 additions & 2 deletions recipe/verl_agent/common/agent_loop_verl_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,9 @@
from verl.workers.rollout.replica import TokenOutput

from xtuner.v1.data_proto.rl_data import RolloutState, SampleParams, Status
from xtuner.v1.rl.agent_loop import AgentLoop, AgentLoopConfig
from xtuner.v1.rl.judger import Judger
from xtuner.v1.rl.rollout.controller import RolloutControllerProxy
from xtuner.v1.rl.agent_loop import AgentLoop, AgentLoopConfig


class VerlToolAgentLoopConfig(AgentLoopConfig):
Expand Down Expand Up @@ -138,6 +138,7 @@ async def generate_sample(self, rollout_state: RolloutState) -> RolloutState:
rollout_state.response_ids = output.response_ids
rollout_state.logprobs = output.response_logprobs
rollout_state.routed_experts = output.routed_experts
rollout_state.routed_experts_owner = "rollout" if output.routed_experts is not None else None
rollout_state.response_mask = output.response_mask
rollout_state.status = Status.COMPLETED
rollout_state.extra_fields.update(output.extra_fields)
Expand All @@ -150,5 +151,5 @@ async def generate_sample(self, rollout_state: RolloutState) -> RolloutState:
# judge rollout_state
if self.judger is not None:
rollout_state = await self.judger.judge(rollout_state)

return rollout_state
1 change: 1 addition & 0 deletions tests/rl/test_producer.py
Original file line number Diff line number Diff line change
Expand Up @@ -322,6 +322,7 @@ async def test_discard_rollout_state_keeps_required_fields_valid(self):
# 验证 discard 不破坏 RolloutState 的必填字段契约,同时释放可丢弃的重字段。
item = make_rollout_state(42, status=Status.COMPLETED, reward_score=1.0)
item.routed_experts = MagicMock()
item.routed_experts_owner = "rollout"
item.extra_fields = {"large": [1, 2, 3]}

discarded = discard_rollout_state(item)
Expand Down
83 changes: 83 additions & 0 deletions tests/rl/test_replay_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,89 @@ async def save_and_resume(


class TestReplayBuffer(unittest.IsolatedAsyncioTestCase):
async def test_retryable_stale_rollout_refs_are_released(self):
"""Retryable expiry releases direct-rollout refs before resetting
state."""
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()
routed_refs = [object()]
stale = make_rollout_state(
1,
response_model_steps=[0],
routed_experts=routed_refs,
)
stale.routed_experts_owner = "rollout"

with patch("xtuner.v1.rl.utils.ray_utils.free_object_refs") as free_refs:
status = replay_buffer._apply_staleness_lifecycle(
[stale],
current_train_step=5,
stale_threshold=3,
token_stale_threshold=None,
expired_groups_retryable=True,
)

self.assertEqual(status, Status.EXPIRED)
free_refs.assert_called_once_with(routed_refs)
self.assertIsNone(stale.routed_experts)
self.assertIsNone(stale.routed_experts_owner)

async def test_retryable_stale_trace_store_refs_are_only_detached(self):
"""A single stale segment must not release a shared TraceStore
session."""
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()
stale = make_rollout_state(
1,
response_model_steps=[0],
routed_experts=[object()],
)
stale.routed_experts_owner = "trace_store"

with patch("xtuner.v1.rl.utils.ray_utils.free_object_refs") as free_refs:
status = replay_buffer._apply_staleness_lifecycle(
[stale],
current_train_step=5,
stale_threshold=3,
token_stale_threshold=None,
expired_groups_retryable=True,
)

self.assertEqual(status, Status.EXPIRED)
free_refs.assert_not_called()
self.assertIsNone(stale.routed_experts)
self.assertIsNone(stale.routed_experts_owner)

async def test_retryable_stale_unowned_refs_are_released(self):
"""Refs without an owner tag (legacy-checkpoint restores) must not leak
when the state is reset."""
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()
routed_refs = [object()]
stale = make_rollout_state(
1,
response_model_steps=[0],
routed_experts=routed_refs,
)
stale.routed_experts_owner = None

with patch("xtuner.v1.rl.utils.ray_utils.free_object_refs") as free_refs:
status = replay_buffer._apply_staleness_lifecycle(
[stale],
current_train_step=5,
stale_threshold=3,
token_stale_threshold=None,
expired_groups_retryable=True,
)

self.assertEqual(status, Status.EXPIRED)
free_refs.assert_called_once_with(routed_refs)
self.assertIsNone(stale.routed_experts)
self.assertIsNone(stale.routed_experts_owner)

async def test_common_query_count_and_take_batch_contract(self):
# ReplayBuffer 的公共读写契约:按 task/status 隔离统计,并且 take_batch 会消费已取出的数据。
for config_name, replay_buffer_config_cls in REPLAY_BUFFER_CONFIGS:
Expand Down
120 changes: 109 additions & 11 deletions tests/rl/test_rollout_logic.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@
- SGLangWorker pause/continue 对 abort flag 和 server request 的控制。
- RolloutWorker abort、abort request timeout 和 in-flight request 取消语义。
- RolloutHealthManager 对 inactive/unhealthy worker 的生命周期标记逻辑。
- PartialRolloutHandler 拼接 routed_experts 后释放旧 Ray ObjectRef 的逻辑。
- PartialRolloutHandler 拼接 routed_experts 后由调用方显式释放旧 Ray ObjectRef 的逻辑。

旧 test_rollout_utils.py 中的 TestRolloutControllerRecover 需要真实 Ray controller / lmdeploy backend,
不属于 PR-fast,后续应放到 PR-real smoke 或 nightly。
Expand Down Expand Up @@ -1271,12 +1271,8 @@ def test_pending_weight_update_rechecks_rollout_phase_after_lock_acquire(self):
rollout_controller=MagicMock(),
rollout_config=SimpleNamespace(weight_transport_type="checkpoint_engine"),
)
manager._pending_rollout_weight_update_stop_event = SimpleNamespace(
wait=MagicMock(side_effect=(False, True))
)
manager._rollout_resources_available = SimpleNamespace(
is_set=MagicMock(side_effect=(True, False))
)
manager._pending_rollout_weight_update_stop_event = SimpleNamespace(wait=MagicMock(side_effect=(False, True)))
manager._rollout_resources_available = SimpleNamespace(is_set=MagicMock(side_effect=(True, False)))
manager._rollout_weight_update_lock = SimpleNamespace(
acquire=MagicMock(return_value=True),
release=MagicMock(),
Expand Down Expand Up @@ -1660,6 +1656,7 @@ def fake_ray_get(refs, timeout=None):
self.assertEqual(actor.offload.calls, [()])
self.assertEqual(actor.restore_skip_load_weights.calls, [()])


class TestPartialRolloutHandler(unittest.IsolatedAsyncioTestCase):
async def test_preprocess_and_postprocess_preserve_response_prefix(self):
# partial rollout 续写时应复用 prompt+历史 response,并把新 response token 追加到历史后面。
Expand Down Expand Up @@ -1744,8 +1741,9 @@ async def test_multi_round_partial_rollout_never_exceeds_max_tokens(self):
self.assertEqual(rollout_state.response_ids, [101, 102, 201, 202, 301])
self.assertLessEqual(len(rollout_state.response_ids), max_tokens)

async def test_postprocess_frees_old_routed_expert_refs_after_concat(self):
# partial rollout 拼接 routed_experts 后,应释放历史和当前 ObjectRef,避免长期占用对象存储。
async def test_postprocess_frees_input_refs_owned_by_rollout_after_concat(self):
# concat 之后 history/current 两个输入 ref 都被 direct rollout 路径释放。

class FakeObjectRef:
def __init__(self, value):
self.value = value
Expand All @@ -1765,6 +1763,7 @@ async def _resolve():
response_ids=[1, 2],
logprobs=[0.1, 0.2],
routed_experts=history_ref,
routed_experts_owner="rollout",
status=Status.ABORTED,
)

Expand All @@ -1787,6 +1786,105 @@ async def _resolve():

self.assertIs(out.routed_experts, concat_ref)
self.assertEqual(ray_put.call_args.args[0].tolist(), [[1], [2], [3]])
free_object_refs.assert_any_call([history_ref])
free_object_refs.assert_any_call([cur_ref])
free_object_refs.assert_any_call(cur_ref)
free_object_refs.assert_any_call(history_ref)
self.assertEqual(free_object_refs.call_count, 2)

async def test_postprocess_does_not_free_trace_store_history_ref(self):
"""History refs borrowed from the TraceStore must stay valid for
sibling segments; only the current ref is released."""

class FakeObjectRef:
def __init__(self, value):
self.value = value

def __await__(self):
async def _resolve():
return self.value

return _resolve().__await__()

history_ref = FakeObjectRef([[1], [2]])
cur_ref = FakeObjectRef([[1], [2], [3]])
concat_ref = FakeObjectRef(None)
rollout_state = RolloutState(
message=[],
response="old",
response_ids=[1, 2],
logprobs=[0.1, 0.2],
routed_experts=history_ref,
routed_experts_owner="trace_store",
status=Status.ABORTED,
)

with (
patch("xtuner.v1.rl.rollout.utils.RayObjectRef", FakeObjectRef),
patch("xtuner.v1.rl.rollout.utils.ray.put", return_value=concat_ref),
patch("xtuner.v1.rl.rollout.utils.free_object_refs") as free_object_refs,
):
out = await PartialRolloutHandler().postprocess(
rollout_state,
response="new",
response_ids=[3],
logprobs=[0.3],
routed_experts=cur_ref,
finish_reason="abort",
status=Status.ABORTED,
prompt_tokens=3,
completion_tokens=1,
)

self.assertIs(out.routed_experts, concat_ref)
free_object_refs.assert_called_once_with(cur_ref)
self.assertEqual(out.routed_experts_owner, "rollout")

async def test_postprocess_frees_unowned_history_ref(self):
"""History refs without an owner tag (legacy-checkpoint restores)
follow the same rule as ``rl_data``: anything not borrowed from the
TraceStore is released once the concatenation replaces it."""

class FakeObjectRef:
def __init__(self, value):
self.value = value

def __await__(self):
async def _resolve():
return self.value

return _resolve().__await__()

history_ref = FakeObjectRef([[1], [2]])
cur_ref = FakeObjectRef([[1], [2], [3]])
concat_ref = FakeObjectRef(None)
rollout_state = RolloutState(
message=[],
response="old",
response_ids=[1, 2],
logprobs=[0.1, 0.2],
routed_experts=history_ref,
routed_experts_owner=None,
status=Status.ABORTED,
)

with (
patch("xtuner.v1.rl.rollout.utils.RayObjectRef", FakeObjectRef),
patch("xtuner.v1.rl.rollout.utils.ray.put", return_value=concat_ref),
patch("xtuner.v1.rl.rollout.utils.free_object_refs") as free_object_refs,
):
out = await PartialRolloutHandler().postprocess(
rollout_state,
response="new",
response_ids=[3],
logprobs=[0.3],
routed_experts=cur_ref,
finish_reason="abort",
status=Status.ABORTED,
prompt_tokens=3,
completion_tokens=1,
)

self.assertIs(out.routed_experts, concat_ref)
free_object_refs.assert_any_call(cur_ref)
free_object_refs.assert_any_call(history_ref)
self.assertEqual(free_object_refs.call_count, 2)
self.assertEqual(out.routed_experts_owner, "rollout")
67 changes: 67 additions & 0 deletions tests/rl/test_staleness_policy.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,12 +4,15 @@
"""

import unittest
from unittest.mock import patch

import ray
from pydantic import ValidationError

from xtuner.v1.data_proto.rl_data import (
RolloutState,
calculate_group_effective_response_masks,
discard_rollout_state,
reset_rollout_response,
)
from xtuner.v1.rl.agent_loop_manager import (
Expand Down Expand Up @@ -115,6 +118,70 @@ def test_rerolled_state_without_semantic_mask_uses_token_staleness_only(self):
self.assertIsNone(state.response_mask)
self.assertEqual(masks, [[1, 1]])

def test_reset_rollout_response_releases_owned_and_detaches_borrowed_refs(self):
"""Reset releases rollout-owned refs and only detaches TraceStore
borrows."""
from xtuner.v1.data_proto.rl_data import _release_owned_routed_experts

class FakeObjectRef:
pass

direct = self._state(response_model_steps=[0])
direct_ref = FakeObjectRef()
direct.routed_experts = direct_ref
direct.routed_experts_owner = "rollout"
borrowed = self._state(response_model_steps=[0])
borrowed.routed_experts = FakeObjectRef()
borrowed.routed_experts_owner = "trace_store"

with (
patch.object(ray, "ObjectRef", FakeObjectRef),
patch("xtuner.v1.rl.utils.ray_utils.free_object_refs") as free_refs,
):
reset_rollout_response(direct)
_release_owned_routed_experts(borrowed)
reset_rollout_response(borrowed)

free_refs.assert_called_once_with(direct_ref)
self.assertIsNone(direct.routed_experts)
self.assertIsNone(direct.routed_experts_owner)
self.assertIsNone(borrowed.routed_experts)
self.assertIsNone(borrowed.routed_experts_owner)

def test_reset_rollout_response_releases_unowned_refs(self):
"""Refs without an owner tag (legacy-checkpoint restores) must not leak
when the state is reset."""

class FakeObjectRef:
pass

state = self._state(response_model_steps=[0])
unowned_ref = FakeObjectRef()
state.routed_experts = unowned_ref
state.routed_experts_owner = None

with (
patch.object(ray, "ObjectRef", FakeObjectRef),
patch("xtuner.v1.rl.utils.ray_utils.free_object_refs") as free_refs,
):
reset_rollout_response(state)

free_refs.assert_called_once_with(unowned_ref)
self.assertIsNone(state.routed_experts)
self.assertIsNone(state.routed_experts_owner)

def test_discard_trace_store_state_detaches_without_freeing_trace_ref(self):
state = self._state(response_model_steps=[0])
state.routed_experts = object()
state.routed_experts_owner = "trace_store"

with patch("xtuner.v1.rl.utils.ray_utils.free_object_refs") as free_refs:
discarded = discard_rollout_state(state)

free_refs.assert_not_called()
self.assertIsNone(discarded.routed_experts)
self.assertIsNone(discarded.routed_experts_owner)

@staticmethod
def _state(
*,
Expand Down
Loading
Loading