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
36 changes: 36 additions & 0 deletions tests/experimental/orchestrator/batch_assembly_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -82,6 +82,42 @@ def test_padded_batch_assembler(self):
self.assertEqual(payload.action_mask.shape, (2, 8))
self.assertEqual(payload.advantages.shape, (2, 8))

def test_grpo_train_example_assembler(self):
payload = datatypes.RLTrainerPayload(
token_ids=np.array([10, 11, 20, 21, 22], dtype=np.int32),
token_mask=np.ones(5, dtype=np.float32),
loss_mask=np.array([0, 0, 1, 1, 0], dtype=np.float32),
action_mask=np.array([0, 0, 1, 1, 0], dtype=np.float32),
advantages=np.array([0, 0, 2, 2, 2], dtype=np.float32),
prompt_ids=np.array([10, 11], dtype=np.int32),
prompt_mask=np.ones(2, dtype=np.float32),
completion_ids=np.array([20, 21, 22], dtype=np.int32),
completion_mask=np.array([1, 1, 0], dtype=np.float32),
)

assembler = batch_assembly.GRPOTrainExampleAssembler(
batch_size=2,
max_prompt_length=4,
max_response_length=5,
pad_id=0,
)
train_example = assembler.pack([payload])[0]

self.assertEqual(train_example.prompt_ids.shape, (2, 4))
self.assertEqual(train_example.completion_ids.shape, (2, 5))
np.testing.assert_array_equal(
train_example.prompt_ids[0], np.array([0, 0, 10, 11])
)
np.testing.assert_array_equal(
train_example.completion_ids[0], np.array([20, 21, 22, 0, 0])
)
np.testing.assert_array_equal(
train_example.completion_mask[0], np.array([1, 1, 0, 0, 0])
)
np.testing.assert_array_equal(
train_example.advantages[0], np.array([2, 2, 2, 0, 0])
)


if __name__ == "__main__":
absltest.main()
101 changes: 99 additions & 2 deletions tests/experimental/orchestrator/distributed_rl_engine_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ def __init__(self, *args, **kwargs):
self.poll_responses = mock.AsyncMock()
self.weight_sync = mock.AsyncMock()
self.fwd_bwd = mock.AsyncMock()
self.update = mock.AsyncMock()
self.prepare_weight_sync = mock.AsyncMock()
self.score = mock.AsyncMock()
self.per_token_logps = mock.AsyncMock()
Expand Down Expand Up @@ -90,6 +91,62 @@ async def _run():

asyncio.run(_run())

def test_generate_uses_explicit_generation_args(self):
async def _run():
resp = datatypes.RolloutResponse(
request_id="r1", status="COMPLETED", env_reward=1.0
)
self.mock_rollout_1.generate.return_value = [resp]
results = await self.engine.generate(
["p1"],
generation_args=datatypes.GenerationArgs(
max_generation_steps=8,
temperature=0.5,
return_logprobs=False,
),
)
self.assertLen(results, 1)
self.mock_rollout_1.generate.assert_called_once_with(
prompts=["p1"],
max_generation_steps=8,
temperature=0.5,
return_logprobs=False,
)
asyncio.run(_run())

def test_generate_rejects_legacy_generation_kwargs(self):
async def _run():
with self.assertRaisesRegex(TypeError, "GenerationArgs"):
await self.engine.generate(["p1"], temperature=0.5)
asyncio.run(_run())

def test_generate_routes_rollout_requests(self):
async def _run():
request = datatypes.RolloutRequest(
request_id="r1",
prompt="p1",
prompt_id="prompt_1",
generation_kwargs={"max_generation_steps": 8},
metadata={"prefix_hash": 0},
)
resp = datatypes.RolloutResponse(
request_id="r1",
prompt_id="prompt_1",
status="COMPLETED",
env_reward=1.0,
)
self.mock_rollout_1.generate.return_value = [resp]

results = await self.engine.generate([request])

self.assertLen(results, 1)
self.mock_rollout_1.generate.assert_called_once()
self.assertEqual(
self.mock_rollout_1.generate.call_args.kwargs["requests"], [request]
)

asyncio.run(_run())

