From 50733cc71206057fd592069102f866a1e7072b53 Mon Sep 17 00:00:00 2001 From: matrix72c Date: Fri, 14 Aug 2026 00:00:14 +0800 Subject: [PATCH 1/5] [Fix] Preserve concurrent rollout trace sessions --- tests/rl/test_producer.py | 32 +++++++++- tests/rl/test_replay_buffer.py | 64 +++++++++++++------ tests/rl/test_rl_disaggregated_trainer.py | 32 ++++++++++ .../v1/rl/agent_loop_manager/produce_utils.py | 7 ++ xtuner/v1/rl/replay_buffer.py | 23 ++++++- xtuner/v1/rl/rollout/trace_store.py | 38 +++++++++++ xtuner/v1/train/rl_trainer.py | 43 ++++++++++--- 7 files changed, 207 insertions(+), 32 deletions(-) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 55280dbe8..9bd6a8d76 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -20,7 +20,7 @@ import asyncio import unittest -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch from xtuner.v1.data_proto.rl_data import RolloutState, Status, discard_rollout_state from xtuner.v1.rl.agent_loop_manager import ( @@ -383,6 +383,36 @@ def is_valid_sample_fn(samples): self.assertEqual(result.failed_samples, 1) self.assertEqual(result.filtered_samples, 1) + async def test_put_generated_group_releases_terminal_trace_sessions(self): + # FAILED / FILTERED 不进入 replay buffer,必须在丢弃 RolloutState 前释放对应 trace session。 + cases = ( + (Status.FAILED, True, 101), + (Status.COMPLETED, False, 102), + ) + for status, is_valid, session_id in cases: + with self.subTest(status=status): + strategy = SyncProduceStrategyConfig(is_valid_sample_fn=lambda _samples: is_valid).build() + ctx = self._build_context( + strategy, + f"terminal_{status.name.lower()}", + self._build_agent_loop(), + self._build_sampler(), + batch_size=1, + ) + item = make_rollout_state(session_id, status=status, reward_score=1.0) + item.session_id = session_id + item.routed_experts = MagicMock() + + with patch( + "xtuner.v1.rl.agent_loop_manager.produce_utils.release_existing_sessions", + new=AsyncMock(return_value={str(session_id)}), + ) as release_sessions: + self.assertFalse(await ctx.put_generated_group([item])) + + release_sessions.assert_awaited_once_with([str(session_id)]) + self.assertIsNone(item.session_id) + self.assertIsNone(item.routed_experts) + async def test_put_generated_group_records_raw_rewards_before_filtering(self): # 验证 raw reward 在过滤前统计,filtered group 仍能贡献生成侧 reward 指标。 task_name = "test_raw_reward_before_filter" diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index 69b5c6a2c..e37e7d0e6 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -24,6 +24,7 @@ import tempfile import unittest from pathlib import Path +from unittest.mock import AsyncMock, patch import numpy as np import ray @@ -41,6 +42,7 @@ def make_rollout_state( uid: int, *, + session_id: int | None = None, status: Status = Status.COMPLETED, seq_staleness: int = 0, prompt_ids: list[int] | None = None, @@ -66,6 +68,7 @@ def make_rollout_state( group_id=uid, message=[{"role": "user", "content": f"prompt {uid}"}], prompt_ids=prompt_ids, + session_id=session_id, tokens=list(tokens) if tokens is not None else list(prompt_ids), response=response if response is not None else f"response {uid}", response_ids=response_ids, @@ -193,6 +196,7 @@ async def test_common_put_drops_expired_group_when_tail_batch_is_disabled(self): replay_buffer = replay_buffer_config_cls().build() stale = make_rollout_state( 1, + session_id=101, prompt_ids=[101, 102], tokens=[999], response="stale response", @@ -205,14 +209,19 @@ async def test_common_put_drops_expired_group_when_tail_batch_is_disabled(self): extra_fields={"train_prompt_ids": [101, 102]}, ) - await replay_buffer.put( - [stale], - "task", - current_train_step=5, - stale_threshold=3, - expired_groups_retryable=False, - ) + with patch( + "xtuner.v1.rl.replay_buffer.release_existing_sessions", + new=AsyncMock(return_value={"101"}), + ) as release_sessions: + await replay_buffer.put( + [stale], + "task", + current_train_step=5, + stale_threshold=3, + expired_groups_retryable=False, + ) + release_sessions.assert_awaited_once_with(["101"]) assert stale.status == Status.EXPIRED assert stale.prompt_ids is None assert stale.tokens is None @@ -230,6 +239,7 @@ async def test_common_put_defaults_to_retryable_expired_group(self): pixel_values = np.ones((2, 3), dtype=np.float32) stale = make_rollout_state( 1, + session_id=102, prompt_ids=[101, 102], tokens=[999], response="stale response", @@ -243,13 +253,18 @@ async def test_common_put_defaults_to_retryable_expired_group(self): extra_fields={"train_prompt_ids": [101, 102]}, ) - await replay_buffer.put( - [stale], - "task", - current_train_step=5, - stale_threshold=3, - ) + with patch( + "xtuner.v1.rl.replay_buffer.release_existing_sessions", + new=AsyncMock(), + ) as release_sessions: + await replay_buffer.put( + [stale], + "task", + current_train_step=5, + stale_threshold=3, + ) + release_sessions.assert_not_awaited() expired = await replay_buffer.get(1, "task", Status.EXPIRED) reusable = expired[0][0] assert reusable.status == Status.EXPIRED @@ -419,11 +434,13 @@ async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self) replay_buffer = replay_buffer_config_cls().build() terminal_stale = make_rollout_state( 1, + session_id=201, response_model_steps=[1], mm_info={"pixel_values": np.ones((2, 3), dtype=np.float32)}, ) retryable_stale = make_rollout_state( 2, + session_id=202, response_model_steps=[1], mm_info={"pixel_values": np.ones((2, 3), dtype=np.float32)}, ) @@ -431,15 +448,20 @@ async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self) await replay_buffer.put([retryable_stale], "retryable_task") assert len(replay_buffer) == 2 - expired_counts = await replay_buffer.refresh_staleness( - task_stale_thresholds={"terminal_task": 2, "retryable_task": 2}, - expired_groups_retryable_by_task={ - "terminal_task": False, - "retryable_task": True, - }, - current_train_step=4, - ) + with patch( + "xtuner.v1.rl.replay_buffer.release_existing_sessions", + new=AsyncMock(return_value={"201"}), + ) as release_sessions: + expired_counts = await replay_buffer.refresh_staleness( + task_stale_thresholds={"terminal_task": 2, "retryable_task": 2}, + expired_groups_retryable_by_task={ + "terminal_task": False, + "retryable_task": True, + }, + current_train_step=4, + ) + release_sessions.assert_awaited_once_with(["201"]) assert expired_counts == {"terminal_task": 1, "retryable_task": 1} assert terminal_stale.status == Status.EXPIRED assert terminal_stale.prompt_ids is None diff --git a/tests/rl/test_rl_disaggregated_trainer.py b/tests/rl/test_rl_disaggregated_trainer.py index 80b96c6da..8780d126e 100644 --- a/tests/rl/test_rl_disaggregated_trainer.py +++ b/tests/rl/test_rl_disaggregated_trainer.py @@ -305,6 +305,38 @@ def blocking_train_one_batch(*args, **kwargs): self.assertIn("produce_loop_tick_during_training", manager.calls) self.assertEqual(trainer._cur_step, 1) + def test_train_batch_releases_only_consumed_trace_sessions_for_disaggregated_rollout(self): + # 后台 producer 在 learner 训练期间仍会创建 trace;训练结束只能释放当前消费 batch 的 session。 + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=101) + trainer = self._make_trainer(_FakeManager([])) + trainer._release_trace_store = RLDisaggregatedTrainer._release_trace_store.__get__( + trainer, RLDisaggregatedTrainer + ) + live_session_ids = {"101", "202"} + + def release_sessions(session_ids): + released = [session_id for session_id in session_ids if session_id in live_session_ids] + live_session_ids.difference_update(released) + return released + + store = SimpleNamespace( + release_sessions=SimpleNamespace(remote=MagicMock(side_effect=release_sessions)), + ) + + with ( + patch("xtuner.v1.rl.rollout.trace_store.get_existing_store", return_value=store), + patch("xtuner.v1.train.rl_trainer.ray.get", side_effect=lambda value: value), + ): + trainer._train_one_batch( + [[train_sample]], + train_step=1, + step_timer_dict={}, + release_only_consumed_trace_sessions=True, + ) + + store.release_sessions.remote.assert_called_once_with(["101"]) + self.assertEqual(live_session_ids, {"202"}) + def test_fit_observes_background_producer_failure_before_training_waited_batch(self): # 后台 producer 异常是终止性失败;前台 get_batch 还在等待时必须立刻暴露,不能先训练随后才失败。 train_sample = SimpleNamespace(group_id=1, rollout_id=1) diff --git a/xtuner/v1/rl/agent_loop_manager/produce_utils.py b/xtuner/v1/rl/agent_loop_manager/produce_utils.py index ef9151bbc..0a19dee53 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -21,6 +21,7 @@ ) from xtuner.v1.rl.agent_loop import AgentLoopSpec from xtuner.v1.rl.replay_buffer import ReplayBuffer +from xtuner.v1.rl.rollout.trace_store import release_existing_sessions from xtuner.v1.rl.utils import ( AGENT_LOOP_PAUSE_REQUEST_TIMEOUT_S, PRODUCER_PAUSE_PENDING_TASK_TIMEOUT_S, @@ -197,7 +198,13 @@ async def put_generated_group(self, group: list[RolloutState]) -> bool: # 失败样本和业务过滤样本都不进入 replay buffer。 self.progress.add_produced(self.task_name, samples=len(group), tokens=produced_tokens) self.progress.add_discarded(self.task_name, discard_status, samples=len(group)) + released_session_ids = await release_existing_sessions( + [str(item.session_id) for item in group if item.session_id is not None] + ) for item in group: + if item.session_id is not None and str(item.session_id) in released_session_ids: + # TraceStore.release_sessions() already freed these routed-expert refs. + item.routed_experts = None discard_rollout_state(item) return False diff --git a/xtuner/v1/rl/replay_buffer.py b/xtuner/v1/rl/replay_buffer.py index a4c8f45e3..e5d882ab3 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -20,6 +20,7 @@ reset_rollout_response, update_sample_version, ) +from xtuner.v1.rl.rollout.trace_store import release_existing_sessions from xtuner.v1.rl.utils import ( BetweenNode, ConditionNode, @@ -488,11 +489,27 @@ def _apply_staleness_lifecycle( reset_rollout_response(item) else: for item in group: - discard_rollout_state(item) item.status = Status.EXPIRED return Status.EXPIRED + @staticmethod + async def _discard_terminal_expired_groups(groups: list[list[RolloutState]]) -> None: + """Release terminal trace sessions in one RPC, then discard groups.""" + if not groups: + return + + released_session_ids = await release_existing_sessions( + [str(item.session_id) for group in groups for item in group if item.session_id is not None] + ) + for group in groups: + for item in group: + if item.session_id is not None and str(item.session_id) in released_session_ids: + # TraceStore.release_sessions() already freed these routed-expert refs. + item.routed_experts = None + discard_rollout_state(item) + item.status = Status.EXPIRED + async def put( self, items: list[RolloutState], @@ -518,6 +535,7 @@ async def put( ) staleness = max(item.seq_staleness for item in items) if status == Status.EXPIRED and not expired_groups_retryable: + await self._discard_terminal_expired_groups([items]) return storage_item = StorageItem( item=items, @@ -558,6 +576,7 @@ async def refresh_staleness( expired_counts: dict[str, int] = {} retryable_by_task = expired_groups_retryable_by_task or {} token_stale_thresholds = task_token_stale_thresholds or {} + terminal_expired_groups: list[list[RolloutState]] = [] async with self._lock: updated_records: list[StorageItem] = [] deleted_uids: list[int] = [] @@ -583,12 +602,14 @@ async def refresh_staleness( if status == Status.EXPIRED: expired_count += 1 if not retryable: + terminal_expired_groups.append(record.item) deleted_uids.append(record.uid) continue updated_records.append(replace(record, status=status, staleness=staleness)) expired_counts[task_name] = expired_count await self._storage.delete(deleted_uids) await self._storage.update(updated_records) + await self._discard_terminal_expired_groups(terminal_expired_groups) return expired_counts async def is_ready( diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 8fed1119a..9dc5f53c5 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -335,6 +335,23 @@ def release(self, session_id: str, key: str | None = None): trie = self.sessions.pop(session_id) if key is None else self.sessions[session_id] trie.release(key) + def release_sessions(self, session_ids: list[str]) -> list[str]: + """Release existing trace sessions in one actor call. + + Args: + session_ids (list[str]): Session identifiers that no longer own live rollout state. + + Returns: + list[str]: Session identifiers that existed and were released. + """ + released_session_ids = [] + for session_id in dict.fromkeys(session_ids): + if session_id not in self.sessions: + continue + self.release(session_id) + released_session_ids.append(session_id) + return released_session_ids + def release_all(self): """Release all sessions and free associated resources.""" for session_id in list(self.sessions): @@ -462,6 +479,8 @@ def get_existing_store(): global _handle_cache if _handle_cache is not None: return _handle_cache + if not ray.is_initialized(): + return None try: _handle_cache = ray.get_actor(_STORE_NAME, namespace=_STORE_NAMESPACE) @@ -470,6 +489,25 @@ def get_existing_store(): return _handle_cache +async def release_existing_sessions(session_ids: list[str]) -> set[str]: + """Release trace sessions that exist without creating the singleton store. + + Args: + session_ids (list[str]): Candidate trace session identifiers. + + Returns: + set[str]: Session identifiers that existed and were released. + """ + if not session_ids: + return set() + + store = get_existing_store() + if store is None: + return set() + + return set(await store.release_sessions.remote(session_ids)) + + if __name__ == "__main__": print("=== 评估使用 Trie 加速 tokenize.py 避免多轮对话重复 tokenization ===") diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 1730d68ab..2109da193 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -882,20 +882,37 @@ async def _run_initial_evaluate(self) -> None: finally: self._release_trace_store() - def _release_trace_store(self) -> None: + def _release_trace_store(self, train_batch: list[list[RolloutState]] | None = None) -> None: from xtuner.v1.rl.rollout.trace_store import get_existing_store store = get_existing_store() if store is None: return - self.logger.info("Release all sessions and free associated resources") - ray.get(store.release_all.remote()) - keys = ray.get(store.list_sessions.remote()) - # NOTE: previously asserted ``len(keys) == 0`` here, but a leftover session key should not crash the whole - # fit() at teardown. Warn instead so the leak stays visible without aborting the run. - if keys: - self.logger.warning(f"Trace store keys not released after release_all: {keys}") + if train_batch is None: + self.logger.info("Release all sessions and free associated resources") + ray.get(store.release_all.remote()) + keys = ray.get(store.list_sessions.remote()) + # A leftover session key should stay visible without crashing fit() + # during teardown. + if keys: + self.logger.warning(f"Trace store keys not released after release_all: {keys}") + return + + session_ids = { + str(rollout_state.session_id) + for group in train_batch + for rollout_state in group + if rollout_state.session_id is not None + } + if not session_ids: + return + + released_session_ids = ray.get(store.release_sessions.remote(sorted(session_ids))) + self.logger.info( + "Release consumed trace sessions and preserve concurrent rollout sessions: " + f"released={len(released_session_ids)}, requested={len(session_ids)}" + ) def _train_one_batch( self, @@ -905,6 +922,7 @@ def _train_one_batch( *, offload_rollout_before_train: bool = False, onload_train_before_train: bool = False, + release_only_consumed_trace_sessions: bool = False, raw_rewards_sum: float = 0.0, raw_rewards_count: int = 0, ) -> TrainInfo: @@ -949,7 +967,13 @@ def _train_one_batch( rollout_idx=train_step, ) - self._release_trace_store() + if release_only_consumed_trace_sessions: + # Disaggregated rollout keeps producing while learner.fit() runs. Release + # only the sessions consumed by this batch so concurrent rollout traces + # remain valid until their own batches are trained. + self._release_trace_store(train_batch) + else: + self._release_trace_store() return { "data_info": data_info, @@ -1983,6 +2007,7 @@ async def _fit(self): train_batch, train_step, step_timer_dict, + release_only_consumed_trace_sessions=True, raw_rewards_sum=produce_result.raw_rewards_sum, raw_rewards_count=produce_result.raw_rewards_count, ) From 12b1ebcef69e654191817290959eebe2b9669fc9 Mon Sep 17 00:00:00 2001 From: matrix72c Date: Mon, 17 Aug 2026 10:45:06 +0800 Subject: [PATCH 2/5] [Fix] Address trace-session lifecycle review --- tests/rl/test_replay_buffer.py | 46 ++++++++++++ tests/rl/test_rl_colocate_trainer.py | 5 +- tests/rl/test_rl_disaggregated_trainer.py | 80 +++++++++++++++++---- tests/rl/test_rl_trainer_checkpoint.py | 4 +- tests/rl/test_trace_store.py | 77 ++++++++++++++++++++ xtuner/v1/rl/rollout/trace_store.py | 1 + xtuner/v1/train/rl_trainer.py | 85 +++++++++++++---------- 7 files changed, 247 insertions(+), 51 deletions(-) create mode 100644 tests/rl/test_trace_store.py diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index e37e7d0e6..865166712 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -21,6 +21,7 @@ # 11. save/resume 保留 Ray ObjectRef:直接 ObjectRef 和 dict(dict(ObjectRef)) 嵌套结构恢复后, # 解引用得到的内容都应与保存前一致。 +import asyncio import tempfile import unittest from pathlib import Path @@ -475,6 +476,51 @@ async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self) assert len(replay_buffer) == 1 assert await replay_buffer.get(1, "terminal_task", Status.EXPIRED) == [] + async def test_refresh_staleness_batches_terminal_release_outside_lock(self): + replay_buffer = AsyncReplayBufferConfig().build() + first = make_rollout_state(1, session_id=301, response_model_steps=[1]) + second = make_rollout_state(2, session_id=302, response_model_steps=[1]) + await replay_buffer.put([first], "terminal_task") + await replay_buffer.put([second], "terminal_task") + + release_started = asyncio.Event() + allow_release = asyncio.Event() + + async def delayed_release(session_ids): + release_started.set() + await allow_release.wait() + return set(session_ids) + + with patch( + "xtuner.v1.rl.replay_buffer.release_existing_sessions", + new=AsyncMock(side_effect=delayed_release), + ) as release_sessions: + refresh_task = asyncio.create_task( + replay_buffer.refresh_staleness( + task_stale_thresholds={"terminal_task": 2}, + expired_groups_retryable_by_task={"terminal_task": False}, + current_train_step=4, + ) + ) + await asyncio.wait_for(release_started.wait(), timeout=1.0) + try: + # The terminal records are already removed and the buffer lock is + # available while the trace-store RPC is still blocked. + count = await asyncio.wait_for( + replay_buffer.count("terminal_task", Status.COMPLETED), + timeout=1.0, + ) + assert count == 0 + finally: + allow_release.set() + expired_counts = await refresh_task + + release_sessions.assert_awaited_once_with(["301", "302"]) + assert expired_counts == {"terminal_task": 2} + assert first.status == Status.EXPIRED + assert second.status == Status.EXPIRED + assert len(replay_buffer) == 0 + async def test_common_refresh_staleness_contract(self): # refresh_staleness 同时覆盖默认刷新 completed/aborted,以及 status filter 只刷新指定状态。 for config_name, replay_buffer_config_cls in REPLAY_BUFFER_CONFIGS: diff --git a/tests/rl/test_rl_colocate_trainer.py b/tests/rl/test_rl_colocate_trainer.py index da9a349a9..b31f1dd84 100644 --- a/tests/rl/test_rl_colocate_trainer.py +++ b/tests/rl/test_rl_colocate_trainer.py @@ -142,7 +142,8 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_ trainer._benchmark_training_samples = 0 trainer._benchmark_training_tokens = 0 trainer._save_trajectories = MagicMock() - trainer._release_trace_store = MagicMock() + trainer._release_trace_sessions = MagicMock(return_value=set()) + trainer._release_all_trace_sessions = MagicMock() trainer._sync_weights_and_save = MagicMock( side_effect=lambda train_step, step_timer_dict: train_step % trainer._sync_weights_interval == 0 ) @@ -208,6 +209,8 @@ def test_fit_accepts_async_strategy_manager_on_colocate_path(self): trainer.rollout_controller.offload.remote.assert_called_once_with() trainer.train_controller.onload.assert_called_once_with(target="all") trainer.train_controller.fit.assert_called_once() + trainer._release_all_trace_sessions.assert_called_once_with() + trainer._release_trace_sessions.assert_not_called() self.assertEqual(trainer._cur_step, 1) def test_fit_requires_non_empty_batch_from_manager(self): diff --git a/tests/rl/test_rl_disaggregated_trainer.py b/tests/rl/test_rl_disaggregated_trainer.py index 8780d126e..a0b966754 100644 --- a/tests/rl/test_rl_disaggregated_trainer.py +++ b/tests/rl/test_rl_disaggregated_trainer.py @@ -135,7 +135,8 @@ def _make_trainer(self, agent_loop_manager): ) trainer._save_trajectories = MagicMock() trainer._save_eval_trajectories = MagicMock() - trainer._release_trace_store = MagicMock() + trainer._release_trace_sessions = MagicMock(return_value=set()) + trainer._release_all_trace_sessions = MagicMock() trainer._log_step = MagicMock() trainer._maybe_save_checkpoint = AsyncMock() trainer._maybe_save_hf = MagicMock() @@ -185,7 +186,7 @@ def _minimal_train_info(self, *, training_samples: int, training_tokens: int, be def test_fit_persists_checkpoint_for_completed_model_step(self): # 验证 checkpoint 以 fit 完成的 model_step 为准,并通过 async manager.save 落盘。 - train_sample = SimpleNamespace(group_id=1, rollout_id=1) + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=None) manager = _FakeManager([ProduceBatchResult(rollout_states=[[train_sample]])]) manager.save = AsyncMock() trainer = self._make_trainer(manager) @@ -217,7 +218,7 @@ def test_fit_persists_checkpoint_for_completed_model_step(self): def test_fit_retries_same_step_after_empty_expired_skip(self): # 验证空 expired batch 只同步上一版模型,不推进 train_step,并重试同一步。 - train_sample = SimpleNamespace(group_id=1, rollout_id=1) + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=None) manager = _FakeManager( [ ProduceBatchResult(rollout_states=[], status=ProduceBatchStatus.EXPIRED_BATCH), @@ -243,7 +244,7 @@ def test_fit_retries_same_step_after_empty_expired_skip(self): def test_fit_trains_non_empty_expired_batch_then_syncs_current_step(self): # 验证非空 expired batch 仍会训练,并用当前完成的 model_step 恢复 producer。 - train_sample = SimpleNamespace(group_id=1, rollout_id=1) + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=None) manager = _FakeManager( [ProduceBatchResult(rollout_states=[[train_sample]], status=ProduceBatchStatus.EXPIRED_BATCH)] ) @@ -258,7 +259,7 @@ def test_fit_trains_non_empty_expired_batch_then_syncs_current_step(self): def test_fit_rebinds_weight_update_with_rollout_update_address(self): # 验证非共卡后续同步权重时继续沿用 rollout config 中的 NCCL update 地址。 - train_sample = SimpleNamespace(group_id=1, rollout_id=1) + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=None) manager = _FakeManager([ProduceBatchResult(rollout_states=[[train_sample]])]) trainer = self._make_trainer(manager) trainer._rollout_config = SimpleNamespace(weight_update_host="10.0.0.1", weight_update_port=23456) @@ -280,7 +281,7 @@ def test_fit_rebinds_weight_update_with_rollout_update_address(self): def test_fit_keeps_background_producer_running_while_training_blocks(self): # 验证非共卡训练阻塞在同步训练 batch 时,后台 producer 仍能继续调度。 - train_sample = SimpleNamespace(group_id=1, rollout_id=1) + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=None) training_started = threading.Event() producer_ticked = threading.Event() manager = _TickingManager( @@ -309,7 +310,7 @@ def test_train_batch_releases_only_consumed_trace_sessions_for_disaggregated_rol # 后台 producer 在 learner 训练期间仍会创建 trace;训练结束只能释放当前消费 batch 的 session。 train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=101) trainer = self._make_trainer(_FakeManager([])) - trainer._release_trace_store = RLDisaggregatedTrainer._release_trace_store.__get__( + trainer._release_trace_sessions = RLDisaggregatedTrainer._release_trace_sessions.__get__( trainer, RLDisaggregatedTrainer ) live_session_ids = {"101", "202"} @@ -331,15 +332,68 @@ def release_sessions(session_ids): [[train_sample]], train_step=1, step_timer_dict={}, - release_only_consumed_trace_sessions=True, ) store.release_sessions.remote.assert_called_once_with(["101"]) self.assertEqual(live_session_ids, {"202"}) + def test_evaluation_releases_only_eval_sessions_and_preserves_leftovers(self): + eval_sample = SimpleNamespace(group_id=2, rollout_id=2, session_id="eval") + trainer = self._make_trainer(_FakeManager([])) + trainer._release_trace_sessions = RLDisaggregatedTrainer._release_trace_sessions.__get__( + trainer, RLDisaggregatedTrainer + ) + trainer.eval_agent_loop_manager.produce_batch = AsyncMock( + return_value=ProduceBatchResult(rollout_states=[[eval_sample]]) + ) + live_session_ids = {"leftover", "eval"} + + def release_sessions(session_ids): + released = [session_id for session_id in session_ids if session_id in live_session_ids] + live_session_ids.difference_update(released) + return released + + store = SimpleNamespace( + release_sessions=SimpleNamespace(remote=MagicMock(side_effect=release_sessions)), + ) + with ( + patch("xtuner.v1.rl.rollout.trace_store.get_existing_store", return_value=store), + patch("xtuner.v1.train.rl_trainer.ray.get", side_effect=lambda value: value), + ): + metrics = asyncio.run(trainer._run_evaluation(train_step=1)) + + self.assertEqual(metrics, {"acc": 1.0}) + store.release_sessions.remote.assert_called_once_with(["eval"]) + self.assertEqual(live_session_ids, {"leftover"}) + + def test_evaluation_releases_eval_sessions_when_evaluator_fails(self): + eval_sample = SimpleNamespace(group_id=2, rollout_id=2, session_id="eval") + trainer = self._make_trainer(_FakeManager([])) + trainer.eval_agent_loop_manager.produce_batch = AsyncMock( + return_value=ProduceBatchResult(rollout_states=[[eval_sample]]) + ) + trainer.evaluator.run = MagicMock(side_effect=RuntimeError("evaluation failed")) + + with self.assertRaisesRegex(RuntimeError, "evaluation failed"): + asyncio.run(trainer._run_evaluation(train_step=1)) + + trainer._release_trace_sessions.assert_called_once_with(["eval"]) + + def test_initial_evaluation_releases_only_its_sessions(self): + eval_sample = SimpleNamespace(group_id=2, rollout_id=2, session_id="initial-eval") + trainer = self._make_trainer(_FakeManager([])) + trainer.eval_agent_loop_manager.produce_batch = AsyncMock( + return_value=ProduceBatchResult(rollout_states=[[eval_sample]]) + ) + + asyncio.run(trainer._run_initial_evaluate()) + + trainer._release_trace_sessions.assert_called_once_with(["initial-eval"]) + trainer._release_all_trace_sessions.assert_not_called() + def test_fit_observes_background_producer_failure_before_training_waited_batch(self): # 后台 producer 异常是终止性失败;前台 get_batch 还在等待时必须立刻暴露,不能先训练随后才失败。 - train_sample = SimpleNamespace(group_id=1, rollout_id=1) + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=None) manager = _FailingProducerManager([ProduceBatchResult(rollout_states=[[train_sample]])]) trainer = self._make_trainer(manager) @@ -352,11 +406,9 @@ def test_fit_observes_background_producer_failure_before_training_waited_batch(s def test_fit_runs_eval_before_reset_and_stops_producer(self): # 验证 eval 在 producer 恢复前执行,避免生产侧提前抢占 rollout 资源。 # 确定性排序依赖 RolloutState 的 group_id 和 rollout_id,测试用轻量对象模拟即可。 - train_sample = SimpleNamespace(group_id=1, rollout_id=1) - eval_sample = SimpleNamespace(group_id=2, rollout_id=2) - manager = _FakeManager( - [ProduceBatchResult(rollout_states=[[train_sample]], status=ProduceBatchStatus.NORMAL)] - ) + train_sample = SimpleNamespace(group_id=1, rollout_id=1, session_id=None) + eval_sample = SimpleNamespace(group_id=2, rollout_id=2, session_id=None) + manager = _FakeManager([ProduceBatchResult(rollout_states=[[train_sample]], status=ProduceBatchStatus.NORMAL)]) trainer = self._make_trainer(manager) trainer._enable_evaluate = True events: list[str] = [] diff --git a/tests/rl/test_rl_trainer_checkpoint.py b/tests/rl/test_rl_trainer_checkpoint.py index cb2977b6c..19c836657 100644 --- a/tests/rl/test_rl_trainer_checkpoint.py +++ b/tests/rl/test_rl_trainer_checkpoint.py @@ -45,6 +45,7 @@ from xtuner.v1.rl.utils import AcceleratorResourcesConfig from xtuner.v1.train.rl_trainer import RLColocateTrainerConfig, RLDisaggregatedTrainerConfig + QWEN3_4B_PATH = os.environ.get("QWEN3_4B_PATH") CHECKPOINT_DIR = "checkpoints" TRAIN_STATE_PATH = "train_state.json" @@ -237,7 +238,8 @@ def build_rollout_controller(rollout_cfg, placement_group): patch("xtuner.v1.train.rl_trainer.set_cpu_resource_manager", lambda manager: None), patch("xtuner.v1.train.rl_trainer.get_rollout_engine_version", return_value={}), patch("xtuner.v1.train.rl_trainer.ray.get", side_effect=lambda obj, timeout=None: obj), - patch("xtuner.v1.train.rl_trainer.BaseRLTrainer._release_trace_store", return_value=None), + patch("xtuner.v1.train.rl_trainer.BaseRLTrainer._release_trace_sessions", return_value=set()), + patch("xtuner.v1.train.rl_trainer.BaseRLTrainer._release_all_trace_sessions", return_value=None), patch.object(WorkerConfig, "build", autospec=True, side_effect=build_train_controller), patch.object(RolloutConfig, "build", autospec=True, side_effect=build_rollout_controller), ): diff --git a/tests/rl/test_trace_store.py b/tests/rl/test_trace_store.py new file mode 100644 index 000000000..bf00ead99 --- /dev/null +++ b/tests/rl/test_trace_store.py @@ -0,0 +1,77 @@ +import asyncio +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +import ray + +from xtuner.v1.rl.rollout.trace_store import ( + RolloutTraceStore, + _free_ray_refs, + get_existing_store, + release_existing_sessions, +) + + +class TestRolloutTraceStore(unittest.TestCase): + @classmethod + def setUpClass(cls): + cls.started_ray = False + try: + if not ray.is_initialized(): + ray.init(address="local", num_cpus=1, include_dashboard=False, ignore_reinit_error=True) + cls.started_ray = True + except Exception as exc: + raise unittest.SkipTest(f"Ray init failed for trace-store tests: {exc}") from exc + + @classmethod + def tearDownClass(cls): + if cls.started_ray and ray.is_initialized(): + ray.shutdown() + + def test_release_sessions_deduplicates_and_skips_missing_ids(self): + store = RolloutTraceStore.remote() + try: + ray.get(store.insert.remote("a", "prompt-a", {"value": 1})) + ray.get(store.insert.remote("b", "prompt-b", {"value": 2})) + + released = ray.get(store.release_sessions.remote(["a", "missing", "a"])) + + self.assertEqual(released, ["a"]) + self.assertEqual(ray.get(store.list_sessions.remote()), ["b"]) + finally: + ray.kill(store) + + def test_release_existing_sessions_stably_deduplicates_before_rpc(self): + release_remote = AsyncMock(return_value=["one"]) + store = SimpleNamespace(release_sessions=SimpleNamespace(remote=release_remote)) + with patch( + "xtuner.v1.rl.rollout.trace_store.get_existing_store", + return_value=store, + ): + released = asyncio.run(release_existing_sessions(["one", "one", "missing"])) + + self.assertEqual(released, {"one"}) + release_remote.assert_awaited_once_with(["one", "missing"]) + + def test_release_existing_sessions_handles_empty_input_and_missing_store(self): + with patch("xtuner.v1.rl.rollout.trace_store.get_existing_store") as get_store: + self.assertEqual(asyncio.run(release_existing_sessions([])), set()) + get_store.assert_not_called() + + with patch( + "xtuner.v1.rl.rollout.trace_store.get_existing_store", + return_value=None, + ): + self.assertEqual(asyncio.run(release_existing_sessions(["missing"])), set()) + + def test_get_existing_store_returns_none_when_ray_is_uninitialized(self): + with patch("xtuner.v1.rl.rollout.trace_store.ray.is_initialized", return_value=False): + self.assertIsNone(get_existing_store()) + + def test_free_ray_refs_recurses_into_nested_containers(self): + object_ref = ray.put({"payload": [1, 2, 3]}) + with patch.object(ray.internal, "free") as free: + _free_ray_refs({"outer": [({"inner": object_ref},)]}) + + free.assert_called_once_with([object_ref], local_only=False) diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 9dc5f53c5..8f1a3726e 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -498,6 +498,7 @@ async def release_existing_sessions(session_ids: list[str]) -> set[str]: Returns: set[str]: Session identifiers that existed and were released. """ + session_ids = list(dict.fromkeys(str(session_id) for session_id in session_ids)) if not session_ids: return set() diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 2109da193..60fd1cb72 100644 --- a/xtuner/v1/train/rl_trainer.py +++ b/xtuner/v1/train/rl_trainer.py @@ -95,7 +95,7 @@ def _trainer_config_requires_rollout_proxy(cfg: "BaseRLTrainerConfig") -> bool: ) or _agent_loop_manager_requires_rollout_proxy(cfg.eval_agent_loop_manager_cfg) -# 在使用了 trace_store 情况下,我们不能提前释放 obj ref 而是由 _release_trace_store 统一释放 +# 在使用了 trace_store 情况下,我们不能提前释放 obj ref,而是由 trainer 在消费后统一释放 # 这样可以确保一拆多情况下正确。判断逻辑和 rollout_proxy 一致。 def _agent_loop_manager_uses_trace_store( cfg: AgentLoopManagerConfig | DisaggAgentLoopManagerConfig | None, @@ -105,6 +105,18 @@ def _agent_loop_manager_uses_trace_store( return _agent_loop_manager_requires_rollout_proxy(cfg) +def _trace_session_ids(rollout_batches: list[list[RolloutState]]) -> list[str]: + """Return stable, unique trace session ids owned by rollout batches.""" + return list( + dict.fromkeys( + str(rollout_state.session_id) + for group in rollout_batches + for rollout_state in group + if rollout_state.session_id is not None + ) + ) + + def check_fa3(): if os.environ.get("XTUNER_USE_FA3", "0") != "1": return @@ -859,6 +871,7 @@ def _maybe_save_hf(self, cur_step: int): self.tokenizer.save_pretrained(str(save_hf_path)) async def _run_initial_evaluate(self) -> None: + eval_batch: list[list[RolloutState]] = [] try: eval_produce_result = await self.eval_agent_loop_manager.produce_batch( self.evaluator.eval_batch_size, @@ -880,39 +893,44 @@ async def _run_initial_evaluate(self) -> None: tb_scores = {f"eval/{k}": v for k, v in eval_metrics.items()} self._exp_tracker.add_scalars(tag_scalar_dict=tb_scores, global_step=0) finally: - self._release_trace_store() + self._release_trace_sessions(_trace_session_ids(eval_batch)) - def _release_trace_store(self, train_batch: list[list[RolloutState]] | None = None) -> None: + def _release_trace_sessions(self, session_ids: list[str]) -> set[str]: from xtuner.v1.rl.rollout.trace_store import get_existing_store + session_ids = list(dict.fromkeys(str(session_id) for session_id in session_ids)) + if not session_ids: + return set() + store = get_existing_store() if store is None: - return - - if train_batch is None: - self.logger.info("Release all sessions and free associated resources") - ray.get(store.release_all.remote()) - keys = ray.get(store.list_sessions.remote()) - # A leftover session key should stay visible without crashing fit() - # during teardown. - if keys: - self.logger.warning(f"Trace store keys not released after release_all: {keys}") - return + return set() - session_ids = { - str(rollout_state.session_id) - for group in train_batch - for rollout_state in group - if rollout_state.session_id is not None - } - if not session_ids: - return - - released_session_ids = ray.get(store.release_sessions.remote(sorted(session_ids))) + released_session_ids = ray.get(store.release_sessions.remote(session_ids)) self.logger.info( - "Release consumed trace sessions and preserve concurrent rollout sessions: " + "Release owned trace sessions and preserve sessions held by other consumers: " f"released={len(released_session_ids)}, requested={len(session_ids)}" ) + return set(released_session_ids) + + def _release_all_trace_sessions(self) -> None: + from xtuner.v1.rl.rollout.trace_store import get_existing_store + + store = get_existing_store() + if store is None: + return + + self.logger.info("Release all sessions and free associated resources") + ray.get(store.release_all.remote()) + keys = ray.get(store.list_sessions.remote()) + # A leftover session key should stay visible without crashing fit() + # during teardown. + if keys: + self.logger.warning(f"Trace store keys not released after release_all: {keys}") + + def _release_trace_sessions_after_train_batch(self, train_batch: list[list[RolloutState]]) -> None: + """Release training traces when no concurrent rollout owner exists.""" + self._release_all_trace_sessions() def _train_one_batch( self, @@ -922,7 +940,6 @@ def _train_one_batch( *, offload_rollout_before_train: bool = False, onload_train_before_train: bool = False, - release_only_consumed_trace_sessions: bool = False, raw_rewards_sum: float = 0.0, raw_rewards_count: int = 0, ) -> TrainInfo: @@ -967,13 +984,7 @@ def _train_one_batch( rollout_idx=train_step, ) - if release_only_consumed_trace_sessions: - # Disaggregated rollout keeps producing while learner.fit() runs. Release - # only the sessions consumed by this batch so concurrent rollout traces - # remain valid until their own batches are trained. - self._release_trace_store(train_batch) - else: - self._release_trace_store() + self._release_trace_sessions_after_train_batch(train_batch) return { "data_info": data_info, @@ -981,6 +992,7 @@ def _train_one_batch( } async def _run_evaluation(self, train_step: int) -> dict[str, float]: + eval_batch: list[list[RolloutState]] = [] try: eval_produce_result = await self.eval_agent_loop_manager.produce_batch( self.evaluator.eval_batch_size, @@ -1000,7 +1012,7 @@ async def _run_evaluation(self, train_step: int) -> dict[str, float]: self.logger.info(f"Train step {train_step} eval trajectories saved to {eval_trajectory_path}") return eval_metrics finally: - self._release_trace_store() + self._release_trace_sessions(_trace_session_ids(eval_batch)) def _save_debug_rollout_batch(self, train_batch: list[list[RolloutState]], train_step: int) -> None: assert self._debug_rollout_dir is not None @@ -1874,6 +1886,10 @@ def __init__(self, cfg: RLDisaggregatedTrainerConfig): self._cpu_resource_manager.log_registered_summary() + def _release_trace_sessions_after_train_batch(self, train_batch: list[list[RolloutState]]) -> None: + """Release only consumed traces while the background producer runs.""" + self._release_trace_sessions(_trace_session_ids(train_batch)) + def _build_disaggregated_placement_groups( self, train_resources: AcceleratorResourcesConfig, @@ -2007,7 +2023,6 @@ async def _fit(self): train_batch, train_step, step_timer_dict, - release_only_consumed_trace_sessions=True, raw_rewards_sum=produce_result.raw_rewards_sum, raw_rewards_count=produce_result.raw_rewards_count, ) From 94e4d45c540ebf26d5d208566292b457331091d5 Mon Sep 17 00:00:00 2001 From: matrix72c Date: Tue, 18 Aug 2026 13:41:00 +0800 Subject: [PATCH 3/5] refactor(rl): centralize terminal trace cleanup --- tests/rl/test_producer.py | 2 +- tests/rl/test_replay_buffer.py | 8 ++--- tests/rl/test_trace_store.py | 30 +++++++++++++++++++ .../v1/rl/agent_loop_manager/produce_utils.py | 12 ++------ xtuner/v1/rl/replay_buffer.py | 11 ++----- xtuner/v1/rl/rollout/trace_store.py | 24 +++++++++++++-- 6 files changed, 60 insertions(+), 27 deletions(-) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 9bd6a8d76..c3cf48df8 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -404,7 +404,7 @@ async def test_put_generated_group_releases_terminal_trace_sessions(self): item.routed_experts = MagicMock() with patch( - "xtuner.v1.rl.agent_loop_manager.produce_utils.release_existing_sessions", + "xtuner.v1.rl.rollout.trace_store.release_existing_sessions", new=AsyncMock(return_value={str(session_id)}), ) as release_sessions: self.assertFalse(await ctx.put_generated_group([item])) diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index 865166712..b06ccf315 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -211,7 +211,7 @@ async def test_common_put_drops_expired_group_when_tail_batch_is_disabled(self): ) with patch( - "xtuner.v1.rl.replay_buffer.release_existing_sessions", + "xtuner.v1.rl.rollout.trace_store.release_existing_sessions", new=AsyncMock(return_value={"101"}), ) as release_sessions: await replay_buffer.put( @@ -255,7 +255,7 @@ async def test_common_put_defaults_to_retryable_expired_group(self): ) with patch( - "xtuner.v1.rl.replay_buffer.release_existing_sessions", + "xtuner.v1.rl.rollout.trace_store.release_existing_sessions", new=AsyncMock(), ) as release_sessions: await replay_buffer.put( @@ -450,7 +450,7 @@ async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self) assert len(replay_buffer) == 2 with patch( - "xtuner.v1.rl.replay_buffer.release_existing_sessions", + "xtuner.v1.rl.rollout.trace_store.release_existing_sessions", new=AsyncMock(return_value={"201"}), ) as release_sessions: expired_counts = await replay_buffer.refresh_staleness( @@ -492,7 +492,7 @@ async def delayed_release(session_ids): return set(session_ids) with patch( - "xtuner.v1.rl.replay_buffer.release_existing_sessions", + "xtuner.v1.rl.rollout.trace_store.release_existing_sessions", new=AsyncMock(side_effect=delayed_release), ) as release_sessions: refresh_task = asyncio.create_task( diff --git a/tests/rl/test_trace_store.py b/tests/rl/test_trace_store.py index bf00ead99..62efcb34d 100644 --- a/tests/rl/test_trace_store.py +++ b/tests/rl/test_trace_store.py @@ -9,10 +9,40 @@ RolloutTraceStore, _free_ray_refs, get_existing_store, + release_and_discard_rollout_groups, release_existing_sessions, ) +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) + routed_experts_seen_by_discard = {} + + def record_discard(item): + routed_experts_seen_by_discard[item.session_id] = item.routed_experts + + with ( + patch( + "xtuner.v1.rl.rollout.trace_store.release_existing_sessions", + new=AsyncMock(return_value={"trace-owned"}), + ) as release_sessions, + patch( + "xtuner.v1.rl.rollout.trace_store.discard_rollout_state", + side_effect=record_discard, + ) as discard, + ): + asyncio.run(release_and_discard_rollout_groups([[trace_owned, rollout_owned]])) + + 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.assertEqual(discard.call_count, 2) + + class TestRolloutTraceStore(unittest.TestCase): @classmethod def setUpClass(cls): diff --git a/xtuner/v1/rl/agent_loop_manager/produce_utils.py b/xtuner/v1/rl/agent_loop_manager/produce_utils.py index 0a19dee53..6639256c0 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -16,12 +16,11 @@ RolloutState, Status, calculate_group_effective_response_masks, - discard_rollout_state, get_group_status, ) from xtuner.v1.rl.agent_loop import AgentLoopSpec from xtuner.v1.rl.replay_buffer import ReplayBuffer -from xtuner.v1.rl.rollout.trace_store import release_existing_sessions +from xtuner.v1.rl.rollout.trace_store import release_and_discard_rollout_groups from xtuner.v1.rl.utils import ( AGENT_LOOP_PAUSE_REQUEST_TIMEOUT_S, PRODUCER_PAUSE_PENDING_TASK_TIMEOUT_S, @@ -198,14 +197,7 @@ async def put_generated_group(self, group: list[RolloutState]) -> bool: # 失败样本和业务过滤样本都不进入 replay buffer。 self.progress.add_produced(self.task_name, samples=len(group), tokens=produced_tokens) self.progress.add_discarded(self.task_name, discard_status, samples=len(group)) - released_session_ids = await release_existing_sessions( - [str(item.session_id) for item in group if item.session_id is not None] - ) - for item in group: - if item.session_id is not None and str(item.session_id) in released_session_ids: - # TraceStore.release_sessions() already freed these routed-expert refs. - item.routed_experts = None - discard_rollout_state(item) + await release_and_discard_rollout_groups([group]) return False # ABORTED 保持可重试;EXPIRED 由 task 的 retryability 决定保留或丢弃。 diff --git a/xtuner/v1/rl/replay_buffer.py b/xtuner/v1/rl/replay_buffer.py index e5d882ab3..59b80d923 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -14,13 +14,12 @@ RolloutState, Status, calculate_group_effective_response_masks, - discard_rollout_state, get_group_status, refresh_seq_staleness, reset_rollout_response, update_sample_version, ) -from xtuner.v1.rl.rollout.trace_store import release_existing_sessions +from xtuner.v1.rl.rollout.trace_store import release_and_discard_rollout_groups from xtuner.v1.rl.utils import ( BetweenNode, ConditionNode, @@ -499,15 +498,9 @@ async def _discard_terminal_expired_groups(groups: list[list[RolloutState]]) -> if not groups: return - released_session_ids = await release_existing_sessions( - [str(item.session_id) for group in groups for item in group if item.session_id is not None] - ) + await release_and_discard_rollout_groups(groups) for group in groups: for item in group: - if item.session_id is not None and str(item.session_id) in released_session_ids: - # TraceStore.release_sessions() already freed these routed-expert refs. - item.routed_experts = None - discard_rollout_state(item) item.status = Status.EXPIRED async def put( diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 8f1a3726e..903cfe9da 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -5,6 +5,7 @@ import ray from pydantic import BaseModel, ConfigDict, Field +from xtuner.v1.data_proto.rl_data import RolloutState, discard_rollout_state from xtuner.v1.utils import get_logger @@ -354,9 +355,7 @@ def release_sessions(self, session_ids: list[str]) -> list[str]: def release_all(self): """Release all sessions and free associated resources.""" - for session_id in list(self.sessions): - self.release(session_id) - self.sessions.clear() + self.release_sessions(list(self.sessions)) self.objects.clear() self.updated_at.clear() @@ -509,6 +508,25 @@ async def release_existing_sessions(session_ids: list[str]) -> set[str]: return set(await store.release_sessions.remote(session_ids)) +async def release_and_discard_rollout_groups(groups: list[list[RolloutState]]) -> None: + """Release trace-owned resources before discarding terminal rollouts. + + Sessions released by the trace store have already freed their routed-expert + references. Detach those references before the generic rollout-state + cleanup so it does not explicitly free them a second time. Rollouts whose + sessions are absent from the store retain their references for the generic + cleanup path. + """ + released_session_ids = await release_existing_sessions( + [str(item.session_id) for group in groups for item in group if item.session_id is not None] + ) + for group in groups: + 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) + + if __name__ == "__main__": print("=== 评估使用 Trie 加速 tokenize.py 避免多轮对话重复 tokenization ===") From eebda23ef238708105d23eddf04fe305ede10827 Mon Sep 17 00:00:00 2001 From: matrix72c Date: Tue, 18 Aug 2026 15:42:29 +0800 Subject: [PATCH 4/5] refactor(rl): use non-retryable cleanup terminology --- tests/rl/test_producer.py | 4 ++-- tests/rl/test_replay_buffer.py | 42 +++++++++++++++++----------------- xtuner/v1/rl/replay_buffer.py | 12 +++++----- 3 files changed, 29 insertions(+), 29 deletions(-) diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index c3cf48df8..89f8674ff 100644 --- a/tests/rl/test_producer.py +++ b/tests/rl/test_producer.py @@ -383,7 +383,7 @@ def is_valid_sample_fn(samples): self.assertEqual(result.failed_samples, 1) self.assertEqual(result.filtered_samples, 1) - async def test_put_generated_group_releases_terminal_trace_sessions(self): + async def test_put_generated_group_releases_non_retryable_trace_sessions(self): # FAILED / FILTERED 不进入 replay buffer,必须在丢弃 RolloutState 前释放对应 trace session。 cases = ( (Status.FAILED, True, 101), @@ -394,7 +394,7 @@ async def test_put_generated_group_releases_terminal_trace_sessions(self): strategy = SyncProduceStrategyConfig(is_valid_sample_fn=lambda _samples: is_valid).build() ctx = self._build_context( strategy, - f"terminal_{status.name.lower()}", + f"non_retryable_{status.name.lower()}", self._build_agent_loop(), self._build_sampler(), batch_size=1, diff --git a/tests/rl/test_replay_buffer.py b/tests/rl/test_replay_buffer.py index b06ccf315..737ca09d4 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -428,12 +428,12 @@ async def test_common_refresh_token_expiry_moves_mixed_group_to_expired_pool(sel self.assertEqual(group[1].response_ids, [12]) self.assertEqual(group[1].reward, {"score": 0.9}) - async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self): - # 同一轮 refresh 仍统计两类过期;只删除 terminal EXPIRED,保留 tail batch 可重试项。 + async def test_common_refresh_staleness_drops_only_non_retryable_expired_groups(self): + # 同一轮 refresh 仍统计两类过期;只删除 non-retryable EXPIRED,保留 tail batch 可重试项。 for config_name, replay_buffer_config_cls in REPLAY_BUFFER_CONFIGS: with self.subTest(replay_buffer_config=config_name): replay_buffer = replay_buffer_config_cls().build() - terminal_stale = make_rollout_state( + non_retryable_stale = make_rollout_state( 1, session_id=201, response_model_steps=[1], @@ -445,7 +445,7 @@ async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self) response_model_steps=[1], mm_info={"pixel_values": np.ones((2, 3), dtype=np.float32)}, ) - await replay_buffer.put([terminal_stale], "terminal_task") + await replay_buffer.put([non_retryable_stale], "non_retryable_task") await replay_buffer.put([retryable_stale], "retryable_task") assert len(replay_buffer) == 2 @@ -454,34 +454,34 @@ async def test_common_refresh_staleness_drops_only_terminal_expired_groups(self) new=AsyncMock(return_value={"201"}), ) as release_sessions: expired_counts = await replay_buffer.refresh_staleness( - task_stale_thresholds={"terminal_task": 2, "retryable_task": 2}, + task_stale_thresholds={"non_retryable_task": 2, "retryable_task": 2}, expired_groups_retryable_by_task={ - "terminal_task": False, + "non_retryable_task": False, "retryable_task": True, }, current_train_step=4, ) release_sessions.assert_awaited_once_with(["201"]) - assert expired_counts == {"terminal_task": 1, "retryable_task": 1} - assert terminal_stale.status == Status.EXPIRED - assert terminal_stale.prompt_ids is None - assert terminal_stale.mm_info is None + assert expired_counts == {"non_retryable_task": 1, "retryable_task": 1} + assert non_retryable_stale.status == Status.EXPIRED + assert non_retryable_stale.prompt_ids is None + assert non_retryable_stale.mm_info is None assert retryable_stale.status == Status.EXPIRED assert retryable_stale.prompt_ids == [2, 1002] assert retryable_stale.mm_info is not None - assert await replay_buffer.count("terminal_task", Status.COMPLETED) == 0 - assert await replay_buffer.count("terminal_task", Status.EXPIRED) == 0 + assert await replay_buffer.count("non_retryable_task", Status.COMPLETED) == 0 + assert await replay_buffer.count("non_retryable_task", Status.EXPIRED) == 0 assert await replay_buffer.count("retryable_task", Status.EXPIRED) == 1 assert len(replay_buffer) == 1 - assert await replay_buffer.get(1, "terminal_task", Status.EXPIRED) == [] + assert await replay_buffer.get(1, "non_retryable_task", Status.EXPIRED) == [] - async def test_refresh_staleness_batches_terminal_release_outside_lock(self): + async def test_refresh_staleness_batches_non_retryable_release_outside_lock(self): replay_buffer = AsyncReplayBufferConfig().build() first = make_rollout_state(1, session_id=301, response_model_steps=[1]) second = make_rollout_state(2, session_id=302, response_model_steps=[1]) - await replay_buffer.put([first], "terminal_task") - await replay_buffer.put([second], "terminal_task") + await replay_buffer.put([first], "non_retryable_task") + await replay_buffer.put([second], "non_retryable_task") release_started = asyncio.Event() allow_release = asyncio.Event() @@ -497,17 +497,17 @@ async def delayed_release(session_ids): ) as release_sessions: refresh_task = asyncio.create_task( replay_buffer.refresh_staleness( - task_stale_thresholds={"terminal_task": 2}, - expired_groups_retryable_by_task={"terminal_task": False}, + task_stale_thresholds={"non_retryable_task": 2}, + expired_groups_retryable_by_task={"non_retryable_task": False}, current_train_step=4, ) ) await asyncio.wait_for(release_started.wait(), timeout=1.0) try: - # The terminal records are already removed and the buffer lock is + # The non-retryable records are already removed and the buffer lock is # available while the trace-store RPC is still blocked. count = await asyncio.wait_for( - replay_buffer.count("terminal_task", Status.COMPLETED), + replay_buffer.count("non_retryable_task", Status.COMPLETED), timeout=1.0, ) assert count == 0 @@ -516,7 +516,7 @@ async def delayed_release(session_ids): expired_counts = await refresh_task release_sessions.assert_awaited_once_with(["301", "302"]) - assert expired_counts == {"terminal_task": 2} + assert expired_counts == {"non_retryable_task": 2} assert first.status == Status.EXPIRED assert second.status == Status.EXPIRED assert len(replay_buffer) == 0 diff --git a/xtuner/v1/rl/replay_buffer.py b/xtuner/v1/rl/replay_buffer.py index 59b80d923..7beba6822 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -493,8 +493,8 @@ def _apply_staleness_lifecycle( return Status.EXPIRED @staticmethod - async def _discard_terminal_expired_groups(groups: list[list[RolloutState]]) -> None: - """Release terminal trace sessions in one RPC, then discard groups.""" + async def _discard_non_retryable_expired_groups(groups: list[list[RolloutState]]) -> None: + """Release non-retryable trace sessions in one RPC, then discard groups.""" if not groups: return @@ -528,7 +528,7 @@ async def put( ) staleness = max(item.seq_staleness for item in items) if status == Status.EXPIRED and not expired_groups_retryable: - await self._discard_terminal_expired_groups([items]) + await self._discard_non_retryable_expired_groups([items]) return storage_item = StorageItem( item=items, @@ -569,7 +569,7 @@ async def refresh_staleness( expired_counts: dict[str, int] = {} retryable_by_task = expired_groups_retryable_by_task or {} token_stale_thresholds = task_token_stale_thresholds or {} - terminal_expired_groups: list[list[RolloutState]] = [] + non_retryable_expired_groups: list[list[RolloutState]] = [] async with self._lock: updated_records: list[StorageItem] = [] deleted_uids: list[int] = [] @@ -595,14 +595,14 @@ async def refresh_staleness( if status == Status.EXPIRED: expired_count += 1 if not retryable: - terminal_expired_groups.append(record.item) + non_retryable_expired_groups.append(record.item) deleted_uids.append(record.uid) continue updated_records.append(replace(record, status=status, staleness=staleness)) expired_counts[task_name] = expired_count await self._storage.delete(deleted_uids) await self._storage.update(updated_records) - await self._discard_terminal_expired_groups(terminal_expired_groups) + await self._discard_non_retryable_expired_groups(non_retryable_expired_groups) return expired_counts async def is_ready( From 3dfed2b019d60fe32f043274811778ada31aea8c Mon Sep 17 00:00:00 2001 From: matrix72c Date: Tue, 18 Aug 2026 15:54:46 +0800 Subject: [PATCH 5/5] style(rl): apply docformatter --- xtuner/v1/rl/replay_buffer.py | 3 ++- xtuner/v1/rl/rollout/trace_store.py | 8 +++----- 2 files changed, 5 insertions(+), 6 deletions(-) diff --git a/xtuner/v1/rl/replay_buffer.py b/xtuner/v1/rl/replay_buffer.py index 7beba6822..61a2abacc 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -494,7 +494,8 @@ def _apply_staleness_lifecycle( @staticmethod async def _discard_non_retryable_expired_groups(groups: list[list[RolloutState]]) -> None: - """Release non-retryable trace sessions in one RPC, then discard groups.""" + """Release non-retryable trace sessions in one RPC, then discard + groups.""" if not groups: return diff --git a/xtuner/v1/rl/rollout/trace_store.py b/xtuner/v1/rl/rollout/trace_store.py index 903cfe9da..045f094c9 100644 --- a/xtuner/v1/rl/rollout/trace_store.py +++ b/xtuner/v1/rl/rollout/trace_store.py @@ -511,11 +511,9 @@ async def release_existing_sessions(session_ids: list[str]) -> set[str]: async def release_and_discard_rollout_groups(groups: list[list[RolloutState]]) -> None: """Release trace-owned resources before discarding terminal rollouts. - Sessions released by the trace store have already freed their routed-expert - references. Detach those references before the generic rollout-state - cleanup so it does not explicitly free them a second time. Rollouts whose - sessions are absent from the store retain their references for the generic - cleanup path. + Sessions released by the trace store have already freed their routed-expert references. Detach those references + before the generic rollout-state cleanup so it does not explicitly free them a second time. Rollouts whose sessions + are absent from the store retain their references for the generic cleanup path. """ released_session_ids = await release_existing_sessions( [str(item.session_id) for group in groups for item in group if item.session_id is not None]