From 3a805732ccb085483711b7abddedf7cfdbaffe96 Mon Sep 17 00:00:00 2001 From: Shadi Noghabi Date: Fri, 21 Aug 2026 14:30:10 -0700 Subject: [PATCH] Refactor rollout dispatch in DistributedRLEngine with group_size expansion and typed GenerationArgs. PiperOrigin-RevId: 968693092 --- .../orchestrator/async_rl_program_test.py | 183 +++++++++++ .../distributed_rl_engine_test.py | 151 ++++++++- .../orchestrator/async_rl_program.py | 290 ++++++++++++++++++ .../orchestrator/distributed_rl_engine.py | 97 +++--- .../orchestrator/rl_engine_interface.py | 18 +- 5 files changed, 687 insertions(+), 52 deletions(-) create mode 100644 tests/experimental/orchestrator/async_rl_program_test.py create mode 100644 tunix/experimental/orchestrator/async_rl_program.py diff --git a/tests/experimental/orchestrator/async_rl_program_test.py b/tests/experimental/orchestrator/async_rl_program_test.py new file mode 100644 index 000000000..804bb3669 --- /dev/null +++ b/tests/experimental/orchestrator/async_rl_program_test.py @@ -0,0 +1,183 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Tests for AsyncRLProgram and StandardRLProgram.""" + +import asyncio +from unittest import mock + +from absl.testing import absltest +import numpy as np +from tunix.experimental.common import datatypes +from tunix.experimental.orchestrator import algorithm_adapter +from tunix.experimental.orchestrator import async_rl_program +from tunix.experimental.orchestrator import batch_assembly +from tunix.experimental.orchestrator import distributed_rl_engine + + +def _create_rollout_response( + request_id: str, + prompt_id: str, + group_id: str, + pair_index: int = 0, + policy_version: int = 0, + reward: float = 1.0, +) -> datatypes.RolloutResponse: + return datatypes.RolloutResponse( + request_id=request_id, + prompt_id=prompt_id, + status="COMPLETED", + env_reward=reward, + policy_version=policy_version, + prompt_tokens=np.array([1, 2], dtype=np.int32), + segments=[ + datatypes.TokenSegment( + source="assistant", + tokens=np.array([3, 4], dtype=np.int32), + loss_mask=np.array([1, 1], dtype=np.int32), + ) + ], + metadata={ + "group_id": group_id, + "pair_index": pair_index, + }, + ) + + +class AsyncRLProgramTest(absltest.TestCase): + + def setUp(self): + super().setUp() + self.mock_engine = mock.MagicMock(spec=distributed_rl_engine.DistributedRLEngine) + self.mock_engine.dispatch_rollouts = mock.AsyncMock() + self.mock_engine.train_step = mock.AsyncMock(return_value="step_done") + async def _mock_poll(*args, **kwargs): + await asyncio.sleep(0.01) + return [] + + self.mock_engine.sync_weights = mock.AsyncMock(return_value=1) + self.mock_engine.poll_rollouts = mock.AsyncMock(side_effect=_mock_poll) + self.mock_algo = mock.MagicMock(spec=algorithm_adapter.AlgorithmAdapter) + self.mock_algo.group_size = 2 + self.mock_algo.mini_batch_size = 1 + self.mock_algo.max_turns = 1 + self.mock_algo.max_packed_len = 16 + self.mock_algo.requires_reference_kl = False + + mock_payload = datatypes.RLTrainerPayload( + token_ids=np.array([1, 2, 3, 4], dtype=np.int32), + token_mask=np.array([0, 0, 1, 1], dtype=np.float32), + loss_mask=np.array([0, 0, 1, 1], dtype=np.float32), + advantages=np.full(4, 1.0, dtype=np.float32), + action_mask=np.array([0, 0, 1, 1], dtype=np.float32), + ) + self.mock_algo.create_trainer_payloads.return_value = [mock_payload, mock_payload] + self.assembler = batch_assembly.SequencePackedBatchAssembler(max_packed_len=16) + + def test_initialization(self): + program = async_rl_program.StandardRLProgram( + dataset=["prompt_1"], + algo=self.mock_algo, + reward_fns=[lambda x: 1.0], + assembler=self.assembler, + ) + self.assertEqual(program.step, 0) + self.assertEqual(program.group_size, 2) + self.assertEqual(program.mini_batch_size, 1) + self.assertIsNotNone(program.raw_q) + self.assertIsNotNone(program.scored_q) + + 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): + 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 + + begin_steps = [] + end_steps = [] + + def on_begin(step): + begin_steps.append(step) + + def on_end(step, result): + end_steps.append((step, result)) + + program = async_rl_program.StandardRLProgram( + 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, + ) + + await program.run_async(self.mock_engine, num_steps=1) + + self.assertEqual(program.step, 1) + self.assertEqual(begin_steps, [0]) + self.assertEqual(end_steps, [(1, "step_done")]) + self.mock_engine.dispatch_rollouts.assert_called_once_with( + [{"prompt": "prompt_data_0", "prompt_id": "prompt_0"}], + group_size=2, + policy_version=0, + ) + self.mock_engine.train_step.assert_called_once() + self.mock_engine.sync_weights.assert_called_once_with( + role=datatypes.Role.ACTOR + ) + + asyncio.run(_run()) + + def test_stage_exception_aborts_queue_and_propagates(self): + class FailingProgram(async_rl_program.StandardRLProgram): + + async def rollout_dispatch_stage(self, engine): + del engine + raise RuntimeError("Rollout worker cluster down!") + + async def _run(): + prog = FailingProgram( + dataset=["prompt"], + algo=self.mock_algo, + assembler=self.assembler, + ) + with self.assertRaises(RuntimeError) as cm: + await prog.run_async(self.mock_engine, num_steps=1) + self.assertIn("Rollout worker cluster down!", str(cm.exception)) + + asyncio.run(_run()) + + +if __name__ == "__main__": + absltest.main() diff --git a/tests/experimental/orchestrator/distributed_rl_engine_test.py b/tests/experimental/orchestrator/distributed_rl_engine_test.py index 8fb59bfa6..71b1a5c55 100644 --- a/tests/experimental/orchestrator/distributed_rl_engine_test.py +++ b/tests/experimental/orchestrator/distributed_rl_engine_test.py @@ -73,14 +73,18 @@ def setUp(self): def test_generate_load_balances_across_rollout_workers(self): async def _run(): - resp1 = datatypes.RolloutResponse(request_id="r1", status="COMPLETED", env_reward=1.0) - resp2 = datatypes.RolloutResponse(request_id="r2", status="COMPLETED", env_reward=2.0) + resp1 = datatypes.RolloutResponse( + request_id="r1", status="COMPLETED", env_reward=1.0 + ) + resp2 = datatypes.RolloutResponse( + request_id="r2", status="COMPLETED", env_reward=2.0 + ) self.mock_rollout_1.generate.return_value = [resp1] self.mock_rollout_2.generate.return_value = [resp2] results = await self.engine.generate(["p1", "p2"]) - self.assertEqual(len(results), 2) + self.assertLen(results, 2) rewards = {res.traj.reward for res in results} self.assertEqual(rewards, {1.0, 2.0}) @@ -258,8 +262,7 @@ async def _run(): asyncio.run(_run()) - def test_balancer_prefix_routing(self): - + def test_dispatch_rollout_requests_with_prefix_routing(self): async def _run(): req1 = datatypes.RolloutRequest( request_id="1", @@ -274,9 +277,10 @@ async def _run(): metadata={"prefix_hash": 1}, ) - await self.engine.dispatch_rollouts([req1, req2]) + req_ids = await self.engine.dispatch_rollout_requests([req1, req2]) + self.assertEqual(req_ids, ["1", "2"]) - # Due to deterministic round-robin / hash logic, req1 goes to rollout_1 and req2 goes to rollout_2 + # Due to deterministic hash logic, req1 -> rollout_1 and req2 -> rollout_2 self.mock_rollout_1.generate.assert_called_once() dispatched_req1 = self.mock_rollout_1.generate.call_args.kwargs[ "requests" @@ -291,18 +295,137 @@ async def _run(): asyncio.run(_run()) - def test_dispatch_rollouts_requires_strict_kwargs(self): + def test_dispatch_rollouts_delegates_to_dispatch_rollout_requests(self): async def _run(): - with self.assertRaisesRegex(ValueError, "prompt_ids' must be provided"): - await self.engine.dispatch_rollouts(["p1", "p2"], policy_version=0) + req1 = datatypes.RolloutRequest( + request_id="1", + prompt="p1", + prompt_id="1", + metadata={"prefix_hash": 0}, + ) + req2 = datatypes.RolloutRequest( + request_id="2", + prompt="p2", + prompt_id="2", + metadata={"prefix_hash": 1}, + ) + + req_ids = await self.engine.dispatch_rollouts([req1, req2]) + self.assertEqual(req_ids, ["1", "2"]) + + self.mock_rollout_1.generate.assert_called_once() + self.mock_rollout_2.generate.assert_called_once() - with self.assertRaisesRegex(ValueError, "match the length of prompts"): - await self.engine.dispatch_rollouts(["p1", "p2"], prompt_ids=["id1"], policy_version=0) + asyncio.run(_run()) - with self.assertRaisesRegex(ValueError, "policy_version' must be provided"): - await self.engine.dispatch_rollouts(["p1", "p2"], prompt_ids=["id1", "id2"]) + def test_dispatch_rollouts_expands_group_size(self): + async def _run(): + req_ids = await self.engine.dispatch_rollouts( + [ + {"prompt": "p1", "prompt_id": "p1"}, + {"prompt": "p2", "prompt_id": "p2"}, + ], + group_size=3, + policy_version=5, + ) + self.assertLen(req_ids, 6) + + # 2 calls to generate (1 per worker in pool via prefix hash / round-robin) + total_dispatched = 0 + for mock_w in (self.mock_rollout_1, self.mock_rollout_2): + for call in mock_w.generate.call_args_list: + reqs = call.kwargs["requests"] + total_dispatched += len(reqs) + for r in reqs: + self.assertEqual(r.target_policy_version, 5) + self.assertIn(r.metadata["pair_index"], (0, 1, 2)) + self.assertEqual(total_dispatched, 6) asyncio.run(_run()) + def test_dispatch_rollouts_auto_extracts_prompt_and_group_ids(self): + async def _run(): + dict_item = { + "prompt": "Solve math", + "prompt_id": "math_1", + "group_id": "grp_1", + } + req_ids = await self.engine.dispatch_rollouts( + [dict_item], group_size=2, policy_version=1 + ) + self.assertLen(req_ids, 2) + + all_dispatched = [] + for mock_w in (self.mock_rollout_1, self.mock_rollout_2): + for c in mock_w.generate.call_args_list: + all_dispatched.extend(c.kwargs["requests"]) + + self.assertLen(all_dispatched, 2) + pair_indices = {r.metadata["pair_index"] for r in all_dispatched} + self.assertEqual(pair_indices, {0, 1}) + self.assertTrue(all(r.prompt_id == "math_1" for r in all_dispatched)) + self.assertTrue( + all(r.metadata["group_id"] == "grp_1" for r in all_dispatched) + ) + + asyncio.run(_run()) + + def test_dispatch_rollouts_passes_generation_args_and_route_metadata(self): + async def _run(): + gen_args = datatypes.GenerationArgs( + temperature=0.7, max_generation_steps=128 + ) + req_ids = await self.engine.dispatch_rollouts( + [{"prompt": "p1", "prompt_id": "p1"}], + group_size=1, + generation_args=gen_args, + route_metadata={"prefix_hash": "cache_key_1"}, + ) + self.assertLen(req_ids, 1) + + mock_call = ( + self.mock_rollout_1.generate.call_args + or self.mock_rollout_2.generate.call_args + ) + dispatched = mock_call.kwargs["requests"][0] + self.assertEqual( + dispatched.generation_kwargs, + {"temperature": 0.7, "max_generation_steps": 128}, + ) + self.assertEqual(dispatched.metadata["prefix_hash"], "cache_key_1") + + asyncio.run(_run()) + + def test_dispatch_rollouts_generates_deterministic_request_ids(self): + async def _run(): + req_ids = await self.engine.dispatch_rollouts( + [{"prompt": "Hello", "prompt_id": "p_123"}], + group_size=2, + policy_version=3, + ) + self.assertEqual(req_ids, ["req_p_123_0_v3", "req_p_123_1_v3"]) + + asyncio.run(_run()) + + def test_dispatch_rollouts_handles_none_metadata(self): + async def _run(): + req_ids = await self.engine.dispatch_rollouts( + [{"prompt": "p1", "prompt_id": "p1"}], + group_size=1, + metadata=None, + route_metadata=None, + ) + self.assertLen(req_ids, 1) + + asyncio.run(_run()) + + def test_dispatch_rollouts_raises_without_prompt_id(self): + async def _run(): + with self.assertRaisesRegex(ValueError, "lacks 'prompt_id'"): + await self.engine.dispatch_rollouts(["raw_prompt_without_id"]) + + asyncio.run(_run()) + + if __name__ == "__main__": absltest.main() diff --git a/tunix/experimental/orchestrator/async_rl_program.py b/tunix/experimental/orchestrator/async_rl_program.py new file mode 100644 index 000000000..fcc3b4cef --- /dev/null +++ b/tunix/experimental/orchestrator/async_rl_program.py @@ -0,0 +1,290 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Layer 3: Workflow & Program (async_rl_program.py) following Orchestrator V2. + +Contains: +- AsyncRLProgram: Base class for multi-stage concurrent DAG workflows. +- StandardRLProgram: Single standard program handling 95% of use cases with + long-polling rollout collector and streaming gradient accumulation. +""" + +import asyncio +from collections.abc import Callable, Iterable, Sequence +from typing import Any + +from absl import logging + +from tunix.experimental.common import datatypes +from tunix.experimental.orchestrator import algorithm_adapter +from tunix.experimental.orchestrator import batch_assembly +from tunix.experimental.orchestrator import rl_engine_interface +from tunix.experimental.queue_manager import trajectory_queue_manager +from tunix.rl import common as rl_common + +# _response_to_trajectory_item has been moved to distributed_rl_engine.py + + +class AsyncRLProgram: + """Base class for asynchronous multi-stage DAG workflows.""" + + def __init__(self): + self._is_running = False + self.policy_version = 0 + + @property + def step(self) -> int: + return self.policy_version + + +class StandardRLProgram(AsyncRLProgram): + """Single standard program handling 95% of use cases with long-polling rollouts. + + Runs 4 concurrent stages: + 1. Rollout dispatch stage: Fire-and-forget requests across worker pool. + 2. Polling stage: Long-polls completed rollout responses into grouping queue. + 3. Critique stage: Scores rewards, PRMs, and reference KL logprobs. + 4. Train stage: Streaming gradient accumulation over microbatches. + """ + + def __init__( + self, + dataset: Iterable[Any], + algo: algorithm_adapter.AlgorithmAdapter, + reward_fns: Sequence[Callable[..., Any]] | None = None, + assembler: batch_assembly.BatchAssembler | None = None, + group_size: int = 8, + mini_batch_size: int = 4, + max_staleness: int | None = None, + on_step_begin: Callable[[int], None] | None = None, + on_step_end: Callable[[int, Any], None] | None = None, + ): + super().__init__() + self.dataset = dataset + self.algo = algo + self.reward_fns = list(reward_fns) if reward_fns else [] + self.group_size = getattr(algo, "group_size", group_size) + self.mini_batch_size = getattr(algo, "mini_batch_size", mini_batch_size) + self.assembler = assembler or batch_assembly.SequencePackedBatchAssembler( + max_packed_len=getattr(algo, "max_packed_len", 8192) + ) + self.on_step_begin = on_step_begin + self.on_step_end = on_step_end + + self.raw_q = trajectory_queue_manager.TrajectoryQueueManager.create( + group_size=self.group_size, + max_staleness=max_staleness, + current_policy_version=lambda: self.policy_version, + ) + self.scored_q = trajectory_queue_manager.TrajectoryQueueManager.create( + group_size=self.group_size + ) + + async def rollout_dispatch_stage( + self, engine: rl_engine_interface.AbstractRLEngine + ) -> None: + """Stage 1A: Dispatches rollout requests across workers asynchronously.""" + for prompt_idx, prompt_item in enumerate(self.dataset): + if isinstance(prompt_item, dict): + prompt_item = dict(prompt_item) + prompt_item.setdefault("prompt_id", f"prompt_{prompt_idx}") + elif not hasattr(prompt_item, "prompt_id"): + prompt_item = { + "prompt": prompt_item, + "prompt_id": f"prompt_{prompt_idx}", + } + + await engine.dispatch_rollouts( + [prompt_item], + group_size=self.group_size, + policy_version=self.policy_version, + ) + + async def polling_stage( + self, engine: rl_engine_interface.AbstractRLEngine + ) -> None: + """Stage 1B: Long-polls completed worker rollout responses into the queue.""" + while True: + try: + completed = await engine.poll_rollouts(timeout_s=0.1) + if isinstance(completed, list) and completed: + for item in completed: + await self.raw_q.put(item) + + except asyncio.CancelledError: + break + except Exception as exc: # pylint: disable=broad-exception-caught + logging.warning("Error in polling_stage: %s", exc) + await asyncio.sleep(0.01) + + async def critique_stage( + self, engine: rl_engine_interface.AbstractRLEngine + ) -> None: + """Stage 2: Scores rewards, PRMs, and reference KL logprobs.""" + while True: + try: + group = await self.raw_q.get_group() + except asyncio.CancelledError: + break + except Exception: + break + + rewards = [] + for item in group: + if self.reward_fns: + r = sum(fn(item) for fn in self.reward_fns) + else: + r = getattr(item.traj, "reward", 0.0) + rewards.append(float(r)) + + trainer_payloads = self.algo.create_trainer_payloads( + group, rewards=rewards + ) + for idx, payload in enumerate(trainer_payloads): + adv = payload.advantages + reward_val = ( + float(adv[0]) # pyrefly: ignore[bad-index] + if hasattr(adv, "__len__") and len(adv) > 0 # pyrefly: ignore[bad-argument-type] + else float(adv) # pyrefly: ignore[bad-argument-type] + ) + item = datatypes.TrajectoryItem( + pair_index=idx, + group_id=getattr(group[0], "group_id", "default"), + start_step=0, + traj=datatypes.Trajectory(reward=reward_val), + # TODO: Stream RLTrainerPayload directly instead of re-wrapping in TrajectoryItem. + ) + item.payload = payload # pyrefly: ignore[missing-attribute] + await self.scored_q.put(item) + + async def train_stage( + self, engine: rl_engine_interface.AbstractRLEngine, num_steps: int | None = None + ) -> None: + """Stage 3: Streaming gradient accumulation with RLTrainerPayloads.""" + step = 0 + while num_steps is None or step < num_steps: + if self.on_step_begin: + self.on_step_begin(self.step) + + uncommitted_groups = [] + step_result = None + + for group_idx in range(self.mini_batch_size): + scored_items = await self.scored_q.get_batch(num_groups=1) + if not scored_items: + break + uncommitted_groups.append(scored_items) + + payloads = [getattr(item, "payload", None) for item in scored_items] + # TODO: Implement streaming microbatch assembly to overlap packing with trainer execution. + microbatches = self.assembler.pack(payloads) # pyrefly: ignore[bad-argument-type] + if getattr(self.algo, "requires_reference_kl", False): + scored_microbatches = [] + for batch in microbatches: + if not isinstance(batch, rl_common.TrainExample): + raise TypeError( + "Reference KL requires an assembler that returns " + "rl_common.TrainExample microbatches; got " + f"{type(batch).__name__}." + ) + ref_logps = await engine.per_token_logps( + datatypes.Role.REFERENCE, items=batch + ) + scored_microbatches.append( + batch_assembly.with_ref_per_token_logps(batch, ref_logps) + ) + microbatches = scored_microbatches + + is_final = group_idx == self.mini_batch_size - 1 + for batch in microbatches: + step_result = await engine.train_step( + batch, + role=datatypes.Role.ACTOR, + accumulate_gradients=True, + apply_optimizer=is_final, + ) + + new_version = await engine.sync_weights(role=datatypes.Role.ACTOR) + self.policy_version = new_version if new_version else self.step + 1 + self.scored_q.commit(step, groups=uncommitted_groups) + + if self.on_step_end: + self.on_step_end(self.step, step_result) + step += 1 + + async def run_async( + self, + engine: rl_engine_interface.AbstractRLEngine, + num_steps: int | None = None, + **kwargs: Any, + ) -> None: + """Launches all stages concurrently on event loop.""" + del kwargs + logging.info("Starting StandardRLProgram concurrent stages...") + + train_task = asyncio.create_task(self.train_stage(engine, num_steps)) + tasks = [ + asyncio.create_task(self.rollout_dispatch_stage(engine)), + asyncio.create_task(self.polling_stage(engine)), + asyncio.create_task(self.critique_stage(engine)), + train_task, + ] + + try: + while not train_task.done(): + done, _ = await asyncio.wait( + tasks, return_when=asyncio.FIRST_COMPLETED, timeout=0.05 + ) + for task in done: + if task.exception(): + raise task.exception() # pyrefly: ignore[bad-raise] + if train_task.exception(): + raise train_task.exception() # pyrefly: ignore[bad-raise] + except Exception as exc: + logging.error("Exception in StandardRLProgram execution: %s", exc) + await self.raw_q.abort(exc) + await self.scored_q.abort(exc) + raise + finally: + for task in tasks: + if not task.done(): + task.cancel() + + def run( + self, + engine: rl_engine_interface.AbstractRLEngine, + num_steps: int | None = None, + **kwargs: Any, + ) -> None: + """Synchronous entry point running all stages on an event loop.""" + try: + loop = asyncio.get_running_loop() + except RuntimeError: + loop = None + + def _retrieve_task_exception(t: asyncio.Task[Any]) -> None: + try: + t.result() + except Exception: # pylint: disable=broad-except + # Exception is already logged inside run_async, we just need to + # retrieve it so asyncio doesn't complain about unretrieved exceptions. + pass + + if loop and loop.is_running(): + self._bg_task = asyncio.create_task( + self.run_async(engine, num_steps, **kwargs) + ) + self._bg_task.add_done_callback(_retrieve_task_exception) + else: + asyncio.run(self.run_async(engine, num_steps, **kwargs)) diff --git a/tunix/experimental/orchestrator/distributed_rl_engine.py b/tunix/experimental/orchestrator/distributed_rl_engine.py index e7367c25c..0c8c10a15 100644 --- a/tunix/experimental/orchestrator/distributed_rl_engine.py +++ b/tunix/experimental/orchestrator/distributed_rl_engine.py @@ -127,55 +127,80 @@ async def _invoke_worker( return await res return res + async def dispatch_rollout_requests( + self, + requests: Sequence[datatypes.RolloutRequest], + ) -> list[str]: + """Dispatches pre-formed RolloutRequests across rollout workers using prefix routing.""" + for req in requests: + route_key = (req.metadata or {}).get("prefix_hash", req.prompt_id) + worker = self._rollout_pool._get_next_actor( + kwargs={"route_key": route_key} + ) + res = worker.dispatch_task(method_name="generate", requests=[req]) + if inspect.isawaitable(res): + await res + + return [r.request_id for r in requests] + async def dispatch_rollouts( - self, prompts: Sequence[Any], **kwargs: Any + self, + prompts: Sequence[Any], + *, + group_size: int = 1, + policy_version: int = 0, + generation_args: datatypes.GenerationArgs | None = None, + route_metadata: Mapping[str, Any] | None = None, + **kwargs: Any, ) -> list[str]: """Dispatches rollout requests across workers, constructing RolloutRequests internally if needed.""" + base_metadata = { + **(route_metadata or {}), + **(kwargs.get("metadata") or {}), + } + gen_kwargs = generation_args.as_kwargs() if generation_args else {} + version = kwargs.get("policy_version", policy_version) + rollout_reqs: list[datatypes.RolloutRequest] = [] for idx, p in enumerate(prompts): - # TODO: why do we support sending rollout requests directly? shouldn't this be the engine resposibility? if isinstance(p, datatypes.RolloutRequest): rollout_reqs.append(p) - else: - prompt_ids = kwargs.get("prompt_ids") - if not prompt_ids or len(prompt_ids) != len(prompts): - raise ValueError( - "When passing raw prompts, 'prompt_ids' must be provided in" - " kwargs and match the length of prompts." - ) - if "policy_version" not in kwargs: - raise ValueError( - "When passing raw prompts, 'policy_version' must be provided in" - " kwargs." - ) - # TODO: should we autogenerate request_id? - req_id = ( - kwargs.get("request_id") or f"req_{idx}_{uuid.uuid4().hex[:8]}" + continue + + prompt_id = getattr(p, "prompt_id", None) or ( + p.get("prompt_id") if isinstance(p, dict) else None + ) + if not prompt_id: + raise ValueError( + f"Prompt at index {idx} lacks 'prompt_id'. Every prompt item " + "must provide a 'prompt_id' attribute or dict key." ) + + prompt_id = str(prompt_id) + group_id = str( + getattr(p, "group_id", None) + or (p.get("group_id") if isinstance(p, dict) else None) + or prompt_id + ) + raw_prompt = p.get("prompt", p) if isinstance(p, dict) else p + + for g_idx in range(group_size): rollout_reqs.append( datatypes.RolloutRequest( - request_id=req_id, - prompt=p, - prompt_id=prompt_ids[idx], - target_policy_version=kwargs["policy_version"], - metadata=dict(kwargs.get("metadata", {})), + request_id=f"req_{prompt_id}_{g_idx}_v{version}", + prompt=raw_prompt, + prompt_id=prompt_id, + target_policy_version=version, + generation_kwargs=gen_kwargs, + metadata={ + **base_metadata, + "group_id": group_id, + "pair_index": g_idx, + }, ) ) - for req in rollout_reqs: - metadata = req.metadata or {} - route_key = metadata.get("prefix_hash") - if route_key is None: - route_key = req.prompt_id - worker = self._rollout_pool._get_next_actor( - kwargs={"route_key": route_key} - ) - - res = worker.dispatch_task(method_name="generate", requests=[req]) - if inspect.isawaitable(res): - await res - - return [r.request_id for r in rollout_reqs] + return await self.dispatch_rollout_requests(rollout_reqs) async def poll_rollouts( self, timeout_s: float = remote_execution.LONG_POLL_TIMEOUT_S diff --git a/tunix/experimental/orchestrator/rl_engine_interface.py b/tunix/experimental/orchestrator/rl_engine_interface.py index 0fad1f49d..04d7e5931 100644 --- a/tunix/experimental/orchestrator/rl_engine_interface.py +++ b/tunix/experimental/orchestrator/rl_engine_interface.py @@ -23,10 +23,24 @@ class AbstractRLEngine(Protocol): """Stateless compute primitives for distributed worker meshes.""" + async def dispatch_rollout_requests( + self, + requests: Sequence[datatypes.RolloutRequest], + ) -> list[str]: + """Dispatches pre-formed RolloutRequests across rollout workers using prefix routing.""" + ... + async def dispatch_rollouts( - self, prompts: Sequence[Any], **kwargs: Any + self, + prompts: Sequence[Any], + *, + group_size: int = 1, + policy_version: int = 0, + generation_args: datatypes.GenerationArgs | None = None, + route_metadata: Mapping[str, Any] | None = None, + **kwargs: Any, ) -> list[str]: - """Dispatches rollout requests across workers (constructing RolloutRequests internally).""" + """Dispatches rollout requests across workers (expanding group_size and constructing RolloutRequests internally if needed).""" ... async def poll_rollouts(