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
163 changes: 163 additions & 0 deletions tests/experimental/orchestrator/batch_assembly_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
from absl.testing import absltest
import numpy as np
from tunix.experimental.common import datatypes
from tunix.experimental.common import lineage
from tunix.experimental.orchestrator import batch_assembly
from tunix.rl import common as rl_common

Expand Down Expand Up @@ -276,6 +277,119 @@ def test_sequence_packed_assembler_multiple_bins(self):
self.assertEqual(payloads[0].token_ids.shape, (1, 12))
self.assertEqual(payloads[1].token_ids.shape, (1, 12))

def test_pack_merges_lineage_contexts(self):
ctx1 = lineage.LineageContext(
tracking_id="traj_p1_g0", parent_tracking_ids=["p1"]
)
ctx1.add_event("engine.dispatch", "rollout", {"group_index": 0})
ctx1.add_event("worker.rollout", "generate", {"worker_id": "w0"})

ctx2 = lineage.LineageContext(
tracking_id="traj_p1_g1", parent_tracking_ids=["p1"]
)
ctx2.add_event("engine.dispatch", "rollout", {"group_index": 1})
ctx2.add_event("worker.rollout", "generate", {"worker_id": "w1"})

payload1 = datatypes.RLTrainerPayload(
token_ids=np.array([1, 2, 3, 4], dtype=np.int32),
token_mask=np.ones(4, dtype=np.float32),
loss_mask=np.ones(4, dtype=np.float32),
advantages=np.ones(4, dtype=np.float32),
metadata={"lineage": ctx1},
)
payload2 = datatypes.RLTrainerPayload(
token_ids=np.array([5, 6, 7, 8], dtype=np.int32),
token_mask=np.ones(4, dtype=np.float32),
loss_mask=np.ones(4, dtype=np.float32),
advantages=np.ones(4, dtype=np.float32),
metadata={"lineage": ctx2},
)

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

self.assertLen(payloads, 1)
batch_payload = payloads[0]
self.assertIn("lineage", batch_payload.metadata)
batch_ctx = batch_payload.metadata["lineage"]
self.assertEqual(batch_ctx.tracking_id, "batch_0")
self.assertEqual(
sorted(batch_ctx.parent_tracking_ids), ["traj_p1_g0", "traj_p1_g1"]
)
self.assertLen(batch_ctx.events, 1)
merge_event = batch_ctx.events[0]
self.assertEqual(merge_event.component, "orchestrator.assembler")
self.assertEqual(merge_event.operation, "pack")
self.assertEqual(
merge_event.attributes.get("packing_type"), "sequence_packed"
)
self.assertEqual(merge_event.attributes.get("bin_size"), 2)

def test_pack_multi_bin_creates_distinct_lineage_contexts(self):
ctx1 = lineage.LineageContext(
tracking_id="traj_p1_g0", parent_tracking_ids=["p1"]
)
ctx2 = lineage.LineageContext(
tracking_id="traj_p2_g0", parent_tracking_ids=["p2"]
)

p1 = datatypes.RLTrainerPayload(
token_ids=np.arange(10, dtype=np.int32),
token_mask=np.ones(10, dtype=np.float32),
loss_mask=np.ones(10, dtype=np.float32),
advantages=np.ones(10, dtype=np.float32),
metadata={"lineage": ctx1},
)
p2 = datatypes.RLTrainerPayload(
token_ids=np.arange(8, dtype=np.int32),
token_mask=np.ones(8, dtype=np.float32),
loss_mask=np.ones(8, dtype=np.float32),
advantages=np.ones(8, dtype=np.float32),
metadata={"lineage": ctx2},
)

assembler = batch_assembly.SequencePackedBatchAssembler(max_packed_len=12)
payloads = assembler.pack([p1, p2])

self.assertLen(payloads, 2)
self.assertEqual(payloads[0].metadata["lineage"].tracking_id, "batch_0")
self.assertEqual(
payloads[0].metadata["lineage"].parent_tracking_ids, ["traj_p1_g0"]
)
self.assertEqual(payloads[1].metadata["lineage"].tracking_id, "batch_1")
self.assertEqual(
payloads[1].metadata["lineage"].parent_tracking_ids, ["traj_p2_g0"]
)

