Skip to content
Open
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
84 changes: 53 additions & 31 deletions tests/experimental/common/datatypes_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ def _rollout_request_dto() -> datatypes.RolloutRequest:
request_id="req-123",
prompt="Solve 2+2",
prompt_id="req-rollout-42",
group_offset_id="group-1",
group_offset_id="1",
generation_kwargs={"max_tokens": 128, "temperature": 0.5},
max_turns=5,
target_policy_version=3,
Expand Down Expand Up @@ -108,41 +108,20 @@ def test_trajectory_response_round_trips_through_cloudpickle(self):
def test_error_result_round_trips(self):
result = datatypes.RolloutResponse(
request_id="req-2",
status="TIMEOUT",
error=datatypes.ErrorInfo(
error_type="TimeoutError",
message="deadline exceeded",
retryable=True,
),
status="ERROR",
error="Model worker died",
)

restored = cloudpickle.loads(cloudpickle.dumps(result))

self.assertEqual(restored.status, "TIMEOUT")
self.assertEqual(restored.error.error_type, "TimeoutError")
self.assertTrue(restored.error.retryable)
self.assertEqual(restored.prompt_tokens.size, 0)
self.assertEmpty(restored.segments)
self.assertEqual(restored.request_id, "req-2")
self.assertEqual(restored.status, "ERROR")
self.assertEqual(restored.error, "Model worker died")

def test_token_segment_enforces_shapes(self):
with self.assertRaisesRegex(
ValueError, "loss_mask shape .* != tokens shape"
):
datatypes.TokenSegment(
source="env",
tokens=np.array([1, 2]),
loss_mask=np.array([1]),
)

with self.assertRaisesRegex(
ValueError, "logps shape .* != tokens shape"
):
datatypes.TokenSegment(
source="assistant",
tokens=np.array([1, 2]),
loss_mask=np.array([1, 1]),
logps=np.array([0.5]),
)
tokens = np.array([1, 2, 3])
loss_mask = np.array([1, 0])
with self.assertRaises(ValueError):
datatypes.TokenSegment(source="assistant", tokens=tokens, loss_mask=loss_mask)

def test_from_trajectory(self):
step1 = datatypes.Step(
Expand Down Expand Up @@ -198,6 +177,49 @@ def test_health_report_defaults_heartbeat_unix_s_to_current_time(self):
self.assertGreaterEqual(report.heartbeat_unix_s, before)
self.assertLessEqual(report.heartbeat_unix_s, after)

def test_rollout_request_lineage_and_traj_id(self):
req1 = datatypes.RolloutRequest(
prompt_id="prompt_42",
group_offset_id="3",
)
self.assertEqual(req1.traj_id, "traj_prompt_42_3")
self.assertEqual(req1.request_id, "traj_prompt_42_3")

req2 = datatypes.RolloutRequest(
prompt_id="prompt_42",
group_offset_id="sample_0",
)
self.assertEqual(req2.traj_id, "traj_prompt_42_sample_0")

req3 = datatypes.RolloutRequest(
prompt_id="prompt_single",
)
self.assertEqual(req3.traj_id, "traj_prompt_single")

def test_trajectory_item_lineage_initialization(self):
item = datatypes.TrajectoryItem(
prompt_id="prompt_99",
group_offset_id="2",
start_step=0,
traj=datatypes.Trajectory(trajectory_id="traj_prompt_99_2", reward=1.0),
)
self.assertEqual(item.group_id, "prompt_99")
self.assertEqual(item.pair_index, 2)
self.assertEqual(item.trajectory_id, "traj_prompt_99_2")

def test_trajectory_item_legacy_keyword_initialization(self):
item = datatypes.TrajectoryItem(
group_id="legacy_group_42",
pair_index=3,
start_step=0,
traj=datatypes.Trajectory(reward=1.0),
)
self.assertEqual(item.prompt_id, "legacy_group_42")
self.assertEqual(item.group_offset_id, "3")
self.assertEqual(item.group_id, "legacy_group_42")
self.assertEqual(item.pair_index, 3)
self.assertEqual(item.trajectory_id, "traj_legacy_group_42_3")


if __name__ == "__main__":
absltest.main()
16 changes: 8 additions & 8 deletions tests/experimental/orchestrator/algorithm_adapter_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,8 +35,8 @@ def test_grpo_advantage_normalization(self):
def test_grpo_create_trainer_payloads(self):
adapter = algorithm_adapter.GRPOAdapter(group_size=2)
item1 = datatypes.TrajectoryItem(
pair_index=0,
group_id="g1",
prompt_id="g1",
group_offset_id="0",
start_step=0,
traj=datatypes.Trajectory(reward=1.0),
)
Expand All @@ -45,8 +45,8 @@ def test_grpo_create_trainer_payloads(self):
item1.action_mask = np.array([1, 1], dtype=np.float32)

item2 = datatypes.TrajectoryItem(
pair_index=1,
group_id="g1",
prompt_id="g1",
group_offset_id="1",
start_step=0,
traj=datatypes.Trajectory(reward=2.0),
)
Expand All @@ -57,15 +57,15 @@ def test_grpo_create_trainer_payloads(self):
payloads = adapter.create_trainer_payloads([item1, item2], rewards=[1.0, 2.0])
self.assertLen(payloads, 2)
self.assertIsInstance(payloads[0], datatypes.RLTrainerPayload)
self.assertLess(payloads[0].advantages[0], 0.0)
self.assertGreater(payloads[1].advantages[0], 0.0)
self.assertEqual(payloads[0].trajectory_ids, ["traj_g1_0"])
self.assertEqual(payloads[1].trajectory_ids, ["traj_g1_1"])
self.assertEqual(adapter.loss_fn(), algo_core.grpo_loss_fn)

def test_ppo_advantages_and_trainer_payloads(self):
adapter = algorithm_adapter.PPOAdapter(group_size=2, gamma=0.99, lam=0.95)
item = datatypes.TrajectoryItem(
pair_index=0,
group_id="g1",
prompt_id="g1",
group_offset_id="0",
start_step=0,
traj=datatypes.Trajectory(reward=1.0),
)
Expand Down
26 changes: 26 additions & 0 deletions tests/experimental/orchestrator/batch_assembly_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,10 +101,12 @@ def test_grpo_train_example_assembler(self):
max_response_length=5,
pad_id=0,
)
payload.trajectory_ids = ["traj_prompt_1_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))
self.assertEqual(train_example.trajectory_ids, ("traj_prompt_1_0", "__pad__"))
np.testing.assert_array_equal(
train_example.prompt_ids[0], np.array([0, 0, 10, 11])
)
Expand All @@ -118,6 +120,30 @@ def test_grpo_train_example_assembler(self):
train_example.advantages[0], np.array([2, 2, 2, 0, 0])
)

