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
3 changes: 2 additions & 1 deletion 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 Down
3 changes: 2 additions & 1 deletion tests/rl/test_producer.py
Original file line number Diff line number Diff line change
Expand Up @@ -322,9 +322,10 @@ 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)
discarded = discard_rollout_state(item, release_refs=True)

self.assertEqual(discarded.message, [{"role": "user", "content": "prompt 42"}])
self.assertEqual(discarded.status, Status.INIT)
Expand Down
54 changes: 54 additions & 0 deletions tests/rl/test_replay_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,60 @@ async def save_and_resume(


class TestReplayBuffer(unittest.IsolatedAsyncioTestCase):
async def test_retryable_stale_rollout_refs_are_released_by_caller(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()
stale = make_rollout_state(
1,
response_model_steps=[0],
routed_experts=object(),
)
stale.routed_experts_owner = "rollout"

with patch("xtuner.v1.rl.replay_buffer.release_owned_routed_experts") as release_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)
release_refs.assert_called_once_with(stale)
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.replay_buffer.release_owned_routed_experts") as release_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)
release_refs.assert_not_called()
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
54 changes: 52 additions & 2 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 @@ -1725,7 +1725,7 @@ async def test_multi_round_partial_rollout_never_exceeds_max_tokens(self):
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,避免长期占用对象存储
# direct rollout caller 显式要求 partial handler 释放历史和当前 ObjectRef。
class FakeObjectRef:
def __init__(self, value):
self.value = value
Expand All @@ -1745,6 +1745,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 @@ -1763,10 +1764,59 @@ async def _resolve():
status=Status.ABORTED,
prompt_tokens=3,
completion_tokens=1,
release_input_routed_experts=True,
)

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])
self.assertEqual(free_object_refs.call_count, 2)

async def test_postprocess_does_not_free_input_refs_by_default(self):
"""PartialRolloutHandler must leave input refs to its caller by
default."""

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_not_called()
self.assertEqual(out.routed_experts_owner, "rollout")
56 changes: 56 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,59 @@ 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_only_clears_fields(self):
"""Reset must not implicitly release an ObjectRef owned by its
caller."""
state = self._state(response_model_steps=[0, 4])
state.routed_experts = object()
state.routed_experts_owner = "rollout"

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

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

def test_release_owned_routed_experts_only_frees_direct_rollout_refs(self):
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,
):
release_owned_routed_experts(direct)
release_owned_routed_experts(borrowed)

free_refs.assert_called_once_with(direct_ref)
self.assertIsNone(direct.routed_experts)
self.assertIsNone(direct.routed_experts_owner)
self.assertIsNotNone(borrowed.routed_experts)
self.assertEqual(borrowed.routed_experts_owner, "trace_store")

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, release_refs=True)

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

@staticmethod
def _state(
*,
Expand Down
40 changes: 37 additions & 3 deletions tests/rl/test_trace_store.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,12 +19,18 @@ class TestRolloutTraceCleanup(unittest.TestCase):
def test_release_and_discard_detaches_only_trace_owned_refs(self):
trace_owned_ref = object()
rollout_owned_ref = object()
trace_owned = SimpleNamespace(session_id="trace-owned", routed_experts=trace_owned_ref)
rollout_owned = SimpleNamespace(session_id="rollout-owned", routed_experts=rollout_owned_ref)
trace_owned = SimpleNamespace(
session_id="trace-owned", routed_experts=trace_owned_ref, routed_experts_owner="trace_store"
)
rollout_owned = SimpleNamespace(
session_id="rollout-owned", routed_experts=rollout_owned_ref, routed_experts_owner="rollout"
)
routed_experts_seen_by_discard = {}
release_flags = {}

def record_discard(item):
def record_discard(item, **kwargs):
routed_experts_seen_by_discard[item.session_id] = item.routed_experts
release_flags[item.session_id] = kwargs["release_refs"]

with (
patch(
Expand All @@ -41,6 +47,8 @@ def record_discard(item):
release_sessions.assert_awaited_once_with(["trace-owned", "rollout-owned"])
self.assertIsNone(routed_experts_seen_by_discard["trace-owned"])
self.assertIs(routed_experts_seen_by_discard["rollout-owned"], rollout_owned_ref)
self.assertTrue(release_flags["trace-owned"])
self.assertTrue(release_flags["rollout-owned"])
self.assertEqual(discard.call_count, 2)

def test_get_existing_store_returns_none_when_ray_is_uninitialized(self):
Expand Down Expand Up @@ -82,6 +90,32 @@ def test_release_sessions_deduplicates_and_skips_missing_ids(self):
finally:
ray.kill(store)

def test_trie_overwrite_releases_replaced_ref_and_keeps_new_value(self):
trie = trace_store_module.Trie()
old_ref = ray.put({"value": "old"})
new_ref = ray.put({"value": "new"})
trie.insert("turn", {"expert_key": old_ref})

with patch.object(ray.internal, "free") as free:
trie.insert("turn", {"expert_key": new_ref})

free.assert_called_once_with([old_ref], local_only=False)
_, nodes = trie.search("turn", filter_none=True)
self.assertEqual(ray.get(nodes[-1].value["expert_key"]), {"value": "new"})

def test_trie_overwrite_does_not_free_ref_still_shared_by_session(self):
trie = trace_store_module.Trie()
shared_ref = ray.put({"value": "shared"})
replacement_ref = ray.put({"value": "replacement"})
trie.insert("first", {"expert_key": shared_ref})
trie.insert("second", {"expert_key": shared_ref})

with patch.object(ray.internal, "free") as free:
trie.insert("first", {"expert_key": replacement_ref})

free.assert_not_called()
self.assertEqual(ray.get(shared_ref), {"value": "shared"})

def test_release_existing_sessions_stably_deduplicates_before_rpc(self):
release_remote = AsyncMock(return_value=["one"])
store = SimpleNamespace(release_sessions=SimpleNamespace(remote=release_remote))
Expand Down
Loading
Loading