def test_pack_without_lineage_returns_clean_payload(self):
p = datatypes.RLTrainerPayload(
token_ids=np.array([1, 2], dtype=np.int32),
token_mask=np.ones(2, dtype=np.float32),
loss_mask=np.ones(2, dtype=np.float32),
advantages=np.ones(2, dtype=np.float32),
)
assembler = batch_assembly.SequencePackedBatchAssembler(max_packed_len=8)
payloads = assembler.pack([p])
self.assertLen(payloads, 1)
self.assertNotIn("lineage", payloads[0].metadata)

def test_pack_sequential_increments_monotonic_batch_ids(self):
ctx = lineage.LineageContext(
tracking_id="traj_t0_g0", parent_tracking_ids=["t0"]
)
p = datatypes.RLTrainerPayload(
token_ids=np.array([1, 2], dtype=np.int32),
token_mask=np.ones(2, dtype=np.float32),
loss_mask=np.ones(2, dtype=np.float32),
advantages=np.ones(2, dtype=np.float32),
metadata={"lineage": ctx},
)
assembler = batch_assembly.SequencePackedBatchAssembler(max_packed_len=8)
out1 = assembler.pack([p])
out2 = assembler.pack([p])
self.assertEqual(out1[0].metadata["lineage"].tracking_id, "batch_0")
self.assertEqual(out2[0].metadata["lineage"].tracking_id, "batch_1")


class GRPOTrainExampleAssemblerTest(absltest.TestCase):

Expand Down Expand Up @@ -833,6 +947,55 @@ def test_none_advantages_defaults_to_zeros(self):
self.assertEqual(payload.advantages.shape, (2, 5))
np.testing.assert_allclose(payload.advantages[0], [0.0, 0.0, 0.0, 0.0, 0.0])

def test_pack_merges_lineage_contexts(self):
ctx1 = lineage.LineageContext(
tracking_id="traj_p1_g0", parent_tracking_ids=["p1"]
)
ctx1.add_event("worker.rollout", "generate", {"worker_id": "w0"})
ctx2 = lineage.LineageContext(
tracking_id="traj_p1_g1", parent_tracking_ids=["p1"]
)
ctx2.add_event("worker.rollout", "generate", {"worker_id": "w1"})

item1 = _make_payload(2, 2)
item1.metadata = {"lineage": ctx1}
item2 = _make_payload(2, 2)
item2.metadata = {"lineage": ctx2}

assembler = batch_assembly.PaddedBatchAssembler(
batch_size=2, max_prompt_length=4, max_response_length=4, pad_id=0
)
payloads = assembler.pack([item1, item2])

self.assertLen(payloads, 1)
batch_payload = payloads[0]
self.assertIn("lineage", batch_payload.metadata)
batch_ctx = batch_payload.metadata["lineage"]
self.assertEqual(batch_ctx.tracking_id, "batch_0")
self.assertEqual(
sorted(batch_ctx.parent_tracking_ids), ["traj_p1_g0", "traj_p1_g1"]
)
self.assertLen(batch_ctx.events, 1)
merge_event = batch_ctx.events[0]
self.assertEqual(merge_event.component, "orchestrator.assembler")
self.assertEqual(merge_event.operation, "pack")
self.assertEqual(merge_event.attributes.get("packing_type"), "padded")
self.assertEqual(merge_event.attributes.get("chunk_size"), 2)

def test_pack_sequential_increments_monotonic_batch_ids(self):
ctx = lineage.LineageContext(
tracking_id="traj_t0_g0", parent_tracking_ids=["t0"]
)
item = _make_payload(2, 2)
item.metadata = {"lineage": ctx}
assembler = batch_assembly.PaddedBatchAssembler(
batch_size=2, max_prompt_length=4, max_response_length=4, pad_id=0
)
out1 = assembler.pack([item])
out2 = assembler.pack([item])
self.assertEqual(out1[0].metadata["lineage"].tracking_id, "batch_0")
self.assertEqual(out2[0].metadata["lineage"].tracking_id, "batch_1")


if __name__ == "__main__":
absltest.main()
64 changes: 64 additions & 0 deletions tests/experimental/orchestrator/distributed_rl_engine_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -821,6 +821,70 @@ async def _run():

asyncio.run(_run())

def test_train_step_appends_lineage_event(self):
async def _run():
self.mock_actor.fwd_bwd.return_value = datatypes.Response(
metadata={"queued": True}
)
self.mock_actor.update.return_value = 1

