diff --git a/recipe/verl_agent/common/agent_loop_verl_tool.py b/recipe/verl_agent/common/agent_loop_verl_tool.py index f78ef27264..fab2b49090 100644 --- a/recipe/verl_agent/common/agent_loop_verl_tool.py +++ b/recipe/verl_agent/common/agent_loop_verl_tool.py @@ -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): @@ -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) @@ -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 diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 364db9f09c..edf3b23ecf 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -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) diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index 737ca09d4a..b91d7ff85b 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -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: diff --git a/tests/rl/test_rollout_logic.py b/tests/rl/test_rollout_logic.py index de0af7120d..baa0de1bbd 100644 --- a/tests/rl/test_rollout_logic.py +++ b/tests/rl/test_rollout_logic.py @@ -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。 @@ -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(), @@ -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 追加到历史后面。 @@ -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 @@ -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, ) @@ -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") diff --git a/tests/rl/test_staleness_policy.py b/tests/rl/test_staleness_policy.py index b9a746d1c8..f6d4db405e 100644 --- a/tests/rl/test_staleness_policy.py +++ b/tests/rl/test_staleness_policy.py @@ -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 ( @@ -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( *, diff --git a/tests/rl/test_trace_store.py b/tests/rl/test_trace_store.py index 51ad063362..03d8286435 100644 --- a/tests/rl/test_trace_store.py +++ b/tests/rl/test_trace_store.py @@ -19,8 +19,12 @@ 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 = {} def record_discard(item): @@ -82,6 +86,48 @@ def test_release_sessions_deduplicates_and_skips_missing_ids(self): finally: ray.kill(store) + def test_trie_overwrite_drops_handle_without_freeing_borrowed_refs(self): + """Overwrite must not free replaced refs: sibling RolloutStates keep + their own handles alive via Ray's reference counting.""" + 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_not_called() + # The borrower keeps its own handle: dropping the trie's handle must + # not invalidate the object. + self.assertEqual(ray.get(old_ref), {"value": "old"}) + _, nodes = trie.search("turn", filter_none=True) + self.assertEqual(ray.get(nodes[-1].value["expert_key"]), {"value": "new"}) + + def test_trie_release_frees_each_ref_once_across_shared_subtrees(self): + """Refs shared between trie values are freed exactly once at session + release.""" + trie = trace_store_module.Trie() + left_ref = ray.put({"value": "left"}) + right_ref = ray.put({"value": "right"}) + shared_ref = ray.put({"value": "shared"}) + trie.insert("turn-a", {"expert_key": left_ref, "keep": shared_ref}) + trie.insert("turn-b", {"expert_key": right_ref, "keep": shared_ref}) + + freed = [] + + def record_free(refs, **kwargs): + freed.extend(refs) + + with patch.object(ray.internal, "free", side_effect=record_free): + trie.release() + + self.assertEqual(len(freed), len({ref.hex() for ref in freed})) + self.assertEqual( + {ref.hex() for ref in freed}, + {left_ref.hex(), right_ref.hex(), shared_ref.hex()}, + ) + def test_release_existing_sessions_stably_deduplicates_before_rpc(self): release_remote = AsyncMock(return_value=["one"]) store = SimpleNamespace(release_sessions=SimpleNamespace(remote=release_remote)) diff --git a/xtuner/v1/data_proto/rl_data.py b/xtuner/v1/data_proto/rl_data.py index b2e5ef43d1..798a445de5 100644 --- a/xtuner/v1/data_proto/rl_data.py +++ b/xtuner/v1/data_proto/rl_data.py @@ -58,6 +58,9 @@ class Status(Enum): ARCHIVED = "archived" +RoutedExpertsOwner: TypeAlias = Literal["rollout", "trace_store"] + + class MultimodalInfo(TypedDict): # 使用TypedDict给出pixel_values的类型提示 pixel_values: NotRequired[np.ndarray | RayObjectRef | None] @@ -139,6 +142,10 @@ class RolloutState(BaseModel): logprobs: list[float] | None = None teacher_targets: TeacherTargets | None = None routed_experts: np.ndarray | RayObjectRef | list[RayObjectRef] | None = None + # ``routed_experts`` may either be produced by a rollout worker or borrowed + # from the shared TraceStore. The distinction is deliberately explicit: + # an ObjectRef's container type does not convey ownership. + routed_experts_owner: RoutedExpertsOwner | None = None finish_reason: str | None = None # response_mask: 记录response_ids中哪个token算loss, 与response_ids长度相同,每轮rollout在 agent_loop.generate 中覆盖写 response_mask: list[int] | None = None @@ -205,12 +212,20 @@ def clear_object_refs(value: Any) -> Any: return value clear_object_refs(rollout_state) - free_object_refs(refs) + if refs: + free_object_refs(refs) def discard_rollout_state(rollout_state: RolloutState) -> RolloutState: - """Release heavy references and clear fields before dropping a rollout.""" + """Release owned references and clear fields before dropping a rollout. + Routed experts borrowed from the TraceStore are only detached: the + TraceStore session owns them and releases them through its own lifecycle. + """ + + if rollout_state.routed_experts_owner == "trace_store": + rollout_state.routed_experts = None + rollout_state.routed_experts_owner = None free_rollout_state_refs(rollout_state) for field_name, field in type(rollout_state).model_fields.items(): @@ -258,16 +273,38 @@ def update_status_from_finish_reason(finish_reason: str | None) -> Status: return Status.FAILED -def reset_rollout_response(rollout_state: RolloutState) -> RolloutState: - routed_experts = getattr(rollout_state, "routed_experts", None) - if routed_experts is not None: - from ray import ObjectRef +def _release_owned_routed_experts(rollout_state: RolloutState) -> None: + """Release routed-expert refs owned by this rollout state. + + ``RolloutState`` can also contain refs borrowed from ``TraceStore``. This + helper refuses to free those refs; the TraceStore session is their owner + and must release them through its own lifecycle. Refs with no owner tag + (e.g. restored from a legacy checkpoint via ``ray.put``) have no other + borrower and are safe to release here. + """ - from xtuner.v1.rl.utils.ray_utils import free_object_refs + routed_experts = rollout_state.routed_experts + if routed_experts is None: + return - if isinstance(routed_experts, (ObjectRef, list)): - free_object_refs(routed_experts) - rollout_state.routed_experts = None + if rollout_state.routed_experts_owner == "trace_store": + # Expected on retryable stale segments: the borrowed ref stays valid + # until the TraceStore session is released. + logger.debug( + f"Detaching TraceStore-owned routed_experts without freeing (session_id={rollout_state.session_id!r})." + ) + return + + from ray import ObjectRef + + from xtuner.v1.rl.utils.ray_utils import free_object_refs + + if isinstance(routed_experts, (ObjectRef, list)): + free_object_refs(routed_experts) + + +def reset_rollout_response(rollout_state: RolloutState) -> RolloutState: + _release_owned_routed_experts(rollout_state) prompt_ids = getattr(rollout_state, "prompt_ids", None) rollout_state.tokens = list(prompt_ids) if prompt_ids is not None else None rollout_state.response = "" @@ -275,6 +312,7 @@ def reset_rollout_response(rollout_state: RolloutState) -> RolloutState: rollout_state.logprobs = [] rollout_state.teacher_targets = None rollout_state.routed_experts = None + rollout_state.routed_experts_owner = None rollout_state.finish_reason = None rollout_state.response_mask = None rollout_state.response_model_steps = [] diff --git a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py index bdba686585..edd6d2efd0 100644 --- a/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py +++ b/xtuner/v1/rl/agent_loop/localhost_agent_loop/agent_in_localhost_loop.py @@ -264,6 +264,7 @@ async def _fill_rollout_state(self, rollout_state: RolloutState, item: AgentRoll ] rollout_state.logprobs = data["logprobs"] rollout_state.routed_experts = data["routed_experts"] + rollout_state.routed_experts_owner = "trace_store" if data["routed_experts"] is not None else None content = response_message.get("content") rollout_state.response = content if isinstance(content, str) else (str(content) if content is not None else "") @@ -278,6 +279,7 @@ def _fill_eval_rollout_state(self, rollout_state: RolloutState, item: AgentRollo rollout_state.response_ids = None rollout_state.logprobs = None rollout_state.routed_experts = None + rollout_state.routed_experts_owner = None rollout_state.response_mask = None rollout_state.response_model_steps = None rollout_state.extra_fields["agent_status"] = item.status.value diff --git a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py index 14dec5bf74..418d36243f 100644 --- a/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py +++ b/xtuner/v1/rl/agent_loop/sandbox_agent_loop/agent_in_sandbox_loop.py @@ -350,6 +350,7 @@ async def _build_rollout_states(self, rollout_state: RolloutState, item: AgentRo # indistinguishable from a dense trace; catching that would need the MoE-vs-dense signal the loop lacks # (session_server gates on ``enable_return_routed_experts``). segment_state.routed_experts = data["routed_experts"] + segment_state.routed_experts_owner = "trace_store" if data["routed_experts"] is not None else None if segment_state.response_ids: segment_state.response = self.tokenizer.decode(segment_state.response_ids) else: @@ -370,6 +371,7 @@ def _fill_eval_rollout_state(self, rollout_state: RolloutState, item: AgentRollo rollout_state.response_ids = None rollout_state.logprobs = None rollout_state.routed_experts = None + rollout_state.routed_experts_owner = None rollout_state.response_mask = None rollout_state.response_model_steps = None rollout_state.extra_fields["origin_data_source"] = item.data_source diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 36ba406e0a..d4c977bbfb 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -14,32 +14,42 @@ _handle_cache: Any = None -def _free_ray_refs(obj: Any): +def _ray_ref_key(ref: ray.ObjectRef) -> str: + """Return a stable key for de-duplicating one Ray object reference.""" + return ref.hex() + + +def _free_ray_refs(obj: Any, *, _seen: set[str] | None = None): """Recursively free ray.ObjectRef instances trapped inside an object. Args: obj (Any): The object that may contain ray.ObjectRef references (e.g., dict, list, tuple). """ + seen = set() if _seen is None else _seen if isinstance(obj, ray.ObjectRef): + ref_key = _ray_ref_key(obj) + if ref_key in seen: + return + seen.add(ref_key) try: ray.internal.free([obj], local_only=False) except Exception as e: get_logger().error(f"Failed to free Ray ObjectRef {obj}: {e}") elif isinstance(obj, dict): for v in obj.values(): - _free_ray_refs(v) + _free_ray_refs(v, _seen=seen) elif isinstance(obj, (list, tuple)): for v in obj: - _free_ray_refs(v) + _free_ray_refs(v, _seen=seen) elif hasattr(obj, "model_dump"): # Pydantic v2 for v in obj.model_dump().values() if hasattr(obj.model_dump, "__call__") else {}.values(): - _free_ray_refs(v) + _free_ray_refs(v, _seen=seen) elif hasattr(obj, "dict") and callable(getattr(obj, "dict")): # Pydantic v1 for v in obj.dict().values(): - _free_ray_refs(v) + _free_ray_refs(v, _seen=seen) elif hasattr(obj, "__dict__"): for v in vars(obj).values(): - _free_ray_refs(v) + _free_ray_refs(v, _seen=seen) def _common_prefix_len(left: str, right: str) -> int: @@ -173,6 +183,10 @@ def insert(self, key: str, value: Any) -> None: node = node.children[key] break + # A rerolled turn may overwrite an existing key. Do not free the + # old value here: borrowers (sibling RolloutStates) keep their own + # refs alive via Ray's reference counting, and the trie dropping + # its handle never invalidates them. node.value = value def search(self, text: str, filter_none: bool = False) -> Tuple[str, List["TreeNode"]]: @@ -220,16 +234,16 @@ def release(self, key: str | None = None): If None, releases the entire tree. """ - def _free_subtree(node: TreeNode): + def _free_subtree(node: TreeNode, seen: set[str]): for child in node.children.values(): - _free_subtree(child) + _free_subtree(child, seen) if node.value is not None: - _free_ray_refs(node.value) + _free_ray_refs(node.value, _seen=seen) node.value = None node.children.clear() if key is None: - _free_subtree(self.root) + _free_subtree(self.root, set()) return node = self.root @@ -523,6 +537,7 @@ async def release_and_discard_rollout_groups(groups: list[list[RolloutState]]) - for item in group: if item.session_id is not None and str(item.session_id) in released_session_ids: item.routed_experts = None + item.routed_experts_owner = None discard_rollout_state(item) diff --git a/xtuner/v1/rl/rollout/utils.py b/xtuner/v1/rl/rollout/utils.py index 94ccc8023a..d63d77bfa0 100644 --- a/xtuner/v1/rl/rollout/utils.py +++ b/xtuner/v1/rl/rollout/utils.py @@ -103,7 +103,6 @@ async def get_worker(self, session_id: int) -> Optional[Any]: async def _resolve_routed_experts(routed_experts: np.ndarray | RayObjectRef) -> np.ndarray: if isinstance(routed_experts, RayObjectRef): routed_experts_value = await routed_experts - free_object_refs([routed_experts]) else: routed_experts_value = routed_experts assert routed_experts_value is not None, "routed_experts should not be empty after resolution" @@ -151,7 +150,12 @@ async def postprocess( prompt_tokens: int, completion_tokens: int, ) -> RolloutState: - """Postprocess a partial rollout using the default semantics.""" + """Postprocess a partial rollout using the default semantics. + + The handler releases the history ref only when the rollout state owns it; refs borrowed from the TraceStore + must never be freed here. The current ref belongs to the caller, who releases it after this call if it created + it. + """ rollout_state.finish_reason = finish_reason rollout_state.status = status history_response = rollout_state.response or "" @@ -165,7 +169,9 @@ async def postprocess( rollout_state.logprobs = history_logprobs + current_logprobs history_routed_experts = rollout_state.routed_experts + history_routed_experts_owner = rollout_state.routed_experts_owner if history_routed_experts is not None and routed_experts is not None: + history_routed_experts_ref = history_routed_experts routed_experts_expect_len = prompt_tokens + completion_tokens - 1 history_routed_experts_expect_len = prompt_tokens - 1 @@ -209,6 +215,17 @@ async def postprocess( f"prompt_tokens={prompt_tokens}, completion_tokens={completion_tokens}" ) rollout_state.routed_experts = ray.put(concat_routed_experts) + rollout_state.routed_experts_owner = "rollout" + # The concatenated result replaces both inputs. Free the current + # ref unconditionally: postprocess is only called by the direct + # rollout path, which produced it. Free the history ref only when + # this state owned it, never a TraceStore borrow still held by + # sibling segments. + free_object_refs(routed_experts) + if history_routed_experts_owner != "trace_store" and isinstance( + history_routed_experts_ref, (RayObjectRef, list) + ): + free_object_refs(history_routed_experts_ref) end_time = time.perf_counter() self.logger.debug( f"[PartialRolloutHandler] Postprocess routed_experts concatenation time: {end_time - start_time:.4f} seconds" @@ -221,6 +238,7 @@ async def postprocess( f"history_logprobs_len={len(history_logprobs)}, " ) rollout_state.routed_experts = routed_experts + rollout_state.routed_experts_owner = "rollout" elif history_routed_experts is not None and routed_experts is None: # case3: 本次推理为超发的任务, token 还未生成时就被 abort了,所以本次 routed_experts 为空,并且response_ids, logprobs 需要也为空 assert not current_response_ids and not current_logprobs, ( diff --git a/xtuner/v1/rl/rollout/vllm.py b/xtuner/v1/rl/rollout/vllm.py index 1c4ae7fd64..18fde3a5c4 100644 --- a/xtuner/v1/rl/rollout/vllm.py +++ b/xtuner/v1/rl/rollout/vllm.py @@ -476,6 +476,8 @@ async def _handle_non_stream_response(self, rollout_state: RolloutState, respons if validation_errors: error_msg = f"Incomplete rollout data for request {uid}: {', '.join(validation_errors)}" self.logger.error(f"{error_msg}. Raw response: {response_json}") + rollout_state.routed_experts = routed_experts + rollout_state.routed_experts_owner = "rollout" if routed_experts is not None else None rollout_state.status = Status.FAILED rollout_state.error_msg = error_msg return rollout_state @@ -484,6 +486,7 @@ async def _handle_non_stream_response(self, rollout_state: RolloutState, respons rollout_state.response_ids = last_token_ids if len(last_token_ids) > 0 else None rollout_state.logprobs = last_logprobs if len(last_logprobs) > 0 else None rollout_state.routed_experts = routed_experts + rollout_state.routed_experts_owner = "rollout" if routed_experts is not None else None rollout_state.finish_reason = finish_reason rollout_state.status = rollout_status diff --git a/xtuner/v1/rl/rollout/worker.py b/xtuner/v1/rl/rollout/worker.py index 7357abad71..7b5293cbc2 100644 --- a/xtuner/v1/rl/rollout/worker.py +++ b/xtuner/v1/rl/rollout/worker.py @@ -1272,6 +1272,7 @@ async def _safe_handle_response(self, rollout_state: RolloutState, http_response error_msg = f"Incomplete rollout data for msg {uid}: {', '.join(validation_errors)}" self.logger.error(error_msg) rollout_state.routed_experts = routed_experts + rollout_state.routed_experts_owner = "rollout" if routed_experts is not None else None rollout_state.status = Status.FAILED rollout_state.error_msg = error_msg return rollout_state @@ -1279,6 +1280,7 @@ async def _safe_handle_response(self, rollout_state: RolloutState, http_response error_msg = f"Rollout failed for msg {uid} with finish_reason {finish_reason}" self.logger.error(error_msg) rollout_state.routed_experts = routed_experts + rollout_state.routed_experts_owner = "rollout" if routed_experts is not None else None rollout_state.status = Status.FAILED rollout_state.error_msg = error_msg return rollout_state @@ -1302,6 +1304,7 @@ async def _safe_handle_response(self, rollout_state: RolloutState, http_response rollout_state.response_ids = response_ids rollout_state.logprobs = logprobs rollout_state.routed_experts = routed_experts + rollout_state.routed_experts_owner = "rollout" if routed_experts is not None else None rollout_state.finish_reason = finish_reason rollout_state.status = rollout_status return rollout_state diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index 0193dd56ea..9ee024d589 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -53,7 +53,7 @@ kl_penalty, ) from xtuner.v1.rl.model_utils import build_frozen_model -from xtuner.v1.rl.utils import SingleAcceleratorWorker +from xtuner.v1.rl.utils import SingleAcceleratorWorker, free_object_refs from xtuner.v1.rl.weight_update import WeightUpdater from xtuner.v1.train.trainer import LoadCheckpointConfig from xtuner.v1.utils import ( @@ -555,7 +555,7 @@ def _add_rollout_routed_experts( ) storage_dtype = _rollout_routed_experts_storage_dtype(language_cfg.n_routed_experts) - to_free_routed_expert_refs: list[ray.ObjectRef] = [] + to_free_routed_expert_refs: list[ray.ObjectRef | list[ray.ObjectRef]] = [] if isinstance(rollout_routed_experts, list): # list[n,l,e] out_rollout_routed_expert = [] @@ -594,7 +594,7 @@ def _add_rollout_routed_experts( # finish consuming the batch. if self.config.free_rollout_routed_experts_in_worker: if self.sp_mesh is None or self.sp_mesh.size() == 1: - ray.internal.free(rollout_routed_expert_refs, local_only=False) + free_object_refs(rollout_routed_expert_refs) else: if self.sp_mesh.get_local_rank() == 0: # only free once of sp mesh @@ -633,7 +633,7 @@ def _add_rollout_routed_experts( if self.config.free_rollout_routed_experts_in_worker and self.sp_mesh is not None and self.sp_mesh.size() > 1: dist.barrier() for free_routed_expert_refs in to_free_routed_expert_refs: - ray.internal.free(free_routed_expert_refs, local_only=False) + free_object_refs(free_routed_expert_refs) del to_free_routed_expert_refs @contextmanager diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 1bc2dbc6d0..ca160493c4 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -1059,23 +1059,24 @@ def _train_one_batch( self.train_controller.onload(target="all") self.logger.info("Training controller loaded") - with timer("prepare_data", step_timer_dict): - data_batches, data_info = self._prepare_train_data( - train_batch, - self._train_worker_cfg.pack_max_length, - raw_rewards_sum=raw_rewards_sum, - raw_rewards_count=raw_rewards_count, - ) - self.logger.info(f"Prepared {len(data_batches)} training data batches") - - with timer("training", step_timer_dict): - workers_log_item: list[WorkerLogItem] = self.train_controller.fit( - data_batches, - pack_max_length=self._train_worker_cfg.pack_max_length, - rollout_idx=train_step, - ) + try: + with timer("prepare_data", step_timer_dict): + data_batches, data_info = self._prepare_train_data( + train_batch, + self._train_worker_cfg.pack_max_length, + raw_rewards_sum=raw_rewards_sum, + raw_rewards_count=raw_rewards_count, + ) + self.logger.info(f"Prepared {len(data_batches)} training data batches") - self._release_trace_sessions_after_train_batch(train_batch) + with timer("training", step_timer_dict): + workers_log_item: list[WorkerLogItem] = self.train_controller.fit( + data_batches, + pack_max_length=self._train_worker_cfg.pack_max_length, + rollout_idx=train_step, + ) + finally: + self._release_trace_sessions_after_train_batch(train_batch) return { "data_info": data_info,