def test_poll_rollouts_aggregates_worker_responses(self):
async def _run():
resp1 = datatypes.RolloutResponse(
Expand Down Expand Up @@ -123,11 +180,42 @@ async def _run():
self.assertEqual(res, {"loss": 0.5})

self.mock_actor.fwd_bwd.assert_called_once_with(
batch=mock_payload,
payload=mock_payload,
skip_jit=False,
)
self.mock_actor.update.assert_not_called()

asyncio.run(_run())

def test_train_step_applies_optimizer_on_last_microbatch(self):
async def _run():
self.mock_actor.fwd_bwd.return_value = datatypes.Response(
metadata={"queued": True}
)
self.mock_actor.update.return_value = 3
mock_payload = mock.MagicMock(spec=datatypes.RLTrainerPayload)

res = await self.engine.train_step(
mock_payload,
role=datatypes.Role.ACTOR,
accumulate_gradients=True,
apply_optimizer=False,
apply_optimizer=True,
)
self.assertEqual(
res,
{
"fwd_bwd": datatypes.Response(metadata={"queued": True}),
"updated": True,
"train_step": 3,
"accumulated": True,
},
)

self.mock_actor.fwd_bwd.assert_called_once_with(
payload=mock_payload,
skip_jit=False,
)
self.mock_actor.update.assert_called_once_with()

asyncio.run(_run())

Expand All @@ -150,6 +238,15 @@ async def _run():

asyncio.run(_run())

def test_sync_weights_requires_weight_sync_metadata(self):
async def _run():
self.mock_actor.prepare_weight_sync.return_value = datatypes.Response()

with self.assertRaisesRegex(RuntimeError, "WeightSyncMetadata"):
await self.engine.sync_weights(role=datatypes.Role.ACTOR)

asyncio.run(_run())

def test_balancer_prefix_routing(self):

async def _run():
Expand Down
79 changes: 64 additions & 15 deletions tests/experimental/orchestrator/rl_program_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,20 +29,24 @@ class RLProgramTest(absltest.TestCase):
def setUp(self):
super().setUp()
self.mock_engine = mock.MagicMock(spec=rl_engine_interface.AbstractRLEngine)
mock_resp = datatypes.RolloutResponse(
self.mock_request = datatypes.RolloutRequest(
request_id="r1",
status="COMPLETED",
env_reward=1.0,
prompt_id="prompt1",
prompt="prompt1",
)
mock_item = datatypes.TrajectoryItem(
pair_index=0,
group_id="prompt1",
start_step=0,
traj=datatypes.Trajectory(
reward=1.0,
status=datatypes.TrajectoryStatus.SUCCEEDED,
),
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),
)
],
completion_tokens=np.array([3, 4], dtype=np.int32),
action_mask=np.array([1, 1], dtype=np.int32),
)
self.mock_engine.generate = mock.AsyncMock(return_value=[mock_resp])
self.mock_engine.generate = mock.AsyncMock(return_value=[mock_item])
self.mock_engine.train_step = mock.AsyncMock(return_value="mock_train_result")
self.mock_engine.sync_weights = mock.AsyncMock(return_value=1)

Expand Down Expand Up @@ -76,28 +80,73 @@ def on_end(step, result):
on_step_end=on_end,
)

res = program.step_once(prompts=["prompt1"])
res = program.step_once(prompts=[self.mock_request])

self.assertEqual(res, "mock_train_result")
self.mock_engine.generate.assert_called_once_with(prompts=["prompt1"])
self.mock_engine.generate.assert_called_once_with(
prompts=[self.mock_request]
)
self.mock_algo.create_trainer_payloads.assert_called_once()
self.mock_engine.train_step.assert_called_once()
self.mock_engine.sync_weights.assert_called_once_with(role=datatypes.Role.ACTOR)
self.assertEqual(program.step, 1)

self.assertEqual(begin_calls, [0])
self.assertEqual(end_calls, [(1, "mock_train_result")])
self.assertIsNotNone(program.last_step_result)
self.assertEqual(program.last_step_result.num_rollouts, 1)
self.assertEqual(program.last_step_result.num_microbatches, 1)

def test_step_once_can_skip_weight_sync(self):
program = rl_program.SyncRLProgram(
engine=self.mock_engine,
algo=self.mock_algo,
assembler=self.assembler,
sync_weights=False,
)

res = program.step_once(prompts=[self.mock_request])

self.assertEqual(res, "mock_train_result")
self.mock_engine.sync_weights.assert_not_called()
self.assertEqual(program.step, 1)
self.assertIsNotNone(program.last_step_result)
self.assertEqual(program.last_step_result.policy_version, 1)

def test_run_accepts_orchestrator_supplied_engine(self):
program = rl_program.SyncRLProgram(
algo=self.mock_algo,
assembler=self.assembler,
sync_weights=False,
)

program.run(
train_dataset=[[self.mock_request]],
num_steps=1,
engine=self.mock_engine,
)

self.mock_engine.generate.assert_called_once_with(
prompts=[self.mock_request]
)
self.mock_engine.train_step.assert_called_once()
self.assertEqual(program.step, 1)

def test_eval_step_once_flow(self):
program = rl_program.SyncRLProgram(
engine=self.mock_engine,
algo=self.mock_algo,
assembler=self.assembler,
)
res = program.eval_step_once(prompts=["eval_prompt"])
eval_request = datatypes.RolloutRequest(
request_id="eval_r1",
prompt_id="eval_prompt",
prompt="eval_prompt",
)
res = program.eval_step_once(prompts=[eval_request])

