From b68b22106be958ae3b91ad3a0dc8ae0cfed11ea4 Mon Sep 17 00:00:00 2001 From: matrix72c Date: Thu, 3 Sep 2026 16:25:24 +0800 Subject: [PATCH 1/2] fix(rl): transfer routed experts release ownership --- .../verl_agent/common/agent_loop_verl_tool.py | 3 +- tests/rl/test_producer.py | 3 +- tests/rl/test_replay_buffer.py | 54 +++++++++++++ tests/rl/test_rollout_logic.py | 54 ++++++++++++- tests/rl/test_staleness_policy.py | 56 ++++++++++++++ tests/rl/test_trace_store.py | 40 +++++++++- xtuner/v1/data_proto/rl_data.py | 74 ++++++++++++++---- .../agent_in_localhost_loop.py | 2 + .../agent_in_sandbox_loop.py | 2 + xtuner/v1/rl/replay_buffer.py | 6 ++ xtuner/v1/rl/rollout/trace_store.py | 77 ++++++++++++++++--- xtuner/v1/rl/rollout/utils.py | 25 +++++- xtuner/v1/rl/rollout/vllm.py | 3 + xtuner/v1/rl/rollout/worker.py | 8 ++ xtuner/v1/rl/trainer/worker.py | 8 +- xtuner/v1/train/rl_trainer.py | 33 ++++---- 16 files changed, 395 insertions(+), 53 deletions(-) diff --git a/recipe/verl_agent/common/agent_loop_verl_tool.py b/recipe/verl_agent/common/agent_loop_verl_tool.py index 6956d83f8c..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) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 364db9f09c..68dfb6b949 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -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) diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index 737ca09d4a..f0b9da10ee 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -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: diff --git a/tests/rl/test_rollout_logic.py b/tests/rl/test_rollout_logic.py index 95b3898f1c..056e8e9061 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。 @@ -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 @@ -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, ) @@ -1763,6 +1764,7 @@ async def _resolve(): status=Status.ABORTED, prompt_tokens=3, completion_tokens=1, + release_input_routed_experts=True, ) self.assertIs(out.routed_experts, concat_ref) @@ -1770,3 +1772,51 @@ async def _resolve(): 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") diff --git a/tests/rl/test_staleness_policy.py b/tests/rl/test_staleness_policy.py index b9a746d1c8..34c6767d97 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,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( *, diff --git a/tests/rl/test_trace_store.py b/tests/rl/test_trace_store.py index 51ad063362..480669dc5b 100644 --- a/tests/rl/test_trace_store.py +++ b/tests/rl/test_trace_store.py @@ -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( @@ -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): @@ -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)) diff --git a/xtuner/v1/data_proto/rl_data.py b/xtuner/v1/data_proto/rl_data.py index ec00526144..44e1d20967 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] @@ -110,6 +113,10 @@ class RolloutState(BaseModel): response_ids: list[int] | None = None logprobs: list[float] | 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 @@ -176,13 +183,62 @@ 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 release_owned_routed_experts(rollout_state: RolloutState) -> None: + """Release routed-expert refs owned by the direct rollout path. -def discard_rollout_state(rollout_state: RolloutState) -> RolloutState: - """Release heavy references and clear fields before dropping a rollout.""" + ``RolloutState`` can also contain refs borrowed from ``TraceStore``. This + helper is intentionally conservative and refuses to free those refs; the + TraceStore session is their owner and must release them through its own + lifecycle. Callers should invoke this helper only when they are disposing + a direct rollout result. + """ + + routed_experts = rollout_state.routed_experts + if routed_experts is None: + rollout_state.routed_experts_owner = None + return - free_rollout_state_refs(rollout_state) + if rollout_state.routed_experts_owner == "trace_store": + logger.warning( + "Refusing to free TraceStore-owned routed_experts from RolloutState " + f"(session_id={rollout_state.session_id!r}). Release the TraceStore session instead." + ) + return + + if rollout_state.routed_experts_owner != "rollout": + logger.warning( + "Skipping release of routed_experts with unknown owner " + f"{rollout_state.routed_experts_owner!r} (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) + rollout_state.routed_experts = None + rollout_state.routed_experts_owner = None + + +def discard_rollout_state(rollout_state: RolloutState, *, release_refs: bool = False) -> RolloutState: + """Clear a rollout before dropping it. + + Resource release is opt-in so the caller has to make the ownership + decision. TraceStore-owned routed experts are detached before generic + cleanup so ``free_rollout_state_refs`` cannot free them a second time. + """ + + if release_refs: + 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(): if field.is_required(): @@ -230,21 +286,13 @@ def update_status_from_finish_reason(finish_reason: str | None) -> Status: 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 - - from xtuner.v1.rl.utils.ray_utils import free_object_refs - - if isinstance(routed_experts, (ObjectRef, list)): - free_object_refs(routed_experts) - rollout_state.routed_experts = None 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 = "" rollout_state.response_ids = [] rollout_state.logprobs = [] 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 9da3701bef..9b739b6764 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 @@ -261,6 +261,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 "") @@ -275,6 +276,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 75e6e13f23..c3bb8b143e 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 @@ -346,6 +346,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: @@ -366,6 +367,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/replay_buffer.py b/xtuner/v1/rl/replay_buffer.py index 61a2abacce..b10b0a2557 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -16,6 +16,7 @@ calculate_group_effective_response_masks, get_group_status, refresh_seq_staleness, + release_owned_routed_experts, reset_rollout_response, update_sample_version, ) @@ -485,6 +486,11 @@ def _apply_staleness_lifecycle( for item, expired in zip(group, expired_mask): if expired: item.status = Status.EXPIRED + # Direct rollout refs belong to this state/rollout path; + # TraceStore refs remain borrowed until the session is + # released by TraceStore. + if item.routed_experts_owner == "rollout": + release_owned_routed_experts(item) reset_rollout_response(item) else: for item in group: diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 36ba406e0a..f57da48655 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -14,32 +14,68 @@ _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 _collect_ray_ref_keys(obj: Any, keys: set[str] | None = None) -> set[str]: + """Collect all Ray object-reference keys recursively from ``obj``.""" + keys = set() if keys is None else keys + if isinstance(obj, ray.ObjectRef): + keys.add(_ray_ref_key(obj)) + elif isinstance(obj, dict): + for value in obj.values(): + _collect_ray_ref_keys(value, keys) + elif isinstance(obj, (list, tuple, set)): + for value in obj: + _collect_ray_ref_keys(value, keys) + elif hasattr(obj, "model_dump") and callable(getattr(obj, "model_dump")): + _collect_ray_ref_keys(obj.model_dump(), keys) + elif hasattr(obj, "dict") and callable(getattr(obj, "dict")): + _collect_ray_ref_keys(obj.dict(), keys) + elif hasattr(obj, "__dict__"): + _collect_ray_ref_keys(vars(obj), keys) + return keys + + +def _free_ray_refs( + obj: Any, + *, + _seen: set[str] | None = None, + _exclude: 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 + exclude = _exclude or set() if isinstance(obj, ray.ObjectRef): + ref_key = _ray_ref_key(obj) + if ref_key in seen or ref_key in exclude: + 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, _exclude=exclude) elif isinstance(obj, (list, tuple)): for v in obj: - _free_ray_refs(v) + _free_ray_refs(v, _seen=seen, _exclude=exclude) 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, _exclude=exclude) 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, _exclude=exclude) elif hasattr(obj, "__dict__"): for v in vars(obj).values(): - _free_ray_refs(v) + _free_ray_refs(v, _seen=seen, _exclude=exclude) def _common_prefix_len(left: str, right: str) -> int: @@ -173,6 +209,24 @@ def insert(self, key: str, value: Any) -> None: node = node.children[key] break + old_value = node.value + if old_value is not None and old_value is not value: + # A rerolled turn may overwrite an existing key. Release refs + # that are no longer reachable, while preserving refs shared by + # another trie value or by the replacement value itself. + retained_refs = _collect_ray_ref_keys(value) + + def collect_other_values(current: TreeNode) -> None: + if current is not node and current.value is not None: + _collect_ray_ref_keys(current.value, retained_refs) + for child in current.children.values(): + collect_other_values(child) + + collect_other_values(self.root) + if retained_refs: + _free_ray_refs(old_value, _exclude=retained_refs) + else: + _free_ray_refs(old_value) node.value = value def search(self, text: str, filter_none: bool = False) -> Tuple[str, List["TreeNode"]]: @@ -220,16 +274,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,7 +577,8 @@ 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 - discard_rollout_state(item) + item.routed_experts_owner = None + discard_rollout_state(item, release_refs=True) if __name__ == "__main__": diff --git a/xtuner/v1/rl/rollout/utils.py b/xtuner/v1/rl/rollout/utils.py index 94ccc8023a..33505f1c6b 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" @@ -150,8 +149,13 @@ async def postprocess( status: Status, prompt_tokens: int, completion_tokens: int, + release_input_routed_experts: bool = False, ) -> RolloutState: - """Postprocess a partial rollout using the default semantics.""" + """Postprocess a partial rollout using the default semantics. + + The handler only releases the history/current refs when the direct rollout caller explicitly opts in. + TraceStore refs are borrowed and must never be released here. + """ rollout_state.finish_reason = finish_reason rollout_state.status = status history_response = rollout_state.response or "" @@ -165,7 +169,10 @@ 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 + current_routed_experts_ref = routed_experts routed_experts_expect_len = prompt_tokens + completion_tokens - 1 history_routed_experts_expect_len = prompt_tokens - 1 @@ -209,6 +216,19 @@ 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" + if release_input_routed_experts: + if history_routed_experts_owner == "rollout": + free_object_refs( + history_routed_experts_ref + if isinstance(history_routed_experts_ref, list) + else [history_routed_experts_ref] + ) + free_object_refs( + current_routed_experts_ref + if isinstance(current_routed_experts_ref, list) + else [current_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 +241,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 21571c920f..1611549c3f 100644 --- a/xtuner/v1/rl/rollout/vllm.py +++ b/xtuner/v1/rl/rollout/vllm.py @@ -443,6 +443,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 @@ -451,6 +453,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..8bb93c35bc 100644 --- a/xtuner/v1/rl/rollout/worker.py +++ b/xtuner/v1/rl/rollout/worker.py @@ -25,6 +25,7 @@ RolloutState, SampleParams, Status, + release_owned_routed_experts, reset_rollout_response, update_status_from_finish_reason, ) @@ -1031,6 +1032,7 @@ async def generate(self, rollout_state: RolloutState) -> RolloutState: if rollout_state.status == Status.FAILED: error_msg = rollout_state.error_msg status = rollout_state.status + release_owned_routed_experts(rollout_state) reset_rollout_response(rollout_state) rollout_state.status = status rollout_state.error_msg = error_msg @@ -1056,6 +1058,7 @@ def _prepare_request_payload( ``max_tokens``. """ if discard_response: + release_owned_routed_experts(rollout_state) rollout_state = reset_rollout_response(rollout_state) rollout_state.sample_params = rollout_state.sample_params.model_copy( update={"max_tokens": request_max_tokens} @@ -1063,6 +1066,7 @@ def _prepare_request_payload( rollout_state.status = Status.INIT elif not self.enable_partial_rollout and rollout_state.status == Status.ABORTED: # ABORTED samples can be replayed; without partial rollout, rerun from the original prompt. + release_owned_routed_experts(rollout_state) rollout_state = reset_rollout_response(rollout_state) rollout_state.sample_params = rollout_state.sample_params.model_copy( update={"max_tokens": request_max_tokens} @@ -1272,6 +1276,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 +1284,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 @@ -1296,12 +1302,14 @@ async def _safe_handle_response(self, rollout_state: RolloutState, http_response status=rollout_status, prompt_tokens=prompt_tokens, completion_tokens=completion_tokens, + release_input_routed_experts=True, ) else: rollout_state.response = returned_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 1cb2cbe65d..aa4086acea 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -45,7 +45,7 @@ from xtuner.v1.model.utils.misc import ModelForwardExtraLogInfo from xtuner.v1.profiler import profiling_memory, profiling_time from xtuner.v1.rl.loss import BaseRLLossConfig, BaseRLLossContext, finalize_train_policy_metrics, kl_penalty -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 ( @@ -523,7 +523,7 @@ def _add_rollout_routed_experts( else self.config.model_cfg ) - 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 = [] @@ -561,7 +561,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 @@ -596,7 +596,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 49dee2b650..1096748f26 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -1032,23 +1032,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, From c2b2b5d73ad4008e6538c8800439a59726a954d6 Mon Sep 17 00:00:00 2001 From: matrix72c Date: Wed, 9 Sep 2026 22:11:07 +0800 Subject: [PATCH 2/2] ci: retrigger unit tests