diff --git a/tests/rl/test_producer.py b/tests/rl/test_producer.py index 55280dbe8..89f8674ff 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_non_retryable_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"non_retryable_{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.rollout.trace_store.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..737ca09d4 100644 --- a/tests/rl/test_replay_buffer.py +++ b/tests/rl/test_replay_buffer.py @@ -21,9 +21,11 @@ # 11. save/resume 保留 Ray ObjectRef:直接 ObjectRef 和 dict(dict(ObjectRef)) 嵌套结构恢复后, # 解引用得到的内容都应与保存前一致。 +import asyncio import tempfile import unittest from pathlib import Path +from unittest.mock import AsyncMock, patch import numpy as np import ray @@ -41,6 +43,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 +69,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 +197,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 +210,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.rollout.trace_store.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 +240,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 +254,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.rollout.trace_store.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 @@ -412,46 +428,98 @@ 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], 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)}, ) - 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 - 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.rollout.trace_store.release_existing_sessions", + new=AsyncMock(return_value={"201"}), + ) as release_sessions: + expired_counts = await replay_buffer.refresh_staleness( + task_stale_thresholds={"non_retryable_task": 2, "retryable_task": 2}, + expired_groups_retryable_by_task={ + "non_retryable_task": False, + "retryable_task": True, + }, + current_train_step=4, + ) - 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 + release_sessions.assert_awaited_once_with(["201"]) + 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_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], "non_retryable_task") + await replay_buffer.put([second], "non_retryable_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.rollout.trace_store.release_existing_sessions", + new=AsyncMock(side_effect=delayed_release), + ) as release_sessions: + refresh_task = asyncio.create_task( + replay_buffer.refresh_staleness( + 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 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("non_retryable_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 == {"non_retryable_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 只刷新指定状态。 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 80b96c6da..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( @@ -305,9 +306,94 @@ 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_sessions = RLDisaggregatedTrainer._release_trace_sessions.__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={}, + ) + + 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) @@ -320,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..62efcb34d --- /dev/null +++ b/tests/rl/test_trace_store.py @@ -0,0 +1,107 @@ +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_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): + 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/agent_loop_manager/produce_utils.py b/xtuner/v1/rl/agent_loop_manager/produce_utils.py index ef9151bbc..6639256c0 100644 --- a/xtuner/v1/rl/agent_loop_manager/produce_utils.py +++ b/xtuner/v1/rl/agent_loop_manager/produce_utils.py @@ -16,11 +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_and_discard_rollout_groups from xtuner.v1.rl.utils import ( AGENT_LOOP_PAUSE_REQUEST_TIMEOUT_S, PRODUCER_PAUSE_PENDING_TASK_TIMEOUT_S, @@ -197,8 +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)) - for item in group: - 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 a4c8f45e3..61a2abacc 100644 --- a/xtuner/v1/rl/replay_buffer.py +++ b/xtuner/v1/rl/replay_buffer.py @@ -14,12 +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_and_discard_rollout_groups from xtuner.v1.rl.utils import ( BetweenNode, ConditionNode, @@ -488,11 +488,22 @@ 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_non_retryable_expired_groups(groups: list[list[RolloutState]]) -> None: + """Release non-retryable trace sessions in one RPC, then discard + groups.""" + if not groups: + return + + await release_and_discard_rollout_groups(groups) + for group in groups: + for item in group: + item.status = Status.EXPIRED + async def put( self, items: list[RolloutState], @@ -518,6 +529,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_non_retryable_expired_groups([items]) return storage_item = StorageItem( item=items, @@ -558,6 +570,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 {} + non_retryable_expired_groups: list[list[RolloutState]] = [] async with self._lock: updated_records: list[StorageItem] = [] deleted_uids: list[int] = [] @@ -583,12 +596,14 @@ async def refresh_staleness( if status == Status.EXPIRED: expired_count += 1 if not retryable: + 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_non_retryable_expired_groups(non_retryable_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..045f094c9 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 @@ -335,11 +336,26 @@ 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): - self.release(session_id) - self.sessions.clear() + self.release_sessions(list(self.sessions)) self.objects.clear() self.updated_at.clear() @@ -462,6 +478,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 +488,43 @@ 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. + """ + 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 set() + + 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 ===") diff --git a/xtuner/v1/train/rl_trainer.py b/xtuner/v1/train/rl_trainer.py index 1730d68ab..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,9 +893,27 @@ 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) -> 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 set() + + released_session_ids = ray.get(store.release_sessions.remote(session_ids)) + self.logger.info( + "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() @@ -892,11 +923,15 @@ def _release_trace_store(self) -> None: 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. + # 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, train_batch: list[list[RolloutState]], @@ -949,7 +984,7 @@ def _train_one_batch( rollout_idx=train_step, ) - self._release_trace_store() + self._release_trace_sessions_after_train_batch(train_batch) return { "data_info": data_info, @@ -957,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, @@ -976,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 @@ -1850,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,