diff --git a/tests/experimental/orchestrator/rl_program_test.py b/tests/experimental/orchestrator/rl_program_test.py index 8114f7860..2432b82f5 100644 --- a/tests/experimental/orchestrator/rl_program_test.py +++ b/tests/experimental/orchestrator/rl_program_test.py @@ -56,6 +56,46 @@ def _create_rollout_response( ) +def _make_trajectory_group( + prompt_id: str = "prompt_0", + group_id: str = "group_0", + group_size: int = 2, + reward: float = 1.0, +) -> list[datatypes.TrajectoryItem]: + return [ + distributed_rl_engine._response_to_trajectory_item( + _create_rollout_response( + f"req_{prompt_id}_{idx}", + prompt_id, + group_id, + pair_index=idx, + reward=reward, + ) + ) + for idx in range(group_size) + ] + + +def _set_mock_poll_batches( + mock_engine: mock.MagicMock, + *batches: Sequence[datatypes.TrajectoryItem], +) -> None: + call_idx = 0 + batch_list = list(batches) + + async def _mock_poll(timeout_s=0.1): + del timeout_s + nonlocal call_idx + if call_idx < len(batch_list): + res = list(batch_list[call_idx]) + call_idx += 1 + return res + await asyncio.sleep(0.01) + return [] + + mock_engine.poll_rollouts.side_effect = _mock_poll + + class RLProgramTest(absltest.TestCase): def setUp(self): @@ -95,44 +135,6 @@ async def _mock_poll(*args, **kwargs): max_packed_len=16 ) - def _make_trajectory_group( - self, - prompt_id: str = "prompt_0", - group_id: str = "group_0", - group_size: int = 2, - reward: float = 1.0, - ) -> list[datatypes.TrajectoryItem]: - return [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - f"req_{prompt_id}_{idx}", - prompt_id, - group_id, - pair_index=idx, - reward=reward, - ) - ) - for idx in range(group_size) - ] - - def _set_mock_poll_batches( - self, *batches: Sequence[datatypes.TrajectoryItem] - ) -> None: - call_idx = 0 - batch_list = list(batches) - - async def _mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(batch_list): - res = list(batch_list[call_idx]) - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = _mock_poll - def _create_program( self, dataset: Any = ("prompt_0",), @@ -162,34 +164,9 @@ def test_initialization(self): def test_run_async_four_stages_with_long_polling(self): async def _run(): - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "prompt_0", "group_0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "prompt_0", "group_0", pair_index=1 - ) - ), - ], - [], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll + _set_mock_poll_batches( + self.mock_engine, _make_trajectory_group(), [] + ) begin_steps = [] end_steps = [] @@ -200,11 +177,8 @@ def on_begin(step): def on_end(step, result): end_steps.append((step, result)) - program = rl_program.StandardRLProgram( + program = self._create_program( dataset=["prompt_data_0"], - algo=self.mock_algo, - reward_fns=[lambda x: 1.0], - assembler=self.assembler, on_step_begin=on_begin, on_step_end=on_end, ) @@ -234,40 +208,8 @@ def on_end(step, result): def test_step_can_skip_weight_sync(self): async def _run(): - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "prompt_0", "group_0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "prompt_0", "group_0", pair_index=1 - ) - ), - ], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll - - program = rl_program.StandardRLProgram( - algo=self.mock_algo, - reward_fns=[lambda x: 1.0], - assembler=self.assembler, - sync_weights=False, - ) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) + program = self._create_program(sync_weights=False) await program.run_async( self.mock_engine, train_dataset=["override_prompt"], num_steps=1 @@ -391,41 +333,9 @@ async def _run(): asyncio.run(_run()) - def test_run_synchronous_entry_point(self): - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "prompt_0", "group_0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "prompt_0", "group_0", pair_index=1 - ) - ), - ], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll - - program = rl_program.StandardRLProgram( - algo=self.mock_algo, - reward_fns=[lambda x: 2.0], - assembler=self.assembler, - ) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) + program = self._create_program(reward_fns=[lambda x: 2.0]) program.run( self.mock_engine, train_dataset=["sync_prompt"], num_steps=1 @@ -438,39 +348,8 @@ async def mock_poll(timeout_s=0.1): def test_run_with_existing_running_loop(self): async def _run(): - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "prompt_0", "group_0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "prompt_0", "group_0", pair_index=1 - ) - ), - ], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll - - program = rl_program.StandardRLProgram( - algo=self.mock_algo, - reward_fns=[lambda x: 1.0], - assembler=self.assembler, - ) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) + program = self._create_program() program.run( self.mock_engine, train_dataset=["async_prompt"], num_steps=1 @@ -495,39 +374,11 @@ async def _run(): def test_prompt_dictionary_id_and_group_extraction(self): async def _run(): - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "custom_p0", "custom_g0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "custom_p0", "custom_g0", pair_index=1 - ) - ), - ], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll - - program = rl_program.StandardRLProgram( - algo=self.mock_algo, - reward_fns=[lambda x: 1.0], - assembler=self.assembler, + _set_mock_poll_batches( + self.mock_engine, + _make_trajectory_group(prompt_id="custom_p0", group_id="custom_g0"), ) + program = self._create_program() dict_item = { "prompt_id": "custom_p0", @@ -549,51 +400,12 @@ async def mock_poll(timeout_s=0.1): def test_multi_group_mini_batch_gradient_accumulation(self): async def _run(): self.mock_algo.mini_batch_size = 2 - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "prompt_0", "group_0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "prompt_0", "group_0", pair_index=1 - ) - ), - ], - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_1_0", "prompt_1", "group_1", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_1_1", "prompt_1", "group_1", pair_index=1 - ) - ), - ], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll - - program = rl_program.StandardRLProgram( - algo=self.mock_algo, - reward_fns=[lambda x: 1.0], - assembler=self.assembler, + _set_mock_poll_batches( + self.mock_engine, + _make_trajectory_group("prompt_0", "group_0"), + _make_trajectory_group("prompt_1", "group_1"), ) + program = self._create_program() await program.run_async( self.mock_engine, train_dataset=["p0", "p1"], num_steps=1 @@ -629,39 +441,8 @@ async def _run(): return_value=np.array([[-0.1, -0.2]], dtype=np.float32) ) - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "prompt_0", "group_0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "prompt_0", "group_0", pair_index=1 - ) - ), - ], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll - - program = rl_program.StandardRLProgram( - algo=self.mock_algo, - reward_fns=[lambda x: 1.0], - assembler=self.assembler, - ) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) + program = self._create_program() await program.run_async( self.mock_engine, train_dataset=["prompt_0"], num_steps=1 @@ -679,40 +460,8 @@ async def _run(): self.mock_algo.requires_reference_kl = True # Returning a raw dict instead of TrainExample self.assembler.pack = mock.MagicMock(return_value=[{"raw": "batch"}]) - - poll_results = [ - [ - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_0", "prompt_0", "group_0", pair_index=0 - ) - ), - distributed_rl_engine._response_to_trajectory_item( - _create_rollout_response( - "req_0_1", "prompt_0", "group_0", pair_index=1 - ) - ), - ], - ] - call_idx = 0 - - async def mock_poll(timeout_s=0.1): - del timeout_s - nonlocal call_idx - if call_idx < len(poll_results): - res = poll_results[call_idx] - call_idx += 1 - return res - await asyncio.sleep(0.01) - return [] - - self.mock_engine.poll_rollouts.side_effect = mock_poll - - program = rl_program.StandardRLProgram( - algo=self.mock_algo, - reward_fns=[lambda x: 1.0], - assembler=self.assembler, - ) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) + program = self._create_program() with self.assertRaises(TypeError) as cm: await program.run_async( @@ -724,7 +473,7 @@ async def mock_poll(timeout_s=0.1): def test_run_async_handles_early_dispatch_completion(self): async def _run(): - self._set_mock_poll_batches(self._make_trajectory_group()) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) program = self._create_program() await program.run_async(self.mock_engine, num_steps=1) self.assertEqual(program.step, 1) @@ -733,7 +482,7 @@ async def _run(): def test_run_async_propagates_train_stage_exception(self): async def _run(): - self._set_mock_poll_batches(self._make_trajectory_group()) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) self.mock_engine.train_step.side_effect = RuntimeError( "Training worker OOM" ) @@ -747,7 +496,7 @@ async def _run(): def test_run_async_propagates_critique_stage_exception(self): async def _run(): - self._set_mock_poll_batches(self._make_trajectory_group()) + _set_mock_poll_batches(self.mock_engine, _make_trajectory_group()) def failing_reward_fn(_): raise ValueError("Reward model computation failed") @@ -761,7 +510,7 @@ def failing_reward_fn(_): def test_run_async_cancels_background_stages_on_external_cancellation(self): async def _run(): - self._set_mock_poll_batches() # Yields empty and sleeps + _set_mock_poll_batches(self.mock_engine) # Yields empty and sleeps program = self._create_program() task = asyncio.create_task(