Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
521 changes: 521 additions & 0 deletions tests/rl/test_qwen35_vl_moe_recover_e2e.py

Large diffs are not rendered by default.

23 changes: 15 additions & 8 deletions tests/rl/test_rl_colocate_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,7 @@
from xtuner.v1.rl.agent_loop_manager import AsyncProduceStrategyConfig, ProduceBatchResult
from xtuner.v1.rl.agent_loop_manager.agent_loop_manager import AgentLoopManager
from xtuner.v1.rl.agent_loop_manager.produce_utils import _TaskRunner
from xtuner.v1.rl.health_manager import RLHealthManager
from xtuner.v1.rl.replay_buffer import AsyncReplayBufferConfig, SerializedRayObjectRef
from xtuner.v1.train.rl_trainer import RLColocateTrainer, RLThroughputBenchmark

Expand Down Expand Up @@ -112,6 +113,7 @@ def tearDown(self):

def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_weights_interval: int = 1):
trainer = RLColocateTrainer.__new__(RLColocateTrainer)
trainer._rollout_config = SimpleNamespace(weight_transport_type='ipc')
trainer.logger = MagicMock()
trainer._total_train_steps = total_train_steps
trainer._cur_step = 0
Expand Down Expand Up @@ -153,13 +155,13 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_
)

trainer.rollout_controller = SimpleNamespace(
check_and_shutdown_inactive_workers=SimpleNamespace(
remote=MagicMock(return_value="rollout_inactive_workers_shutdown")
shutdown_inactive_workers=SimpleNamespace(
remote=MagicMock(side_effect=_ray_get_none_ref)
),
offload=SimpleNamespace(remote=MagicMock(return_value="rollout_offloaded")),
restart_inactive_workers=SimpleNamespace(remote=MagicMock(return_value="rollout_restarted")),
onload_weights=SimpleNamespace(remote=MagicMock(return_value="weights_loaded")),
onload_kvcache=SimpleNamespace(remote=MagicMock(return_value="kvcache_loaded")),
offload=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)),
restart_inactive_workers=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)),
onload_weights=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)),
onload_kvcache=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)),
validate_registered_workers_to_proxy=SimpleNamespace(remote=MagicMock(side_effect=_ray_get_none_ref)),
)
trainer.train_controller = SimpleNamespace(
Expand All @@ -179,6 +181,11 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_
]
),
)
trainer.rl_health_manager = RLHealthManager(
train_controller=trainer.train_controller,
rollout_controller=trainer.rollout_controller,
rollout_config=trainer._rollout_config,
)
return trainer

def test_fit_accepts_async_strategy_manager_on_colocate_path(self):
Expand Down Expand Up @@ -240,7 +247,7 @@ async def _produce_batch(batch_size, train_step, *, model_step):
return ProduceBatchResult(rollout_states=[[_FakeRolloutState(train_step)]])

trainer = self._make_trainer(SimpleNamespace(produce_batch=_produce_batch))
trainer.rollout_controller.check_and_shutdown_inactive_workers.remote.side_effect = RuntimeError(
trainer.rollout_controller.shutdown_inactive_workers.remote.side_effect = RuntimeError(
"inactive rollout workers after recovery"
)

Expand All @@ -251,7 +258,7 @@ async def _produce_batch(batch_size, train_step, *, model_step):
with self.assertRaisesRegex(RuntimeError, "inactive rollout workers"):
trainer.fit()

trainer.rollout_controller.check_and_shutdown_inactive_workers.remote.assert_called_once_with()
trainer.rollout_controller.shutdown_inactive_workers.remote.assert_called_once_with()
trainer.rollout_controller.offload.remote.assert_not_called()
trainer.train_controller.onload.assert_not_called()
trainer.train_controller.fit.assert_not_called()
Expand Down
2 changes: 1 addition & 1 deletion tests/rl/test_rl_disaggregated_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -148,7 +148,7 @@ def _make_trainer(self, agent_loop_manager):
weight_update=MagicMock(return_value="update"),
)
trainer.rollout_controller = SimpleNamespace(
check_and_shutdown_inactive_workers=SimpleNamespace(
shutdown_inactive_workers=SimpleNamespace(
remote=MagicMock(return_value="rollout_inactive_workers_shutdown")
),
restart_inactive_workers=SimpleNamespace(remote=MagicMock(return_value="rollout_restarted")),
Expand Down
3 changes: 2 additions & 1 deletion tests/rl/test_rl_trainer_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -102,11 +102,12 @@ def __init__(self):
self.continue_generation = _RemoteMethod(async_result=True)
self.flush_cache = _RemoteMethod(return_value="cache_flushed")
self.offload = _RemoteMethod(return_value="rollout_offloaded")
self.check_and_shutdown_inactive_workers = _RemoteMethod(return_value="rollout_inactive_workers_shutdown")
self.shutdown_inactive_workers = _RemoteMethod(return_value="rollout_inactive_workers_shutdown")
self.restart_inactive_workers = _RemoteMethod(return_value="rollout_restarted")
self.onload_weights = _RemoteMethod(return_value="weights_loaded")
self.onload_kvcache = _RemoteMethod(return_value="kvcache_loaded")
self.get_weight_update_targets = _RemoteMethod(return_value=())
self.mark_worker_groups_lifecycle_state = _RemoteMethod(return_value=None)
self.set_enable_partial_rollout = _RemoteMethod(return_value=None)
self.validate_registered_workers_to_proxy = _RemoteMethod(return_value=_AwaitableValue(None))

Expand Down
Loading
Loading