diff --git a/tests/experimental/common/datatypes_test.py b/tests/experimental/common/datatypes_test.py index 58d95baf1..2927fe3f2 100644 --- a/tests/experimental/common/datatypes_test.py +++ b/tests/experimental/common/datatypes_test.py @@ -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, @@ -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( @@ -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() diff --git a/tests/experimental/orchestrator/algorithm_adapter_test.py b/tests/experimental/orchestrator/algorithm_adapter_test.py index c01046f85..b862be3e4 100644 --- a/tests/experimental/orchestrator/algorithm_adapter_test.py +++ b/tests/experimental/orchestrator/algorithm_adapter_test.py @@ -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), ) @@ -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), ) @@ -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), ) diff --git a/tests/experimental/orchestrator/batch_assembly_test.py b/tests/experimental/orchestrator/batch_assembly_test.py index 2969a0290..af5be8212 100644 --- a/tests/experimental/orchestrator/batch_assembly_test.py +++ b/tests/experimental/orchestrator/batch_assembly_test.py @@ -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]) ) @@ -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() diff --git a/tests/experimental/orchestrator/rl_program_test.py b/tests/experimental/orchestrator/rl_program_test.py index c979805cb..49835f376 100644 --- a/tests/experimental/orchestrator/rl_program_test.py +++ b/tests/experimental/orchestrator/rl_program_test.py @@ -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, @@ -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( diff --git a/tests/experimental/queue_manager/trajectory_queue_manager_test.py b/tests/experimental/queue_manager/trajectory_queue_manager_test.py index 6a6714e92..65ef75e2d 100644 --- a/tests/experimental/queue_manager/trajectory_queue_manager_test.py +++ b/tests/experimental/queue_manager/trajectory_queue_manager_test.py @@ -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}, diff --git a/tunix/experimental/common/datatypes.py b/tunix/experimental/common/datatypes.py index 3e93e89cb..214b6f793 100644 --- a/tunix/experimental/common/datatypes.py +++ b/tunix/experimental/common/datatypes.py @@ -36,15 +36,87 @@ TrajectoryStatus = agent_types.TrajectoryStatus -# TODO: Unify this extended TrajectoryItem back into agent_types.TrajectoryItem -# so that all agentic workflows share the same strict token array fields. @dataclasses.dataclass(kw_only=True) -class TrajectoryItem(agent_types.TrajectoryItem): - """Extended TrajectoryItem for Orchestrator with token arrays.""" +class TrajectoryItem: + """Extended TrajectoryItem for Orchestrator with token arrays and lineage tracking. + + Attributes: + prompt_id: Unique identifier for the prompt/task. + group_offset_id: String generation identifier/offset within the prompt's group. + trajectory_id: Semantic trajectory identifier. + start_step: Starting step index within the full trajectory. + traj: The underlying Trajectory object or dict. + prompt_tokens: Unpadded prompt token IDs. + completion_tokens: Unpadded completion token IDs. + action_mask: Binary mask indicating trainable token positions. + policy_version: Version of the policy used to sample the trajectory. + metadata: Additional metadata dictionary. + """ + + prompt_id: str = "" + group_offset_id: str = "" + trajectory_id: str = "" + start_step: int = 0 + traj: Any = None prompt_tokens: np.ndarray | None = None completion_tokens: np.ndarray | None = None action_mask: np.ndarray | None = None policy_version: int = 0 + policy_versions: list[int] = dataclasses.field(default_factory=list) + metadata: dict[str, Any] = dataclasses.field(default_factory=dict) + + # InitVar compatibility for legacy constructor keywords + group_id: dataclasses.InitVar[str | None] = None + pair_index: dataclasses.InitVar[int | None] = None + + # Aliases for backward compatibility with legacy queue managers / algorithms + @property + def group_id(self) -> str: + return self.prompt_id + + @group_id.setter + def group_id(self, val: Any): + self.prompt_id = str(val) + + @property + def pair_index(self) -> int: + try: + return int(self.group_offset_id) + except (ValueError, TypeError): + return 0 + + @pair_index.setter + def pair_index(self, val: int): + self.group_offset_id = str(val) + + def __post_init__( + self, + group_id: str | None = None, + pair_index: int | None = None, + ): + if group_id is not None and not self.prompt_id: + self.prompt_id = str(group_id) + if pair_index is not None and not self.group_offset_id: + self.group_offset_id = str(pair_index) + + if not self.policy_versions: + if hasattr(self.traj, "steps") and self.traj.steps: + self.policy_versions = [ + getattr(step, "policy_version", self.policy_version) + for step in self.traj.steps + ] + elif self.policy_version is not None: + self.policy_versions = [self.policy_version] + + if not self.trajectory_id: + if hasattr(self.traj, "trajectory_id") and self.traj.trajectory_id: + self.trajectory_id = str(self.traj.trajectory_id) + elif self.prompt_id: + self.trajectory_id = ( + f"traj_{self.prompt_id}_{self.group_offset_id}" + if self.group_offset_id + else f"traj_{self.prompt_id}" + ) class Role(str, enum.Enum): """Orchestrator worker roles.""" @@ -233,8 +305,8 @@ class RolloutRequest(Request): prompt: The prompt to generate from (e.g. formatted string, token array, or chat dictionary). prompt_id: Unique identifier for this prompt within a task or dataset. - group_offset_id: Optional identifier for grouping related rollout requests - (e.g. for GRPO). + group_offset_id: String generation identifier/offset within the prompt's group + (e.g. for GRPO generation "0".."G-1"). generation_kwargs: Additional keyword arguments for generation (e.g. sampling parameters like max_tokens and temperature). max_turns: Maximum number of conversation turns for environment interaction. @@ -249,14 +321,16 @@ class RolloutRequest(Request): max_turns: int = 10 target_policy_version: int = 0 + def __post_init__(self): + if not self.request_id: + self.request_id = self.traj_id + @property def traj_id(self) -> str: """Standardized semantic trajectory identifier computed from prompt_id and group_offset_id.""" - return ( - f"traj_{self.prompt_id}_{self.group_offset_id}" - if self.group_offset_id - else f"traj_{self.prompt_id}" - ) + if self.group_offset_id: + return f"traj_{self.prompt_id}_{self.group_offset_id}" + return f"traj_{self.prompt_id}" @dataclasses.dataclass(kw_only=True) @@ -272,12 +346,14 @@ class TokenSegment: loss_mask: Array of ints, 1 where the token is model-emitted (trainable). logps: Array of per-token log-probabilities under the sampling distribution, or None for spans the model did not emit (e.g. env tokens). + policy_version: Weight version used to sample this specific segment. """ source: str tokens: np.ndarray loss_mask: np.ndarray logps: np.ndarray | None = None + policy_version: int = 0 def __post_init__(self): if self.loss_mask.shape != self.tokens.shape: @@ -303,6 +379,8 @@ class RolloutResponse(Response): Attributes: prompt_id: Unique identifier for this prompt within a task or dataset. + group_offset_id: String generation identifier/offset within the prompt's group. + trajectory_id: Semantic trajectory identifier. status: Terminal status name (e.g. a rollout trajectory status, or "CANCELLED"). prompt_tokens: Array of prompt token ids, unpadded, as tokenized by the @@ -311,10 +389,13 @@ class RolloutResponse(Response): call) and environment; concatenated they form the full generated stream. env_reward: Scalar environment reward for the trajectory. policy_version: Weight version used to generate the trajectory. + policy_versions: Sequence of per-step policy versions across multi-turn interactions. error: Failure details when the request did not succeed, else None. """ prompt_id: str = "" + group_offset_id: str = "" + trajectory_id: str = "" status: str prompt_tokens: np.ndarray = dataclasses.field( default_factory=lambda: np.zeros(0, dtype=np.int32) @@ -322,6 +403,7 @@ class RolloutResponse(Response): segments: list[TokenSegment] = dataclasses.field(default_factory=list) env_reward: float = 0.0 policy_version: int = 0 + policy_versions: list[int] = dataclasses.field(default_factory=list) # TODO(b/532722981): capture rollout metrics, e.g., env time. @classmethod @@ -331,6 +413,9 @@ def from_trajectory( traj: Trajectory, prompt_tokens: np.ndarray, policy_version: int, + prompt_id: str = "", + group_offset_id: str = "", + trajectory_id: str = "", ) -> "RolloutResponse": """Constructs a wire-safe RolloutResponse from an internal Trajectory. @@ -342,6 +427,9 @@ def from_trajectory( traj: The internal trajectory to convert. prompt_tokens: Array of prompt token ids. policy_version: Weight version used to generate the trajectory. + prompt_id: Optional prompt identifier. + group_offset_id: Optional string offset identifier within group. + trajectory_id: Optional semantic trajectory identifier. Returns: A wire-safe RolloutResponse. @@ -356,7 +444,11 @@ def _get_step_attr(step, attr): return None segments = [] + step_policy_versions = [] for step in traj.steps: + step_pv = _get_step_attr(step, "policy_version") + step_pv = int(step_pv) if step_pv is not None else policy_version + step_policy_versions.append(step_pv) assistant_tokens = _get_step_attr(step, "assistant_tokens") if assistant_tokens is not None: segments.append( @@ -365,6 +457,7 @@ def _get_step_attr(step, attr): tokens=assistant_tokens, loss_mask=_get_step_attr(step, "assistant_masks"), logps=_get_step_attr(step, "logprobs"), + policy_version=step_pv, ) ) env_tokens = _get_step_attr(step, "env_tokens") @@ -375,19 +468,29 @@ def _get_step_attr(step, attr): tokens=env_tokens, loss_mask=_get_step_attr(step, "env_masks"), logps=None, + policy_version=step_pv, ) ) if hasattr(traj, "status") and traj.status is not None: status_val = getattr(traj.status, "name", str(traj.status)) else: status_val = "COMPLETED" + traj_id_val = ( + trajectory_id + or getattr(traj, "trajectory_id", "") + or request_id + ) return cls( request_id=request_id, + prompt_id=prompt_id or getattr(traj, "task", "") or "", + group_offset_id=group_offset_id, + trajectory_id=traj_id_val, status=status_val, prompt_tokens=prompt_tokens, segments=segments, env_reward=getattr(traj, "reward", 0.0) or 0.0, policy_version=policy_version, + policy_versions=step_policy_versions or [policy_version], ) @@ -524,6 +627,8 @@ class RLTrainerPayload(TrainerPayload): sampler_is_weights: ArrayLike | None = None returns: ArrayLike | None = None old_values: ArrayLike | None = None + trajectory_ids: list[str] = dataclasses.field(default_factory=list) + segment_lineage: dict[int, str] = dataclasses.field(default_factory=dict) metadata: dict[str, Any] = dataclasses.field(default_factory=dict) # TODO: add ppo sepcific fields in a PPO specific fields in PPORLTrainerPayload diff --git a/tunix/experimental/common/test_utils.py b/tunix/experimental/common/test_utils.py index a4b1d7a5b..db670f927 100644 --- a/tunix/experimental/common/test_utils.py +++ b/tunix/experimental/common/test_utils.py @@ -361,9 +361,8 @@ async def collect_rollout_batch( fanned_out_requests = [] for req in requests: for g_idx in range(group_size): - gid = str(g_idx) if group_size > 1 else req.group_offset_id fanned_out_requests.append( - dataclasses.replace(req, group_offset_id=gid) + dataclasses.replace(req, group_offset_id=str(g_idx)) ) tasks: List[Tuple[str, str, Sequence[Any], Dict[str, Any]]] = [ (req.request_id or "req", "generate", (req,), {}) diff --git a/tunix/experimental/orchestrator/algorithm_adapter.py b/tunix/experimental/orchestrator/algorithm_adapter.py index 42cd1c5cd..4a1f7b868 100644 --- a/tunix/experimental/orchestrator/algorithm_adapter.py +++ b/tunix/experimental/orchestrator/algorithm_adapter.py @@ -137,6 +137,15 @@ def create_trainer_payloads( seq_loss_mask = np.concatenate([np.zeros(len(p_arr), dtype=np.float32), act_arr]) seq_adv = np.full(len(seq_tokens), adv_val, dtype=np.float32) + traj_id = ( + getattr(item, "trajectory_id", "") + or ( + f"traj_{item.prompt_id}_{item.group_offset_id}" + if getattr(item, "prompt_id", "") and getattr(item, "group_offset_id", "") + else (f"traj_{item.prompt_id}" if getattr(item, "prompt_id", "") else f"traj_{i}") + ) + ) + payload = datatypes.RLTrainerPayload( token_ids=seq_tokens, token_mask=np.ones_like(seq_tokens, dtype=np.float32), @@ -148,6 +157,7 @@ def create_trainer_payloads( completion_ids=c_arr, completion_mask=act_arr, ref_per_token_logps=np.asarray(ref_lp, dtype=np.float32) if ref_lp is not None else None, + trajectory_ids=[traj_id], ) payloads.append(payload) return payloads @@ -238,6 +248,15 @@ def create_trainer_payloads( seq_loss_mask = np.concatenate([np.zeros(len(p_arr), dtype=np.float32), act_arr]) seq_adv = np.full(len(seq_tokens), adv_val, dtype=np.float32) + traj_id = ( + getattr(item, "trajectory_id", "") + or ( + f"traj_{getattr(item, 'prompt_id', '')}_{getattr(item, 'group_offset_id', str(i))}" + if getattr(item, "prompt_id", "") and getattr(item, "group_offset_id", "") + else (f"traj_{getattr(item, 'prompt_id', '')}" if getattr(item, 'prompt_id', '') else f"traj_{i}") + ) + ) + payload = datatypes.RLTrainerPayload( token_ids=seq_tokens, token_mask=np.ones_like(seq_tokens, dtype=np.float32), @@ -251,6 +270,7 @@ def create_trainer_payloads( old_per_token_logps=np.asarray(old_lp, dtype=np.float32) if old_lp is not None else None, ref_per_token_logps=np.asarray(ref_lp, dtype=np.float32) if ref_lp is not None else None, returns=np.full(len(seq_tokens), vt_val, dtype=np.float32), + trajectory_ids=[traj_id], ) payloads.append(payload) return payloads diff --git a/tunix/experimental/orchestrator/async_rl_program.py b/tunix/experimental/orchestrator/async_rl_program.py index 2364eb7f6..fed8eb4e4 100644 --- a/tunix/experimental/orchestrator/async_rl_program.py +++ b/tunix/experimental/orchestrator/async_rl_program.py @@ -97,22 +97,22 @@ async def rollout_dispatch_stage( for prompt_idx, prompt_item in enumerate(self.dataset): # TODO: Extract prompt_id and group_id from standard tunix data structures # rather than assuming dictionaries or falling back to index strings. - # TODO: the logic of creating group id and prompt id is incorrect and should be fixed. prompt_id = getattr(prompt_item, "prompt_id", f"prompt_{prompt_idx}") - group_id = getattr(prompt_item, "group_id", f"group_{prompt_idx}") if isinstance(prompt_item, dict): prompt_id = prompt_item.get("prompt_id", prompt_id) - group_id = prompt_item.get("group_id", group_id) for g_idx in range(self.group_size): + traj_id = f"traj_{prompt_id}_{g_idx}" await engine.dispatch_rollouts( [prompt_item], - request_id=f"req_{prompt_idx}_{g_idx}", + request_id=traj_id, policy_version=self.policy_version, prompt_ids=[prompt_id], metadata={ - "group_id": group_id, + "group_id": prompt_id, "pair_index": g_idx, + "group_offset_id": str(g_idx), + "trajectory_id": traj_id, }, ) @@ -129,9 +129,6 @@ async def polling_stage( except asyncio.CancelledError: break - except Exception as exc: # pylint: disable=broad-exception-caught - logging.warning("Error in polling_stage: %s", exc) - await asyncio.sleep(0.01) async def critique_stage( self, engine: rl_engine_interface.AbstractRLEngine @@ -162,6 +159,7 @@ async def critique_stage( trainer_payloads = self.algo.create_trainer_payloads( group, rewards=rewards, ref_logps=ref_logps ) + for idx, payload in enumerate(trainer_payloads): adv = payload.advantages reward_val = ( @@ -169,11 +167,17 @@ async def critique_stage( if hasattr(adv, "__len__") and len(adv) > 0 # pyrefly: ignore[bad-argument-type] else float(adv) # pyrefly: ignore[bad-argument-type] ) + traj_id = ( + payload.metadata.get("trajectory_id", "") + or (payload.trajectory_ids[0] if payload.trajectory_ids else "") + or f"traj_{getattr(group[0], 'prompt_id', 'p')}_{idx}" + ) item = datatypes.TrajectoryItem( - pair_index=idx, - group_id=getattr(group[0], "group_id", "default"), + prompt_id=getattr(group[0], "prompt_id", "default"), + group_offset_id=str(idx), + trajectory_id=traj_id, start_step=0, - traj=datatypes.Trajectory(reward=reward_val), + traj=datatypes.Trajectory(trajectory_id=traj_id, reward=reward_val), # TODO: Stream RLTrainerPayload directly instead of re-wrapping in TrajectoryItem. ) item.payload = payload # pyrefly: ignore[missing-attribute] diff --git a/tunix/experimental/orchestrator/batch_assembly.py b/tunix/experimental/orchestrator/batch_assembly.py index cb7320f0c..e299ae54e 100644 --- a/tunix/experimental/orchestrator/batch_assembly.py +++ b/tunix/experimental/orchestrator/batch_assembly.py @@ -152,7 +152,13 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL all_old_logprobs = [] all_ref_logprobs = [] + segment_lineage = {} + trajectory_ids = [] for seg_idx, it in enumerate(b_items, start=1): + traj_id = (it.trajectory_ids[0] if it.trajectory_ids else "") or f"seg_{seg_idx}" + segment_lineage[seg_idx] = traj_id + trajectory_ids.append(traj_id) + toks = ( np.asarray(it.token_ids, dtype=np.int32).reshape(-1) if it.token_ids is not None @@ -233,6 +239,8 @@ 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, :], + trajectory_ids=trajectory_ids, + segment_lineage=segment_lineage, ) payloads.append(payload) @@ -288,7 +296,11 @@ def _pack_chunk( has_ref_logps = any(x.ref_per_token_logps is not None for x in chunk) has_old_logps = any(x.old_per_token_logps is not None for x in chunk) + trajectory_ids = [] for item in chunk: + traj_id = item.trajectory_ids[0] if item.trajectory_ids else "" + trajectory_ids.append(traj_id) + p = np.asarray(item.prompt_ids, dtype=np.int32).reshape(-1) c_full = np.asarray(item.completion_ids, dtype=np.int32).reshape(-1) c_mask_src = ( @@ -355,6 +367,7 @@ def _pack_chunk( ) while len(prompt_ids) < self.batch_size: + trajectory_ids.append("__pad__") prompt_ids.append(np.full(self.max_prompt_length, self.pad_id, np.int32)) prompt_mask.append(np.zeros(self.max_prompt_length, dtype=np.float32)) completion_ids.append( @@ -377,27 +390,33 @@ def _pack_chunk( advantages=jnp.stack(advantages), ref_per_token_logps=jnp.stack(ref_logps) if has_ref_logps else None, old_per_token_logps=jnp.stack(old_logps) if has_old_logps else None, + trajectory_ids=tuple(trajectory_ids), ) -class PaddedBatchAssembler: - """Simple 2D Rectangular Batching: Pads sequences to standard [batch_size, max_seq_len] tensors.""" +class PaddedBatchAssembler(BatchAssembler): + """Pads payloads with uniform maximum sequence lengths into 2D rectangular batches.""" - def __init__(self, batch_size: int = 4, max_seq_len: int = 2048, pad_id: int = 0): - self.batch_size = batch_size + def __init__( + self, + max_seq_len: int, + batch_size: int = 1, + pad_id: int = 0, + ): self.max_seq_len = max_seq_len + self.batch_size = batch_size self.pad_id = pad_id - def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RLTrainerPayload]: - """Pads items into rectangular 2D batches [B, max_seq_len].""" - if not items: + def pack( + self, payloads: Sequence[datatypes.RLTrainerPayload] + ) -> list[datatypes.RLTrainerPayload]: + """Pads rows into uniform rectangular matrices with lineage retention.""" + if not payloads: return [] - item_list = list(items) - payloads: list[datatypes.RLTrainerPayload] = [] - - for i in range(0, len(item_list), self.batch_size): - chunk = item_list[i : i + self.batch_size] + result_payloads = [] + for chunk_start in range(0, len(payloads), self.batch_size): + chunk = payloads[chunk_start : chunk_start + self.batch_size] b_tokens = [] b_loss_masks = [] @@ -405,8 +424,12 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL b_advs = [] b_old_lps = [] b_ref_lps = [] + b_traj_ids = [] for it in chunk: + traj_id = it.trajectory_ids[0] if it.trajectory_ids else "" + b_traj_ids.append(traj_id) + toks = ( np.asarray(it.token_ids, dtype=np.int32).reshape(-1) if it.token_ids is not None @@ -450,6 +473,7 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL # Pad rows up to batch_size while len(b_tokens) < self.batch_size: + b_traj_ids.append("__pad__") b_tokens.append(np.full(self.max_seq_len, self.pad_id, dtype=np.int32)) b_loss_masks.append(np.zeros(self.max_seq_len, dtype=np.float32)) b_action_masks.append(np.zeros(self.max_seq_len, dtype=np.float32)) @@ -467,7 +491,8 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL action_mask=np.stack(b_action_masks), old_per_token_logps=np.stack(b_old_lps) if b_old_lps else None, ref_per_token_logps=np.stack(b_ref_lps) if b_ref_lps else None, + trajectory_ids=b_traj_ids, ) - payloads.append(payload) + result_payloads.append(payload) - return payloads + return result_payloads diff --git a/tunix/experimental/orchestrator/distributed_rl_engine.py b/tunix/experimental/orchestrator/distributed_rl_engine.py index 6edce624f..56953202c 100644 --- a/tunix/experimental/orchestrator/distributed_rl_engine.py +++ b/tunix/experimental/orchestrator/distributed_rl_engine.py @@ -40,11 +40,21 @@ def _response_to_trajectory_item(resp: Any) -> datatypes.TrajectoryItem: if isinstance(resp, datatypes.RolloutResponse): prompt_id = resp.prompt_id or "default_prompt" - metadata = dict(resp.metadata) if resp.metadata else {} - group_id = metadata.get("group_id", prompt_id) - pair_index = metadata.get("pair_index", 0) + metadata = dict(resp.metadata) if getattr(resp, "metadata", None) else {} success_statuses = {"COMPLETED", "SUCCEEDED"} + group_offset_id = ( + resp.group_offset_id + if resp.group_offset_id + else str(metadata.get("pair_index", metadata.get("group_offset_id", ""))) + ) + traj_id = ( + resp.trajectory_id + or metadata.get("trajectory_id", "") + or resp.request_id + or (f"traj_{prompt_id}_{group_offset_id}" if group_offset_id else f"traj_{prompt_id}") + ) traj = datatypes.Trajectory( + trajectory_id=traj_id, reward=resp.env_reward, status=( datatypes.TrajectoryStatus.SUCCEEDED @@ -52,14 +62,22 @@ def _response_to_trajectory_item(resp: Any) -> datatypes.TrajectoryItem: else datatypes.TrajectoryStatus.FAILED ), ) + policy_versions = getattr(resp, "policy_versions", []) + if not policy_versions and getattr(resp, "segments", None): + policy_versions = [ + seg.policy_version for seg in resp.segments if seg.source == "assistant" + ] or [resp.policy_version] + item = datatypes.TrajectoryItem( - pair_index=pair_index, - group_id=group_id, + prompt_id=prompt_id, + group_offset_id=group_offset_id, + trajectory_id=traj_id, start_step=0, traj=traj, metadata=metadata, prompt_tokens=resp.prompt_tokens, policy_version=resp.policy_version, + policy_versions=policy_versions, ) assistant_tokens = [] @@ -77,9 +95,12 @@ def _response_to_trajectory_item(resp: Any) -> datatypes.TrajectoryItem: return item if isinstance(resp, datatypes.Trajectory): + traj_id = getattr(resp, "trajectory_id", "") or "default_traj" + task_id = getattr(resp, "task", "") or "default_group" item = datatypes.TrajectoryItem( - pair_index=0, - group_id=getattr(resp, "task", "default_group"), + prompt_id=task_id, + group_offset_id="0", + trajectory_id=traj_id, start_step=0, traj=resp, policy_version=getattr(resp, "policy_version", 0), diff --git a/tunix/experimental/orchestrator/rl_program.py b/tunix/experimental/orchestrator/rl_program.py index 82a99be17..82cb828d9 100644 --- a/tunix/experimental/orchestrator/rl_program.py +++ b/tunix/experimental/orchestrator/rl_program.py @@ -75,6 +75,8 @@ class RLStepResult: reward_mean: float reward_std: float train_result: Any + global_batch_id: str = "" + trajectory_ids: tuple[str, ...] = () def _default_reward(item: datatypes.TrajectoryItem) -> float: @@ -216,9 +218,17 @@ async def astep_once( else: self.policy_version = current_step + 1 + global_batch_id = f"batch_{current_step:06d}" + all_step_trajectory_ids = tuple( + getattr(r, "trajectory_id", "") or f"traj_{idx}" + for idx, r in enumerate(rollouts) + ) + self.last_step_result = RLStepResult( step=current_step, policy_version=self.policy_version, + global_batch_id=global_batch_id, + trajectory_ids=all_step_trajectory_ids, num_rollouts=len(rollouts), num_microbatches=len(microbatches), reward_mean=float(np.mean(rewards)) if rewards else 0.0, diff --git a/tunix/experimental/rollout/collector.py b/tunix/experimental/rollout/collector.py index a7da1eba8..becb4a7bf 100644 --- a/tunix/experimental/rollout/collector.py +++ b/tunix/experimental/rollout/collector.py @@ -129,10 +129,17 @@ def _convert_to_trajectory(self, rl_traj: Any) -> trajectory_lib.Trajectory: for step in getattr(rl_traj, "steps", []) if getattr(step, "model_response", "") ) + try: + pair_index = int(self.request.group_offset_id) if self.request.group_offset_id else 0 + except (ValueError, TypeError): + pair_index = 0 + metadata.setdefault("text", assistant_text) - metadata.setdefault("prompt_id", self.request.prompt_id) - metadata.setdefault("group_id", self.request.prompt_id) - metadata.setdefault("pair_index", self.request.group_offset_id or 0) + metadata.setdefault("prompt_id", str(self.request.prompt_id or "")) + metadata.setdefault("group_id", str(self.request.prompt_id or "")) + metadata.setdefault("pair_index", pair_index) + metadata.setdefault("group_offset_id", str(self.request.group_offset_id or "")) + metadata.setdefault("trajectory_id", str(self.traj_id)) metadata["prompt_tokens"] = np.asarray( getattr(rl_traj, "prompt_tokens", np.zeros(0, dtype=np.int32)), dtype=np.int32, diff --git a/tunix/experimental/worker/examples/agentic_remote_execution_demo.py b/tunix/experimental/worker/examples/agentic_remote_execution_demo.py index ab25fe78d..2fd3fe937 100644 --- a/tunix/experimental/worker/examples/agentic_remote_execution_demo.py +++ b/tunix/experimental/worker/examples/agentic_remote_execution_demo.py @@ -86,7 +86,7 @@ def _create_agent_env_pair( env_kwargs = { "task": task_data, - "group_id": request.group_offset_id, + "group_id": request.prompt_id, **request.metadata.get("env_kwargs", {}), } env = env_cls(**env_kwargs) @@ -139,13 +139,13 @@ async def run_orchestrator_node( request = datatypes.RolloutRequest( request_id="req_group4_pair0", prompt=single_example, - group_offset_id="group_4", + prompt_id="group_4", + group_offset_id="0", target_policy_version=1, metadata={ "agent_type": "diagnostic", "env_type": "k8s", "system_prompt": "You are an expert K8s agent.", - "pair_index": 0, }, ) diff --git a/tunix/experimental/worker/examples/rl_loop_remote_execution_demo.py b/tunix/experimental/worker/examples/rl_loop_remote_execution_demo.py index 6cfd05846..57e30c033 100644 --- a/tunix/experimental/worker/examples/rl_loop_remote_execution_demo.py +++ b/tunix/experimental/worker/examples/rl_loop_remote_execution_demo.py @@ -98,7 +98,7 @@ def _create_agent_env_pair( env_kwargs = { "task": task_data, - "group_id": request.group_offset_id, + "group_id": request.prompt_id, **request.metadata.get("env_kwargs", {}), } env = env_cls(**env_kwargs) diff --git a/tunix/experimental/worker/rollout_worker.py b/tunix/experimental/worker/rollout_worker.py index d8e4d401b..0fae08d82 100644 --- a/tunix/experimental/worker/rollout_worker.py +++ b/tunix/experimental/worker/rollout_worker.py @@ -304,6 +304,8 @@ def _to_rollout_response( traj=item, # pyrefly: ignore[bad-argument-type] prompt_tokens=prompt_tokens, policy_version=policy_version, + prompt_id=getattr(item, "task", "") or "", + trajectory_id=getattr(item, "trajectory_id", "") or req_id, ) response.prompt_id = str(extra.get("prompt_id", response.prompt_id)) response.env_reward = float(extra.get("reward", response.env_reward)) @@ -344,6 +346,8 @@ def _sampling_to_rollout_response( return datatypes.RolloutResponse( request_id=request.request_id or request.traj_id, prompt_id=request.prompt_id, + group_offset_id=request.group_offset_id, + trajectory_id=request.traj_id, status="COMPLETED", prompt_tokens=prompt_token_arr, segments=[ diff --git a/tunix/rl/agentic/agents/agent_types.py b/tunix/rl/agentic/agents/agent_types.py index d57118c4d..821733d88 100644 --- a/tunix/rl/agentic/agents/agent_types.py +++ b/tunix/rl/agentic/agents/agent_types.py @@ -83,6 +83,7 @@ class Step: env_tokens: Optional[np.ndarray] = None env_masks: Optional[np.ndarray] = None logprobs: Optional[np.ndarray] = None + policy_version: int = 0 class TrajectoryStatus(Enum): @@ -119,6 +120,7 @@ class Trajectory: env_time: Dictionary of environment latency metrics (reset_latency: float, step_latency: list[float] ordered by step index, close_latency: float). """ + trajectory_id: str = "" task: Any = None steps: list[Step] = dataclasses.field(default_factory=list) reward: float = 0.0 @@ -136,6 +138,7 @@ def to_dict(self) -> dict[str, Any]: dict: Serializable dictionary representation of the trajectory. """ return { + "trajectory_id": self.trajectory_id, "task": self.task, "steps": [dataclasses.asdict(step) for step in self.steps], "reward": float(self.reward), diff --git a/tunix/rl/common.py b/tunix/rl/common.py index a1d04765d..c330fd41e 100644 --- a/tunix/rl/common.py +++ b/tunix/rl/common.py @@ -120,6 +120,13 @@ class TrainExample: # to dampen positions where the trainer's recomputed log-probability # diverges from the rollout sampler's. ``None`` disables the correction. sampler_is_weights: jax.Array | None = None + # Lineage tracking metadata (non-pytree fields ignored by JAX JIT). + trajectory_ids: tuple[str, ...] | None = flax.struct.field( + default=None, pytree_node=False + ) + segment_lineage: dict[int, str] | None = flax.struct.field( + default=None, pytree_node=False + ) def compute_kl_divergence(