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
183 changes: 183 additions & 0 deletions tests/experimental/orchestrator/async_rl_program_test.py
Original file line number Diff line number Diff line change
@@ -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()
151 changes: 137 additions & 14 deletions tests/experimental/orchestrator/distributed_rl_engine_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -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})

Expand Down Expand Up @@ -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",
Expand All @@ -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"
Expand All @@ -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()
Loading
Loading