diff --git a/tests/experimental/orchestrator/batch_assembly_test.py b/tests/experimental/orchestrator/batch_assembly_test.py index d8e970d34..ad0fb83d0 100644 --- a/tests/experimental/orchestrator/batch_assembly_test.py +++ b/tests/experimental/orchestrator/batch_assembly_test.py @@ -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 @@ -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): @@ -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() diff --git a/tests/experimental/orchestrator/distributed_rl_engine_test.py b/tests/experimental/orchestrator/distributed_rl_engine_test.py index 8ee716162..5909a1b13 100644 --- a/tests/experimental/orchestrator/distributed_rl_engine_test.py +++ b/tests/experimental/orchestrator/distributed_rl_engine_test.py @@ -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( diff --git a/tests/experimental/worker/trainer_worker_test.py b/tests/experimental/worker/trainer_worker_test.py index f3ba2ce37..6b29debd8 100644 --- a/tests/experimental/worker/trainer_worker_test.py +++ b/tests/experimental/worker/trainer_worker_test.py @@ -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 @@ -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() diff --git a/tunix/experimental/orchestrator/batch_assembly.py b/tunix/experimental/orchestrator/batch_assembly.py index f7e8431ec..92d5460a3 100644 --- a/tunix/experimental/orchestrator/batch_assembly.py +++ b/tunix/experimental/orchestrator/batch_assembly.py @@ -22,10 +22,13 @@ # TODO: Align SequencePackedBatchAssembler with the rest of the ecosystem and potentially move to a common library. """ -from absl import logging +from collections.abc import Mapping, Sequence from typing import Any, Generic, Protocol, Sequence, TypeVar +from typing import Any, Generic, Protocol, TypeVar +from absl import logging import numpy as np from tunix.experimental.common import datatypes +from tunix.experimental.common import lineage from tunix.rl import common as rl_common T = TypeVar("T") @@ -34,8 +37,15 @@ class BatchAssembler(Generic[T], Protocol): """Universal batch assembly protocol for microbatch packing.""" - def pack(self, items: Sequence[T]) -> list[Any]: - """Packs items into hardware-sized microbatch trainer payloads.""" + def pack( + self, + items: Sequence[T], + ) -> list[Any]: + """Packs items into hardware-sized microbatch trainer payloads. + + Args: + items: Sequence of unbatched items to assemble. + """ ... @@ -134,15 +144,60 @@ def with_ref_per_token_logps( ) +def _merge_batch_lineage( + items: Sequence[Any], + *, + batch_id: str, + attributes: Mapping[str, Any] | None = None, +) -> lineage.LineageContext | None: + """Extracts and merges lineage contexts from a sequence of batch items. + + Args: + items: Sequence of items that may carry lineage context in their metadata. + batch_id: Tracking ID to assign to the merged batch context. + attributes: Optional key-value metadata attached to the merge event. + + Returns: + The merged LineageContext, or None if no upstream lineage contexts exist. + """ + lineages = [ + it.metadata["lineage"] + for it in items + if isinstance(getattr(it, "metadata", None), Mapping) + and it.metadata.get("lineage") is not None + ] + if not lineages: + return None + + return lineage.LineageContext.merge( + batch_id=batch_id, + contexts=lineages, + component="orchestrator.assembler", + operation="pack", + attributes=dict(attributes) if attributes else None, + ) + + class SequencePackedBatchAssembler: """1D Sequence Packing: Concatenates items into dense [1, max_packed_len] buffers.""" # TODO: align implementation with current path. def __init__(self, max_packed_len: int = 8192, pad_id: int = 0): self.max_packed_len = max_packed_len self.pad_id = pad_id + self._batch_counter = 0 - def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RLTrainerPayload]: - """Bin-packs items into dense 1D buffers with segment boundaries.""" + def pack( + self, + items: Sequence[datatypes.RLTrainerPayload], + ) -> list[datatypes.RLTrainerPayload]: + """Bin-packs items into dense 1D buffers with segment boundaries. + + Args: + items: Sequence of unbatched RLTrainerPayloads to pack. + + Returns: + List of packed, segment-aligned RLTrainerPayload microbatches. + """ if not items: return [] @@ -169,7 +224,7 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL bin_lengths.append(length) payloads: list[datatypes.RLTrainerPayload] = [] - for b_items in bins: + for b_idx, b_items in enumerate(bins): all_tokens = [] all_loss_masks = [] all_action_masks = [] @@ -250,6 +305,18 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL concat_ref = np.concatenate(all_ref_logprobs) batch_ref_lp = np.pad(concat_ref[: self.max_packed_len], (0, pad_len), constant_values=0.0)[np.newaxis, :] + batch_tracking_id = f"batch_{self._batch_counter}" + + merged_lineage = _merge_batch_lineage( + b_items, + batch_id=batch_tracking_id, + attributes={ + "packing_type": "sequence_packed", + "bin_size": len(b_items), + "packed_len": self.max_packed_len, + }, + ) + payload_metadata = {"lineage": merged_lineage} if merged_lineage else {} payload = datatypes.RLTrainerPayload( token_ids=padded_tokens[np.newaxis, :], token_mask=padded_segment_ids[np.newaxis, :], @@ -260,8 +327,10 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL ref_per_token_logps=batch_ref_lp, segment_ids=padded_segment_ids[np.newaxis, :], segment_positions=padded_segment_positions[np.newaxis, :], + metadata=payload_metadata, ) payloads.append(payload) + self._batch_counter += 1 return payloads @@ -291,7 +360,8 @@ def __init__( self.pad_id = pad_id def pack( - self, items: Sequence[datatypes.RLTrainerPayload] + self, + items: Sequence[datatypes.RLTrainerPayload], ) -> list[rl_common.TrainExample]: item_list = list(items) if not item_list: @@ -434,13 +504,15 @@ def __init__( self.max_prompt_length = max_prompt_length self.max_response_length = max_response_length self.pad_id = pad_id + self._batch_counter = 0 @property def max_seq_len(self) -> int: return self.max_prompt_length + self.max_response_length def pack( - self, items: Sequence[datatypes.RLTrainerPayload] + self, + items: Sequence[datatypes.RLTrainerPayload], ) -> list[datatypes.RLTrainerPayload]: """Pads items into rectangular 2D batches `[B, P + C]`.""" item_list = list(items) @@ -449,7 +521,22 @@ def pack( payloads: list[datatypes.RLTrainerPayload] = [] for i in range(0, len(item_list), self.batch_size): - payloads.append(self._pack_chunk(item_list[i : i + self.batch_size])) + chunk = item_list[i : i + self.batch_size] + payload = self._pack_chunk(chunk) + batch_tracking_id = f"batch_{self._batch_counter}" + merged_lineage = _merge_batch_lineage( + chunk, + batch_id=batch_tracking_id, + attributes={ + "packing_type": "padded", + "chunk_size": len(chunk), + "batch_size": self.batch_size, + }, + ) + if merged_lineage: + payload.metadata["lineage"] = merged_lineage + payloads.append(payload) + self._batch_counter += 1 return payloads def _pack_chunk( diff --git a/tunix/experimental/orchestrator/distributed_rl_engine.py b/tunix/experimental/orchestrator/distributed_rl_engine.py index 13c6e5826..8650ad5fe 100644 --- a/tunix/experimental/orchestrator/distributed_rl_engine.py +++ b/tunix/experimental/orchestrator/distributed_rl_engine.py @@ -499,6 +499,18 @@ async def train_step( apply_optimizer, ) metadata = dict(getattr(payload, "metadata", {}) or {}) + lineage_ctx = metadata.get("lineage") + if lineage_ctx is not None and hasattr(lineage_ctx, "add_event"): + lineage_ctx.add_event( + component="engine.train", + operation="train_step", + attributes={ + "accumulate_gradients": accumulate_gradients, + "apply_optimizer": apply_optimizer, + "policy_version": self._policy_version, + }, + ) + request = datatypes.TrainRequest( request_id=f"train_req_{uuid.uuid4().hex[:8]}", payload=payload, diff --git a/tunix/experimental/worker/trainer_worker.py b/tunix/experimental/worker/trainer_worker.py index 09cdb8ccc..adc786e6c 100644 --- a/tunix/experimental/worker/trainer_worker.py +++ b/tunix/experimental/worker/trainer_worker.py @@ -154,6 +154,17 @@ def fwd_bwd( """Executes one forward/backward pass.""" self._ensure_ready() req_metadata = dict(request.metadata) if request.metadata else {} + if request.metadata and "lineage" in request.metadata: + lineage_ctx = request.metadata["lineage"] + if hasattr(lineage_ctx, "add_event"): + lineage_ctx.add_event( + component="worker.trainer", + operation="fwd_bwd", + attributes={ + "worker_id": self._worker_id, + "policy_version": self._policy_version(), + }, + ) kwargs.pop("skip_jit", None) try: self._trainer.fwd_bwd(request.payload, **kwargs)