def test_sequence_packed_assembler_lineage(self):
payload1 = datatypes.RLTrainerPayload(
token_ids=np.array([1, 2, 3], dtype=np.int32),
token_mask=np.array([1, 1, 1], dtype=np.float32),
loss_mask=np.array([0, 1, 1], dtype=np.float32),
advantages=np.full(3, 1.0, dtype=np.float32),
trajectory_ids=["traj_p1_0"],
)
payload2 = datatypes.RLTrainerPayload(
token_ids=np.array([4, 5], dtype=np.int32),
token_mask=np.array([1, 1], dtype=np.float32),
loss_mask=np.array([1, 1], dtype=np.float32),
advantages=np.full(2, 2.0, dtype=np.float32),
trajectory_ids=["traj_p2_0"],
)

assembler = batch_assembly.SequencePackedBatchAssembler(max_packed_len=8)
payloads = assembler.pack([payload1, payload2])

self.assertLen(payloads, 1)
packed = payloads[0]
self.assertEqual(packed.trajectory_ids, ["traj_p1_0", "traj_p2_0"])
self.assertEqual(packed.segment_lineage, {1: "traj_p1_0", 2: "traj_p2_0"})


if __name__ == "__main__":
absltest.main()
6 changes: 4 additions & 2 deletions tests/experimental/orchestrator/rl_program_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,8 +36,8 @@ def setUp(self):
prompt="prompt1",
)
mock_item = datatypes.TrajectoryItem(
pair_index=0,
group_id="prompt1",
prompt_id="prompt1",
group_offset_id="0",
start_step=0,
traj=datatypes.Trajectory(
reward=1.0,
Expand Down Expand Up @@ -97,6 +97,8 @@ def on_end(step, 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)
self.assertEqual(program.last_step_result.global_batch_id, "batch_000000")
self.assertEqual(program.last_step_result.trajectory_ids, ("traj_prompt1_0",))

def test_step_once_can_skip_weight_sync(self):
program = rl_program.SyncRLProgram(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -31,8 +31,8 @@ def _create_item(
"""Helper to create a TrajectoryItem for testing."""
traj = datatypes.Trajectory(reward=reward)
return datatypes.TrajectoryItem(
pair_index=pair_index,
group_id=group_id,
prompt_id=group_id,
group_offset_id=str(pair_index),
start_step=0,
traj=traj,
metadata={"task_id": task_id},
Expand Down
Loading
Loading