self.assertLen(res, 1)
self.mock_engine.generate.assert_called_once_with(prompts=["eval_prompt"])
self.mock_engine.generate.assert_called_once_with(prompts=[eval_request])
self.mock_algo.create_trainer_payloads.assert_called_once()


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -72,6 +72,7 @@ def test_single_sampling_request(self):
self.assertIsInstance(response, base_sampler_lib.SamplingResponse)
self.assertEqual(response.request_id, "vllm_req_01")
self.assertEqual(response.text, "completion 1")
np.testing.assert_array_equal(response.prompt_token_ids, [1, 2])
np.testing.assert_array_equal(response.token_ids, [10, 20, 30])
self.mock_vllm_sampler.assert_called_once()

Expand Down Expand Up @@ -101,6 +102,8 @@ def test_batch_sampling_requests(self):
self.assertLen(responses, 2)
self.assertEqual(responses[0].text, "completion 1")
self.assertEqual(responses[1].text, "completion 2")
np.testing.assert_array_equal(responses[0].prompt_token_ids, [1, 2])
np.testing.assert_array_equal(responses[1].prompt_token_ids, [3, 4])

def test_weight_sync(self):
mock_weights = {"layer1": "weights"}
Expand Down
1 change: 0 additions & 1 deletion tests/experimental/rollout/rollout_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,7 +48,6 @@ def setUp(self):
self.server = remote_execution.InProcessRemoteExecutionServer(self.service)
self.actor_handle = remote_execution.InProcessActorHandle(self.server)
self.service.start()
self.service.initialize()

def tearDown(self):
super().tearDown()
Expand Down
4 changes: 4 additions & 0 deletions tests/experimental/rollout/sampler_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,7 @@ def _sample_response() -> base_sampler_lib.SamplingResponse:
return base_sampler_lib.SamplingResponse(
request_id="sample-req-1",
text="Hello from Tunix sampler!",
prompt_token_ids=np.array([1, 2], dtype=np.int32),
token_ids=np.array([101, 102, 103], dtype=np.int32),
logprobs=np.array([-0.1, -0.2, -0.05], dtype=np.float32),
finish_reason="stop",
Expand Down Expand Up @@ -104,6 +105,9 @@ def test_sampling_response_round_trips_through_cloudpickle(self):
self.assertEqual(restored.metadata, original.metadata)
self.assertIsNone(restored.error)
np.testing.assert_array_equal(restored.token_ids, original.token_ids)
np.testing.assert_array_equal(
restored.prompt_token_ids, original.prompt_token_ids
)
np.testing.assert_allclose(restored.logprobs, original.logprobs)
np.testing.assert_array_equal(
restored.routed_experts, original.routed_experts
Expand Down
5 changes: 5 additions & 0 deletions tests/experimental/rollout/vanilla_sampler_adapter_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
import asyncio
from absl.testing import absltest
from flax import nnx
import numpy as np
from tunix.experimental.rollout import sampler as base_sampler_lib
from tunix.experimental.rollout import vanilla_sampler_adapter
from tunix.generate import sampler as generate_sampler_lib
Expand Down Expand Up @@ -59,6 +60,7 @@ def test_single_sampling_request(self):
self.assertIsInstance(response, base_sampler_lib.SamplingResponse)
self.assertEqual(response.request_id, "req_01")
self.assertIsNotNone(response.text)
self.assertGreater(response.prompt_token_ids.size, 0)

def test_sampling_request_with_logprobs(self):
req = base_sampler_lib.SamplingRequest(
Expand Down Expand Up @@ -101,6 +103,8 @@ def test_batch_sampling_requests(self):
self.assertEqual(responses[1].request_id, "req_b")
self.assertIsNotNone(responses[0].text)
self.assertIsNotNone(responses[1].text)
self.assertGreater(responses[0].prompt_token_ids.size, 0)
self.assertGreater(responses[1].prompt_token_ids.size, 0)

def test_construct_with_integer_cache_size(self):
sampler_adapter_direct = vanilla_sampler_adapter.VanillaSamplerAdapter(
Expand All @@ -121,6 +125,7 @@ def test_construct_with_integer_cache_size(self):
self.assertIsInstance(response, base_sampler_lib.SamplingResponse)
self.assertEqual(response.request_id, "req_direct")
self.assertIsNotNone(response.text)
self.assertEqual(response.prompt_token_ids.dtype, np.int32)

def test_uninitialized_sampler_raises(self):
uninit_sampler = vanilla_sampler_adapter.VanillaSamplerAdapter(
Expand Down
6 changes: 4 additions & 2 deletions tests/experimental/worker/abstract_worker_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,15 +88,17 @@ def test_lifecycle_state_transitions(self, worker_cls, module_path, kwargs):

with mock.patch(module_path) as mock_response:
# Hook into the Response() creation inside the try-block to check state
def verify_initializing():
def verify_initializing(*args, **kwargs):
del args, kwargs
self.assertEqual(worker.state, WorkerState.INITIALIZING)
return mock.DEFAULT

mock_response.side_effect = verify_initializing
worker.initialize()
self.assertEqual(worker.state, WorkerState.READY)

def verify_compiling():
def verify_compiling(*args, **kwargs):
del args, kwargs
self.assertEqual(worker.state, WorkerState.COMPILING)
return mock.DEFAULT

Expand Down
Loading
Loading