ctx = lineage.LineageContext(
tracking_id="batch_0", parent_tracking_ids=["traj_p1_g0"]
)
ctx.add_event("orchestrator.assembler", "pack")
payload = datatypes.RLTrainerPayload(
advantages=np.array([1.0], dtype=np.float32),
loss_mask=np.array([[1]], dtype=np.int32),
metadata={"lineage": ctx},
)

res = await self.engine.train_step(
payload, accumulate_gradients=False, apply_optimizer=True
)
self.assertTrue(res["updated"])
self.assertEqual(res["train_step"], 1)

# Check TrainRequest passed to actor worker
self.mock_actor.fwd_bwd.assert_called_once()
req = self.mock_actor.fwd_bwd.call_args.kwargs["request"]
self.assertTrue(req.request_id.startswith("train_req_"))
self.assertIs(req.metadata["lineage"], ctx)

# Check lineage event appended
self.assertLen(ctx.events, 2)
self.assertEqual(ctx.events[1].component, "engine.train")
self.assertEqual(ctx.events[1].operation, "train_step")
self.assertFalse(ctx.events[1].attributes["accumulate_gradients"])
self.assertTrue(ctx.events[1].attributes["apply_optimizer"])
self.assertEqual(ctx.events[1].attributes["policy_version"], 0)

asyncio.run(_run())

def test_train_step_without_lineage_context(self):
async def _run():
self.mock_actor.fwd_bwd.return_value = datatypes.Response(
metadata={"queued": True}
)
self.mock_actor.update.return_value = 1

payload = datatypes.RLTrainerPayload(
advantages=np.array([1.0], dtype=np.float32),
loss_mask=np.array([[1]], dtype=np.int32),
metadata={},
)

res = await self.engine.train_step(
payload, accumulate_gradients=False, apply_optimizer=True
)
self.assertTrue(res["updated"])

self.mock_actor.fwd_bwd.assert_called_once()
req = self.mock_actor.fwd_bwd.call_args.kwargs["request"]
self.assertTrue(req.request_id.startswith("train_req_"))
self.assertNotIn("lineage", req.metadata)

asyncio.run(_run())

def test_generate_stamps_lineage_context_for_rollout_requests(self):
async def _run():
req = datatypes.RolloutRequest(
Expand Down
39 changes: 39 additions & 0 deletions tests/experimental/worker/trainer_worker_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
from absl.testing import absltest
import numpy as np
from tunix.experimental.common import datatypes
from tunix.experimental.common import lineage
from tunix.experimental.train import abstract_trainer
from tunix.experimental.worker import trainer_worker

Expand Down Expand Up @@ -126,6 +127,44 @@ def test_update_returns_step_count(self):
step = self.worker.update()
self.assertEqual(step, 11)

def test_fwd_bwd_appends_trainer_worker_lineage_event(self):
ctx = lineage.LineageContext(
tracking_id="batch_0", parent_tracking_ids=["traj_p1_g0"]
)
ctx.add_event("orchestrator.assembler", "pack")

payload = datatypes.RLTrainerPayload(
advantages=np.array([1.0], dtype=np.float32),
loss_mask=np.array([[1]], dtype=np.int32),
)
request = datatypes.TrainRequest(
request_id="req-train-lineage",
payload=payload,
metadata={"lineage": ctx},
)

resp = self.worker.fwd_bwd(request=request)

self.assertIn("lineage", resp.metadata)
self.assertIs(resp.metadata["lineage"], ctx)
self.assertLen(ctx.events, 2)
self.assertEqual(ctx.events[1].component, "worker.trainer")
self.assertEqual(ctx.events[1].operation, "fwd_bwd")
self.assertEqual(ctx.events[1].attributes.get("worker_id"), "trainer_0")
self.assertEqual(ctx.events[1].attributes.get("policy_version"), 3)

def test_fwd_bwd_handles_request_without_lineage(self):
payload = datatypes.RLTrainerPayload(
advantages=np.array([1.0], dtype=np.float32),
loss_mask=np.array([[1]], dtype=np.int32),
)
request = datatypes.TrainRequest(
request_id="req-train-no-lineage",
payload=payload,
metadata={},
)
resp = self.worker.fwd_bwd(request=request)
self.assertNotIn("lineage", resp.metadata)

if __name__ == "__main__":
absltest.main()
Loading
Loading