From 500742bf075de9db92f2af3430e4aff07186e2e0 Mon Sep 17 00:00:00 2001 From: Haoyu Gao Date: Wed, 12 Aug 2026 13:18:40 -0700 Subject: [PATCH] Introduce experimental RL orchestrator and remote worker services for distributed training. PiperOrigin-RevId: 963623649 --- .../orchestrator/batch_assembly_test.py | 36 +++ .../distributed_rl_engine_test.py | 101 ++++++- .../orchestrator/rl_program_test.py | 79 +++++- .../legacy_vllm_sampler_adapter_test.py | 3 + tests/experimental/rollout/rollout_test.py | 1 - tests/experimental/rollout/sampler_test.py | 4 + .../rollout/vanilla_sampler_adapter_test.py | 5 + .../worker/abstract_worker_test.py | 6 +- tunix/experimental/common/datatypes.py | 55 ++-- tunix/experimental/common/test_utils.py | 56 ++-- .../orchestrator/algorithm_adapter.py | 9 + .../orchestrator/batch_assembly.py | 213 +++++++++++++- .../orchestrator/distributed_rl_engine.py | 71 ++++- .../experimental/orchestrator/orchestrator.py | 2 +- .../orchestrator/rl_engine_interface.py | 10 +- tunix/experimental/orchestrator/rl_program.py | 136 +++++++-- .../rollout/legacy_vllm_sampler_adapter.py | 32 +++ tunix/experimental/rollout/sampler.py | 9 + .../rollout/vanilla_sampler_adapter.py | 23 ++ tunix/experimental/train/peft_trainer_v2.py | 6 +- tunix/experimental/worker/inference_worker.py | 114 ++++++++ tunix/experimental/worker/remote_execution.py | 11 +- tunix/experimental/worker/rollout_worker.py | 267 +++++++++++++++++- tunix/experimental/worker/trainer_worker.py | 173 ++++++++++-- tunix/rl/rl_cluster.py | 21 ++ tunix/tests/test_common.py | 16 ++ 26 files changed, 1324 insertions(+), 135 deletions(-) diff --git a/tests/experimental/orchestrator/batch_assembly_test.py b/tests/experimental/orchestrator/batch_assembly_test.py index 7391ead20..2969a0290 100644 --- a/tests/experimental/orchestrator/batch_assembly_test.py +++ b/tests/experimental/orchestrator/batch_assembly_test.py @@ -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() diff --git a/tests/experimental/orchestrator/distributed_rl_engine_test.py b/tests/experimental/orchestrator/distributed_rl_engine_test.py index 1b22da4b7..93aef8af7 100644 --- a/tests/experimental/orchestrator/distributed_rl_engine_test.py +++ b/tests/experimental/orchestrator/distributed_rl_engine_test.py @@ -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() @@ -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( @@ -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()) @@ -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(): diff --git a/tests/experimental/orchestrator/rl_program_test.py b/tests/experimental/orchestrator/rl_program_test.py index 0ccfad167..37fe6c828 100644 --- a/tests/experimental/orchestrator/rl_program_test.py +++ b/tests/experimental/orchestrator/rl_program_test.py @@ -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) @@ -76,10 +80,12 @@ 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) @@ -87,6 +93,44 @@ def on_end(step, result): 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( @@ -94,10 +138,15 @@ def test_eval_step_once_flow(self): 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() diff --git a/tests/experimental/rollout/legacy_vllm_sampler_adapter_test.py b/tests/experimental/rollout/legacy_vllm_sampler_adapter_test.py index a88cc84fd..9734d60c0 100644 --- a/tests/experimental/rollout/legacy_vllm_sampler_adapter_test.py +++ b/tests/experimental/rollout/legacy_vllm_sampler_adapter_test.py @@ -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() @@ -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"} diff --git a/tests/experimental/rollout/rollout_test.py b/tests/experimental/rollout/rollout_test.py index 0c8c3dfd4..725940101 100644 --- a/tests/experimental/rollout/rollout_test.py +++ b/tests/experimental/rollout/rollout_test.py @@ -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() diff --git a/tests/experimental/rollout/sampler_test.py b/tests/experimental/rollout/sampler_test.py index 2a6e74e0b..799e4f17f 100644 --- a/tests/experimental/rollout/sampler_test.py +++ b/tests/experimental/rollout/sampler_test.py @@ -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", @@ -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 diff --git a/tests/experimental/rollout/vanilla_sampler_adapter_test.py b/tests/experimental/rollout/vanilla_sampler_adapter_test.py index 439a28294..8e97cc97a 100644 --- a/tests/experimental/rollout/vanilla_sampler_adapter_test.py +++ b/tests/experimental/rollout/vanilla_sampler_adapter_test.py @@ -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 @@ -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( @@ -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( @@ -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( diff --git a/tests/experimental/worker/abstract_worker_test.py b/tests/experimental/worker/abstract_worker_test.py index 5faf0fbec..3fa05c28e 100644 --- a/tests/experimental/worker/abstract_worker_test.py +++ b/tests/experimental/worker/abstract_worker_test.py @@ -88,7 +88,8 @@ 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 @@ -96,7 +97,8 @@ def 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 diff --git a/tunix/experimental/common/datatypes.py b/tunix/experimental/common/datatypes.py index 234eed583..3e93e89cb 100644 --- a/tunix/experimental/common/datatypes.py +++ b/tunix/experimental/common/datatypes.py @@ -30,17 +30,6 @@ ##### Worker-internal datatypes ##### -# TODO(noghabi): Consolidate Role with rl_cluster.Role. -class Role(enum.Enum): - """Role of the model.""" - - ACTOR = "actor" # policy model - CRITIC = "critic" # value model (only for PPO-style algos, not for GRPO) - REFERENCE = "reference" # kept fixed during training - REWARD = "reward" - ROLLOUT = "rollout" - - # Worker-internal episode representation produced during rollout. Trajectory = agent_types.Trajectory Step = agent_types.Step @@ -56,7 +45,6 @@ class TrajectoryItem(agent_types.TrajectoryItem): completion_tokens: np.ndarray | None = None action_mask: np.ndarray | None = None policy_version: int = 0 - # TODO: trajectory item having the completion tokens, masks, etc is quite redundant since those are in the trainer payload already. class Role(str, enum.Enum): """Orchestrator worker roles.""" @@ -219,6 +207,24 @@ class WorkerInfo: ##### Rollout DTOs ##### +@dataclasses.dataclass(frozen=True, kw_only=True) +class GenerationArgs: + """Typed generation arguments used by the orchestrator generate API.""" + max_generation_steps: int | None = None + temperature: float | None = None + top_p: float | None = None + top_k: int | None = None + seed: int | None = None + return_logprobs: bool | None = None + + def as_kwargs(self) -> dict[str, Any]: + return { + field.name: getattr(self, field.name) + for field in dataclasses.fields(self) + if getattr(self, field.name) is not None + } + + @dataclasses.dataclass(kw_only=True) class RolloutRequest(Request): """Request to generate a rollout from a given prompt. @@ -468,8 +474,9 @@ class TrainerPayload: """Generic trainer payload. Attributes: - token_ids: [B, T] token IDs. By default, structured as left-padded prompt - tokens concatenated with right-padded completion tokens. + token_ids: [B, T] token IDs for a batched trainer payload. By default, + each row is structured as left-padded prompt tokens concatenated with + right-padded completion tokens. token_mask: [B, T] token mask to differentiate padding tokens from valid tokens. segment_ids: Optional [B, T] packing segment ids. @@ -491,6 +498,12 @@ class RLTrainerPayload(TrainerPayload): advantages: [B] or [B, C] advantages. loss_mask: [B, T], 1 where the position contributes to the loss. action_mask: Optional [B, T] or [B, C] mask of policy actions. + prompt_ids: Optional prompt token ids for GRPO-style losses. Unbatched + payloads may carry 1D unpadded rows; batch assembly pads them to [B, P]. + prompt_mask: Optional [B, P] prompt mask. + completion_ids: Optional completion token ids. Unbatched payloads may carry + 1D unpadded rows; batch assembly pads them to [B, C]. + completion_mask: Optional [B, C] completion/action mask. ref_per_token_logps: Optional [B, C] reference model log-probabilities. old_per_token_logps: Optional [B, C] behavior policy log-probabilities. sampler_is_weights: Optional [B, C] importance sampling weights. @@ -502,6 +515,10 @@ class RLTrainerPayload(TrainerPayload): advantages: ArrayLike loss_mask: ArrayLike action_mask: ArrayLike | None = None + prompt_ids: ArrayLike | None = None + prompt_mask: ArrayLike | None = None + completion_ids: ArrayLike | None = None + completion_mask: ArrayLike | None = None ref_per_token_logps: ArrayLike | None = None old_per_token_logps: ArrayLike | None = None sampler_is_weights: ArrayLike | None = None @@ -516,9 +533,9 @@ class LogprobsRequest(Request): """Request to score per-token log-probabilities under a frozen model. Attributes: - prompt_tokens: [B, P], LEFT-padded. - completion_tokens: [B, C], RIGHT-padded; the result aligns to these - completion columns. + prompt_tokens: [B, P] token ids, already LEFT-padded by the caller. + completion_tokens: [B, C] token ids, already RIGHT-padded by the caller; + the result aligns to these completion columns. temperature: Softmax temperature to score under. Mandatory: it must match the temperature the tokens were sampled at, or the log-probs are biased. model_role: Which hosted model to score against (v1: "reference"). @@ -552,8 +569,8 @@ class ScoreRequest(Request): """Request to score scalar rewards/values under a hosted model. Attributes: - prompt_tokens: [B, P], LEFT-padded. - completion_tokens: [B, C], RIGHT-padded. + prompt_tokens: [B, P] token ids, already LEFT-padded by the caller. + completion_tokens: [B, C] token ids, already RIGHT-padded by the caller. model_role: Which hosted model to score against (e.g. "reward"). """ diff --git a/tunix/experimental/common/test_utils.py b/tunix/experimental/common/test_utils.py index 634f9a771..a4b1d7a5b 100644 --- a/tunix/experimental/common/test_utils.py +++ b/tunix/experimental/common/test_utils.py @@ -157,6 +157,9 @@ def __init__( self.tokenizer = MockTokenizer() self.chat_parser = MockChatParser() + def initialize(self) -> None: + pass + async def sample( self, sampling_requests: ( @@ -172,27 +175,38 @@ async def sample( delay = kwargs.get("delay_seconds", self.default_delay) await asyncio.sleep(delay) - req_id_str = "default" - if ( - hasattr(sampling_requests, "request_id") - and sampling_requests.request_id - ): - req_id_str = str(sampling_requests.request_id) - turn = self._turn_counters.get(req_id_str, 0) - self._turn_counters[req_id_str] = turn + 1 - - min_turns = kwargs.get("min_turns", 1) - if turn >= min_turns or kwargs.get("force_finish", False): - ans = kwargs.get("answer", f"result_for_{req_id_str}") - txt = f"FINAL_ANSWER: [{self.sampler_name}] {ans}" - else: - txt = f"TOOL_CALL: search(query='turn {turn} for {req_id_str}')" - tokens = np.array([101, 102], dtype=np.int32) - return base_sampler_lib.SamplingResponse( - text=txt, - token_ids=tokens, - logprobs=np.zeros_like(tokens, dtype=np.float32), - ) + is_sequence = isinstance(sampling_requests, (list, tuple)) + requests = list(sampling_requests) if is_sequence else [sampling_requests] + responses = [] + for req in requests: + request_kwargs = dict(kwargs) + request_kwargs.update(getattr(req, "metadata", {}) or {}) + req_id_str = "default" + if hasattr(req, "request_id") and req.request_id: + req_id_str = str(req.request_id) + prompt = getattr(req, "prompt", req) + turn = self._turn_counters.get(req_id_str, 0) + self._turn_counters[req_id_str] = turn + 1 + + min_turns = request_kwargs.get("min_turns", 1) + if turn >= min_turns or request_kwargs.get("force_finish", False): + ans = request_kwargs.get("answer", f"result_for_{req_id_str}") + txt = f"FINAL_ANSWER: [{self.sampler_name}] {ans}" + else: + txt = f"TOOL_CALL: search(query='turn {turn} for {req_id_str}')" + tokens = np.array([101, 102], dtype=np.int32) + responses.append( + base_sampler_lib.SamplingResponse( + request_id=req_id_str, + text=txt, + prompt_token_ids=np.asarray( + self.tokenizer.encode(str(prompt)), dtype=np.int32 + ), + token_ids=tokens, + logprobs=np.zeros_like(tokens, dtype=np.float32), + ) + ) + return responses if is_sequence else responses[0] async def migrate_kv_cache( self, diff --git a/tunix/experimental/orchestrator/algorithm_adapter.py b/tunix/experimental/orchestrator/algorithm_adapter.py index 04b824bd8..44d874a41 100644 --- a/tunix/experimental/orchestrator/algorithm_adapter.py +++ b/tunix/experimental/orchestrator/algorithm_adapter.py @@ -92,6 +92,7 @@ def __init__( ) self.clip_epsilon = clip_epsilon self.beta_kl = beta_kl + self.requires_reference_kl = beta_kl != 0.0 def compute_advantages( self, @@ -142,6 +143,10 @@ def create_trainer_payloads( loss_mask=seq_loss_mask, advantages=seq_adv, action_mask=seq_loss_mask, + prompt_ids=p_arr, + prompt_mask=np.ones(len(p_arr), dtype=np.float32), + 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, ) payloads.append(payload) @@ -239,6 +244,10 @@ def create_trainer_payloads( loss_mask=seq_loss_mask, advantages=seq_adv, action_mask=seq_loss_mask, + prompt_ids=p_arr, + prompt_mask=np.ones(len(p_arr), dtype=np.float32), + completion_ids=c_arr, + completion_mask=act_arr, 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), diff --git a/tunix/experimental/orchestrator/batch_assembly.py b/tunix/experimental/orchestrator/batch_assembly.py index 7173a1da6..cb7320f0c 100644 --- a/tunix/experimental/orchestrator/batch_assembly.py +++ b/tunix/experimental/orchestrator/batch_assembly.py @@ -22,9 +22,11 @@ # TODO: Align SequencePackedBatchAssembler with the rest of the ecosystem and potentially move to a common library. """ -from typing import Generic, Protocol, Sequence, TypeVar +from typing import Any, Generic, Protocol, Sequence, TypeVar import numpy as np +from jax import numpy as jnp from tunix.experimental.common import datatypes +from tunix.rl import common as rl_common T = TypeVar("T") @@ -37,6 +39,74 @@ def pack(self, items: Sequence[T]) -> list[datatypes.RLTrainerPayload]: ... +def _left_pad( + values: np.ndarray, + length: int, + *, + pad_id: int, +) -> tuple[np.ndarray, np.ndarray]: + arr = np.asarray(values, dtype=np.int32).reshape(-1)[-length:] + out = np.full(length, pad_id, dtype=np.int32) + mask = np.zeros(length, dtype=np.float32) + if arr.size: + out[-arr.size:] = arr + mask[-arr.size:] = 1.0 + return out, mask + + +def _right_pad( + values: np.ndarray, + length: int, + *, + pad_value: float | int = 0, + dtype: Any = np.int32, +) -> tuple[np.ndarray, np.ndarray]: + arr = np.asarray(values, dtype=dtype).reshape(-1)[:length] + out = np.full(length, pad_value, dtype=dtype) + mask = np.zeros(length, dtype=np.float32) + if arr.size: + out[:arr.size] = arr + mask[:arr.size] = 1.0 + return out, mask + + +def _completion_aligned( + values: Any | None, + completion_len: int, + max_response_length: int, + *, + fill_value: float = 0.0, + prompt_len: int | None = None, + full_completion_len: int | None = None, +) -> np.ndarray: + if values is None: + arr = np.full(completion_len, fill_value, dtype=np.float32) + else: + arr = np.asarray(values, dtype=np.float32).reshape(-1) + if arr.size == 1: + arr = np.full(completion_len, float(arr[0]), dtype=np.float32) + elif ( + prompt_len is not None + and full_completion_len is not None + and arr.size == prompt_len + full_completion_len + ): + arr = arr[prompt_len : prompt_len + full_completion_len] + elif full_completion_len is not None and arr.size >= full_completion_len: + arr = arr[:full_completion_len] + elif arr.size >= completion_len: + arr = arr[:completion_len] + else: + arr = np.pad(arr, (0, completion_len - arr.size), constant_values=0.0) + arr = arr[:completion_len] + out, _ = _right_pad( + arr, + max_response_length, + pad_value=0.0, + dtype=np.float32, + ) + return out + + class SequencePackedBatchAssembler: """1D Sequence Packing: Concatenates items into dense [1, max_packed_len] buffers.""" # TODO: align implementation with current path. @@ -169,6 +239,147 @@ def pack(self, items: Sequence[datatypes.RLTrainerPayload]) -> list[datatypes.RL return payloads +class GRPOTrainExampleAssembler: + """Pads GRPO payloads into `TrainExample` microbatches. + + Generic assemblers operate on concatenated token streams. GRPO loss consumes + separate prompt and completion tensors, so this assembler keeps those fields + aligned while still implementing the common `BatchAssembler.pack()` contract. + """ + + def __init__( + self, + *, + batch_size: int, + max_prompt_length: int, + max_response_length: int, + pad_id: int, + ): + if batch_size <= 0: + raise ValueError("train microbatch size must be positive.") + self.batch_size = batch_size + self.max_prompt_length = max_prompt_length + self.max_response_length = max_response_length + self.pad_id = pad_id + + def pack( + self, items: Sequence[datatypes.RLTrainerPayload] + ) -> list[rl_common.TrainExample]: + item_list = list(items) + if not item_list: + return [] + + microbatches = [] + for start in range(0, len(item_list), self.batch_size): + chunk = item_list[start : start + self.batch_size] + microbatches.append(self._pack_chunk(chunk)) + return microbatches + + def _pack_chunk( + self, chunk: Sequence[datatypes.RLTrainerPayload] + ) -> rl_common.TrainExample: + prompt_ids = [] + prompt_mask = [] + completion_ids = [] + completion_mask = [] + advantages = [] + ref_logps = [] + old_logps = [] + 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) + + for item in chunk: + 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 = ( + np.asarray(item.completion_mask, dtype=np.float32).reshape(-1) + if item.completion_mask is not None + else np.ones(c_full.shape, dtype=np.float32) + ) + c = c_full[: self.max_response_length] + c_mask_src = c_mask_src[: c.size] + + p_ids, p_mask = _left_pad( + p, self.max_prompt_length, pad_id=self.pad_id + ) + c_ids, c_default_mask = _right_pad( + c, + self.max_response_length, + pad_value=self.pad_id, + dtype=np.int32, + ) + c_mask = np.zeros(self.max_response_length, dtype=np.float32) + if c_mask_src.size: + c_mask[: c_mask_src.size] = c_mask_src + else: + c_mask = c_default_mask + + prompt_ids.append(p_ids) + prompt_mask.append(p_mask) + completion_ids.append(c_ids) + completion_mask.append(c_mask) + + adv_arr = ( + np.asarray(item.advantages, dtype=np.float32).reshape(-1) + if item.advantages is not None + else None + ) + advantages.append( + _completion_aligned( + adv_arr, + c.size, + self.max_response_length, + fill_value=0.0, + prompt_len=p.size, + full_completion_len=c_full.size, + ) + ) + + if has_ref_logps: + ref_logps.append( + _completion_aligned( + item.ref_per_token_logps, + c.size, + self.max_response_length, + full_completion_len=c_full.size, + ) + ) + if has_old_logps: + old_logps.append( + _completion_aligned( + item.old_per_token_logps, + c.size, + self.max_response_length, + full_completion_len=c_full.size, + ) + ) + + while len(prompt_ids) < self.batch_size: + 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( + np.full(self.max_response_length, self.pad_id, np.int32) + ) + completion_mask.append( + np.zeros(self.max_response_length, dtype=np.float32) + ) + advantages.append(np.zeros(self.max_response_length, dtype=np.float32)) + if has_ref_logps: + ref_logps.append(np.zeros(self.max_response_length, dtype=np.float32)) + if has_old_logps: + old_logps.append(np.zeros(self.max_response_length, dtype=np.float32)) + + return rl_common.TrainExample( + prompt_ids=jnp.stack(prompt_ids), + prompt_mask=jnp.stack(prompt_mask), + completion_ids=jnp.stack(completion_ids), + completion_mask=jnp.stack(completion_mask), + 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, + ) + + class PaddedBatchAssembler: """Simple 2D Rectangular Batching: Pads sequences to standard [batch_size, max_seq_len] tensors.""" diff --git a/tunix/experimental/orchestrator/distributed_rl_engine.py b/tunix/experimental/orchestrator/distributed_rl_engine.py index 91736b745..9d0b9eeb2 100644 --- a/tunix/experimental/orchestrator/distributed_rl_engine.py +++ b/tunix/experimental/orchestrator/distributed_rl_engine.py @@ -31,6 +31,7 @@ from tunix.experimental.orchestrator import rl_engine_interface from tunix.experimental.worker import remote_execution + # TODO: this multi step conversions seem excessive we convert from trajecotry to response then to trajectory item. we should simplify def _response_to_trajectory_item(resp: Any) -> datatypes.TrajectoryItem: """Converts a worker rollout response to an TrajectoryItem.""" @@ -203,25 +204,66 @@ async def _poll_worker(worker: remote_execution.ActorHandle) -> Any: completed.append(_response_to_trajectory_item(it)) return completed - async def generate(self, prompts: Sequence[Any], **kwargs: Any) -> list[datatypes.TrajectoryItem]: + async def generate( + self, + prompts: Sequence[Any], + generation_args: datatypes.GenerationArgs | None = None, + route_metadata: Mapping[str, Any] | None = None, + **kwargs: Any, + ) -> list[datatypes.TrajectoryItem]: """Blocking rollout generation: load-balances prompts across workers and awaits completion.""" if not self._rollout_workers: raise ValueError("DistributedRLEngine has no registered rollout workers.") + if kwargs: + raise TypeError( + "Unexpected generate kwargs: " + f"{sorted(kwargs)}. Use generation_args=GenerationArgs(...) for " + "sampling parameters." + ) + + generation_kwargs = ( + generation_args.as_kwargs() if generation_args is not None else {} + ) + route_metadata_map = dict(route_metadata or {}) worker_to_prompts: dict[Any, list[Any]] = collections.defaultdict(list) + worker_to_requests: dict[Any, list[datatypes.RolloutRequest]] = ( + collections.defaultdict(list) + ) for p in prompts: - metadata = dict(kwargs.get("metadata", {})) - route_key = metadata.get("prefix_hash") or metadata.get("prompt_id") + if isinstance(p, datatypes.RolloutRequest): + request_metadata = dict(p.metadata or {}) + route_key = request_metadata.get("prefix_hash") + if route_key is None: + route_key = p.prompt_id + worker = self._rollout_pool._get_next_actor( + kwargs={"route_key": route_key} + ) + worker_to_requests[worker].append(p) + continue + + route_key = route_metadata_map.get("prefix_hash") + if route_key is None: + route_key = route_metadata_map.get("prompt_id") worker = self._rollout_pool._get_next_actor( kwargs={"route_key": route_key} ) worker_to_prompts[worker].append(p) tasks = [] + for worker, w_requests in worker_to_requests.items(): + if w_requests: + tasks.append( + self._invoke_worker( + worker, "generate", requests=w_requests, **generation_kwargs + ) + ) for worker, w_prompts in worker_to_prompts.items(): if w_prompts: tasks.append( - self._invoke_worker(worker, "generate", prompts=w_prompts, **kwargs) + self._invoke_worker( + worker, "generate", prompts=w_prompts, **generation_kwargs + ) ) if not tasks: @@ -279,16 +321,22 @@ async def train_step( worker = self._trainer_workers.get(role) if worker is None: raise ValueError(f"No trainer worker registered for role {role}") - # TODO: we need on apply_optimizer=apply_optimize steps we need to call update() too. - return await self._invoke_worker( + fwd_bwd_result = await self._invoke_worker( worker, "fwd_bwd", - batch=payload, - accumulate_gradients=accumulate_gradients, - apply_optimizer=apply_optimizer, + payload=payload, skip_jit=skip_jit, **kwargs, ) + if not apply_optimizer: + return fwd_bwd_result + train_step = await self._invoke_worker(worker, "update") + return { + "fwd_bwd": fwd_bwd_result, + "updated": True, + "train_step": train_step, + "accumulated": accumulate_gradients, + } async def sync_weights( # pyrefly: ignore[bad-override] self, @@ -302,6 +350,11 @@ async def sync_weights( # pyrefly: ignore[bad-override] if trainer is None: return 0 sync_metadata = await self._invoke_worker(trainer, "prepare_weight_sync") + if not isinstance(sync_metadata, datatypes.WeightSyncMetadata): + raise RuntimeError( + "prepare_weight_sync must return WeightSyncMetadata; got " + f"{type(sync_metadata).__name__}." + ) tasks = [ self._invoke_worker(w, "weight_sync", metadata=sync_metadata) for w in self._rollout_workers diff --git a/tunix/experimental/orchestrator/orchestrator.py b/tunix/experimental/orchestrator/orchestrator.py index f97b17506..a31f29d24 100644 --- a/tunix/experimental/orchestrator/orchestrator.py +++ b/tunix/experimental/orchestrator/orchestrator.py @@ -94,7 +94,7 @@ def validate_startup(self, alg_config: Any, training_config: Any) -> None: def _get_role_members(self, role: datatypes.Role | str) -> list[Any]: role_key = role.value if isinstance(role, datatypes.Role) else role members = self.registry.group(role_key).members() - + # Fallback in case workers were registered with the enum object directly if not members and isinstance(role, datatypes.Role): members = self.registry.group(role).members() diff --git a/tunix/experimental/orchestrator/rl_engine_interface.py b/tunix/experimental/orchestrator/rl_engine_interface.py index e67dd348f..980f0d5e0 100644 --- a/tunix/experimental/orchestrator/rl_engine_interface.py +++ b/tunix/experimental/orchestrator/rl_engine_interface.py @@ -14,7 +14,7 @@ """The RL engine interface (Layer 1 Compute Routing Protocol) following Orchestrator V2.""" -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from typing import Any, Protocol, runtime_checkable from tunix.experimental.common import datatypes @@ -36,7 +36,11 @@ async def poll_rollouts( ... async def generate( - self, prompts: Sequence[Any], **kwargs: Any + self, + prompts: Sequence[Any], + generation_args: datatypes.GenerationArgs | None = None, + route_metadata: Mapping[str, Any] | None = None, + **kwargs: Any, ) -> list[datatypes.TrajectoryItem]: """Synchronous batched rollout generation over rollout workers.""" ... @@ -61,7 +65,7 @@ async def train_step( apply_optimizer: bool = True, skip_jit: bool = False, **kwargs: Any, - ) -> dict[str, Any]: + ) -> Any: """Executes forward/backward gradient update on trainer workers.""" ... diff --git a/tunix/experimental/orchestrator/rl_program.py b/tunix/experimental/orchestrator/rl_program.py index 1207e0316..90fc3a8a5 100644 --- a/tunix/experimental/orchestrator/rl_program.py +++ b/tunix/experimental/orchestrator/rl_program.py @@ -15,7 +15,8 @@ """Synchronous RL Program (rl_program.py) coordinating Engine, Algo, and Assembler.""" import asyncio -from collections.abc import Callable, Iterable, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence +import dataclasses import inspect from typing import Any, Protocol @@ -26,13 +27,16 @@ from tunix.experimental.orchestrator import batch_assembly from tunix.experimental.orchestrator import rl_engine_interface +RewardFn = Callable[[datatypes.TrajectoryItem], float] + class RLProgram(Protocol): """Standard contract for RL training programs running on ClusterOrchestrator.""" + def run( self, - engine: rl_engine_interface.AbstractRLEngine, - train_dataset: Iterable[Any] | None = None, + engine: rl_engine_interface.AbstractRLEngine | None = None, + train_dataset: Iterable[Sequence[datatypes.RolloutRequest]] | None = None, num_steps: int | None = None, **kwargs: Any, ) -> Any: @@ -53,17 +57,37 @@ def _sync_or_async(coro: Any) -> Any: return coro +@dataclasses.dataclass(frozen=True) +class RLStepResult: + """Summary for the most recent synchronous RL step.""" + + step: int + policy_version: int + num_rollouts: int + num_microbatches: int + reward_mean: float + reward_std: float + train_result: Any + + +def _default_reward(item: datatypes.TrajectoryItem) -> float: + if hasattr(item, "env_reward"): + return float(getattr(item, "env_reward", 0.0)) + return 0.0 + + class SyncRLProgram: """Synchronous RL Program coordinating an iterative RL training loop.""" def __init__( self, - engine: rl_engine_interface.AbstractRLEngine, algo: algorithm_adapter.AlgorithmAdapter, - reward_fns: Sequence[Callable[..., Any]] | None = None, + engine: rl_engine_interface.AbstractRLEngine | None = None, + reward_fns: Sequence[RewardFn] | None = None, assembler: batch_assembly.BatchAssembler | None = None, on_step_begin: Callable[[int], None] | None = None, on_step_end: Callable[[int, Any], None] | None = None, + sync_weights: bool = True, ): self.engine = engine self.algo = algo @@ -73,59 +97,110 @@ def __init__( ) self.on_step_begin = on_step_begin self.on_step_end = on_step_end + self.sync_weights = sync_weights self.policy_version = 0 + # Debug/observability hook for examples and tests; not training state. + self.last_step_result: RLStepResult | None = None @property def step(self) -> int: return self.policy_version + def _resolve_engine( + self, engine: rl_engine_interface.AbstractRLEngine | None = None + ) -> rl_engine_interface.AbstractRLEngine: + active_engine = engine or self.engine + if active_engine is None: + raise ValueError( + "SyncRLProgram requires an engine either at construction time or via " + "ClusterOrchestrator.run_program(engine=...)." + ) + return active_engine + def step_once( self, - prompts: list[str] | list[list[dict[str, str]]], + prompts: Sequence[datatypes.RolloutRequest], + generation_args: datatypes.GenerationArgs | None = None, + route_metadata: Mapping[str, Any] | None = None, **kwargs: Any, ) -> Any: """Executes a single end-to-end RL training step.""" + active_engine = self._resolve_engine() current_step = self.policy_version if self.on_step_begin: self.on_step_begin(current_step) # 1. Generate rollouts - rollouts = _sync_or_async(self.engine.generate(prompts=prompts, **kwargs)) + engine_call_kwargs = dict(kwargs) + if generation_args is not None: + engine_call_kwargs["generation_args"] = generation_args + if route_metadata is not None: + engine_call_kwargs["route_metadata"] = route_metadata + rollouts = _sync_or_async( + active_engine.generate(prompts=prompts, **engine_call_kwargs) + ) # 2. Evaluate rewards rewards = [] for item in rollouts: - r = sum(fn(item) for fn in self.reward_fns) if self.reward_fns else getattr(item, "env_reward", 0.0) + r = ( + sum(fn(item) for fn in self.reward_fns) + if self.reward_fns + else _default_reward(item) + ) rewards.append(float(r)) # 3. Create RLTrainerPayloads via AlgorithmAdapter ref_logps = None if getattr(self.algo, "requires_reference_kl", False): - ref_logps = _sync_or_async(self.engine.per_token_logps(datatypes.Role.REFERENCE, items=rollouts)) + ref_logps = _sync_or_async( + active_engine.per_token_logps(datatypes.Role.REFERENCE, items=rollouts) + ) trainer_payloads = self.algo.create_trainer_payloads( rollouts, rewards=rewards, ref_logps=ref_logps ) # 4. Pack into microbatches microbatches = self.assembler.pack(trainer_payloads) + if not microbatches: + raise RuntimeError("No trainer microbatches were assembled.") # 5. Execute gradient updates step_result = None - for batch in microbatches: + for index, batch in enumerate(microbatches): + is_last = index == len(microbatches) - 1 step_result = _sync_or_async( - self.engine.train_step( + active_engine.train_step( batch, role=datatypes.Role.ACTOR, - accumulate_gradients=False, - apply_optimizer=True, + accumulate_gradients=len(microbatches) > 1, + apply_optimizer=is_last, ) ) # 6. Sync weights to rollout replicas - _sync_or_async(self.engine.sync_weights(role=datatypes.Role.ACTOR)) - - # 7. Increment step - self.policy_version = current_step + 1 + if self.sync_weights: + new_version = _sync_or_async( + active_engine.sync_weights(role=datatypes.Role.ACTOR) + ) + if not isinstance(new_version, int) or new_version <= current_step: + raise RuntimeError( + "sync_weights must return a monotonically increasing int policy " + f"version; got {new_version!r} at step {current_step}." + ) + self.policy_version = new_version + else: + self.policy_version = current_step + 1 + + self.last_step_result = RLStepResult( + step=current_step, + policy_version=self.policy_version, + num_rollouts=len(rollouts), + num_microbatches=len(microbatches), + reward_mean=float(np.mean(rewards)) if rewards else 0.0, + reward_std=float(np.std(rewards)) if rewards else 0.0, + train_result=step_result, + ) if self.on_step_end: self.on_step_end(self.policy_version, step_result) @@ -134,24 +209,43 @@ def step_once( def eval_step_once( self, - prompts: list[str] | list[list[dict[str, str]]], + prompts: Sequence[datatypes.RolloutRequest], + generation_args: datatypes.GenerationArgs | None = None, + route_metadata: Mapping[str, Any] | None = None, **kwargs: Any, ) -> list[datatypes.RLTrainerPayload]: """Executes evaluation step without updating weights.""" - rollouts = _sync_or_async(self.engine.generate(prompts=prompts, **kwargs)) + active_engine = self._resolve_engine() + engine_call_kwargs = dict(kwargs) + if generation_args is not None: + engine_call_kwargs["generation_args"] = generation_args + if route_metadata is not None: + engine_call_kwargs["route_metadata"] = route_metadata + rollouts = _sync_or_async( + active_engine.generate(prompts=prompts, **engine_call_kwargs) + ) rewards = [ - sum(fn(item) for fn in self.reward_fns) if self.reward_fns else getattr(item, "env_reward", 0.0) + ( + sum(fn(item) for fn in self.reward_fns) + if self.reward_fns + else _default_reward(item) + ) for item in rollouts ] return self.algo.create_trainer_payloads(rollouts, rewards=rewards) def run( self, - train_dataset: Iterable[list[str] | list[list[dict[str, str]]]], + engine: rl_engine_interface.AbstractRLEngine | None = None, + train_dataset: Iterable[Sequence[datatypes.RolloutRequest]] | None = None, num_steps: int | None = None, **kwargs: Any, ) -> None: """Runs the RL program training loop over the dataset.""" + active_engine = self._resolve_engine(engine) + self.engine = active_engine + if train_dataset is None: + raise ValueError("SyncRLProgram.run requires a train_dataset.") for idx, prompt_batch in enumerate(train_dataset): if num_steps is not None and idx >= num_steps: break diff --git a/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py b/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py index f255f5e93..c5450075d 100644 --- a/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py +++ b/tunix/experimental/rollout/legacy_vllm_sampler_adapter.py @@ -15,6 +15,7 @@ """Legacy vLLM Sampler adapter integrating with Tunix VllmSampler.""" import abc +import numbers from typing import Any, List, Sequence import numpy as np @@ -78,6 +79,28 @@ def initialize(self) -> None: " instance or tokenizer + config." ) + def _unpadded_prompt_tokens(self, padded_tokens: Any) -> np.ndarray: + """Returns sampler-tokenized prompt ids without backend left padding.""" + arr = np.asarray(padded_tokens, dtype=np.int32).reshape(-1) + pad_id = getattr(self.tokenizer, "pad_token_id", None) + if pad_id is None: + pad_id = getattr(self.tokenizer, "eos_token_id", None) + if not isinstance(pad_id, numbers.Integral): + return arr + non_pad = np.flatnonzero(arr != pad_id) + if non_pad.size == 0: + return np.zeros(0, dtype=np.int32) + return arr[non_pad[0] :] + + def _prompt_tokens_from_request( + self, req: Any, fallback_padded_tokens: Any + ) -> np.ndarray: + """Returns request token ids directly when available, else sampler output.""" + prompt = req.prompt if hasattr(req, "prompt") else req + if not isinstance(prompt, str): + return np.asarray(prompt, dtype=np.int32).reshape(-1) + return self._unpadded_prompt_tokens(fallback_padded_tokens) + # --- Lifecycle & Topology --- async def start(self, **kwargs) -> str | None | Any: """Starts the sampling engine or local loop.""" @@ -210,12 +233,16 @@ async def sample( if toks is not None else np.zeros(0, dtype=np.int32) ) + prompt_token_ids = self._prompt_tokens_from_request( + req, sampler_output.padded_prompt_tokens[i] + ) log_ps = np.array(lps, dtype=np.float32) if lps is not None else None responses.append( base_sampler_lib.SamplingResponse( request_id=req_id, text=txt, + prompt_token_ids=prompt_token_ids, token_ids=tok_ids, logprobs=log_ps, finish_reason="stop", @@ -256,6 +283,11 @@ async def get_transfer_status(self, req_id: Any, **kwargs) -> Any: del req_id, kwargs return "SUCCESS" + async def get_load_info(self, **kwargs) -> base_sampler_lib.LoadInfo: + """Returns best-effort vLLM queue/cache load information.""" + del kwargs + return base_sampler_lib.LoadInfo() + async def post_weight_sync(self, sync_request: Any = None, **kwargs) -> Any: """Finalizes and switches active policy weights after transfer completion.""" del sync_request, kwargs diff --git a/tunix/experimental/rollout/sampler.py b/tunix/experimental/rollout/sampler.py index ecce556c0..be59dc7aa 100644 --- a/tunix/experimental/rollout/sampler.py +++ b/tunix/experimental/rollout/sampler.py @@ -70,6 +70,8 @@ class SamplingResponse(datatypes.Response): Attributes: text: Detokenized completion text string generated by the model. + prompt_token_ids: Array of prompt token IDs as tokenized by the sampler + backend, with backend padding removed. token_ids: Array of generated token IDs (unpadded). logprobs: Array of per-token log-probabilities under the sampling policy (required for RL), or None if not requested. @@ -82,6 +84,9 @@ class SamplingResponse(datatypes.Response): """ text: str = "" + prompt_token_ids: np.ndarray = dataclasses.field( + default_factory=lambda: np.zeros(0, dtype=np.int32) + ) token_ids: np.ndarray = dataclasses.field( default_factory=lambda: np.zeros(0, dtype=np.int32) ) @@ -132,6 +137,10 @@ class Sampler(Protocol): """Protocol defining standard lifecycle, sampling, and weight-sync interface for worker slices.""" # --- Lifecycle & Topology --- + def initialize(self) -> None: + """Initializes backend resources before serving requests.""" + ... + async def start(self, **kwargs) -> str | None | Any: """Starts the sampling server engine or local loop.""" ... diff --git a/tunix/experimental/rollout/vanilla_sampler_adapter.py b/tunix/experimental/rollout/vanilla_sampler_adapter.py index 421433b32..d3dc91224 100644 --- a/tunix/experimental/rollout/vanilla_sampler_adapter.py +++ b/tunix/experimental/rollout/vanilla_sampler_adapter.py @@ -15,6 +15,7 @@ """Vanilla Sampler adapter using Tunix JAX Sampler.""" import abc +import numbers from typing import Any, List, Sequence import numpy as np from tunix.experimental.rollout import sampler as base_sampler_lib @@ -113,6 +114,19 @@ def initialize(self) -> None: " instance or transformer + tokenizer." ) + def _unpadded_prompt_tokens(self, padded_tokens: Any) -> np.ndarray: + """Returns sampler-tokenized prompt ids without backend left padding.""" + arr = np.asarray(padded_tokens, dtype=np.int32).reshape(-1) + pad_id = getattr(self.tokenizer, "pad_token_id", None) + if pad_id is None: + pad_id = getattr(self.tokenizer, "eos_token_id", None) + if not isinstance(pad_id, numbers.Integral): + return arr + non_pad = np.flatnonzero(arr != pad_id) + if non_pad.size == 0: + return np.zeros(0, dtype=np.int32) + return arr[non_pad[0] :] + # --- Lifecycle & Topology --- async def start(self, **kwargs) -> str | None | Any: """Starts the sampling engine or local loop.""" @@ -259,12 +273,16 @@ async def sample( if toks is not None else np.zeros(0, dtype=np.int32) ) + prompt_token_ids = self._unpadded_prompt_tokens( + sampler_output.padded_prompt_tokens[i] + ) log_ps = np.array(lps, dtype=np.float32) if lps is not None else None responses.append( base_sampler_lib.SamplingResponse( request_id=req_id, text=txt, + prompt_token_ids=prompt_token_ids, token_ids=tok_ids, logprobs=log_ps, finish_reason="stop", @@ -298,6 +316,11 @@ async def get_transfer_status(self, req_id: Any, **kwargs) -> Any: del req_id, kwargs return "SUCCESS" + async def get_load_info(self, **kwargs) -> base_sampler_lib.LoadInfo: + """Returns best-effort local sampler load information.""" + del kwargs + return base_sampler_lib.LoadInfo() + async def post_weight_sync(self, sync_request: Any = None, **kwargs) -> Any: """Finalizes and switches active policy weights after transfer completion.""" del sync_request, kwargs diff --git a/tunix/experimental/train/peft_trainer_v2.py b/tunix/experimental/train/peft_trainer_v2.py index df994c284..707fc1512 100644 --- a/tunix/experimental/train/peft_trainer_v2.py +++ b/tunix/experimental/train/peft_trainer_v2.py @@ -1112,7 +1112,11 @@ def restore_checkpoint(self, **kwargs) -> Any: @override def prepare_weight_sync(self, **kwargs) -> None: - pass + if kwargs: + raise ValueError(f"Unexpected prepare_weight_sync kwargs: {sorted(kwargs)}") + raise NotImplementedError( + "PeftTrainer V2 weight sync is not implemented yet." + ) @override def get_metrics(self) -> exp_metrics.MetricsBuffer: diff --git a/tunix/experimental/worker/inference_worker.py b/tunix/experimental/worker/inference_worker.py index 36d39be47..7603df915 100644 --- a/tunix/experimental/worker/inference_worker.py +++ b/tunix/experimental/worker/inference_worker.py @@ -21,6 +21,7 @@ there is no weight sync. """ +from collections.abc import Sequence from typing import Any, Protocol import jax @@ -33,6 +34,32 @@ WorkerState = datatypes.WorkerState +def _left_pad_tokens( + values: Any, + length: int, + *, + pad_id: int, +) -> np.ndarray: + arr = np.asarray(values, dtype=np.int32).reshape(-1)[-length:] + out = np.full(length, pad_id, dtype=np.int32) + if arr.size: + out[-arr.size:] = arr + return out + + +def _right_pad_tokens( + values: Any, + length: int, + *, + pad_id: int, +) -> np.ndarray: + arr = np.asarray(values, dtype=np.int32).reshape(-1)[:length] + out = np.full(length, pad_id, dtype=np.int32) + if arr.size: + out[:arr.size] = arr + return out + + class ReferenceScoringCore(Protocol): """Structural type of the frozen inference core an InferenceWorker wraps.""" @@ -75,6 +102,9 @@ def __init__( eos_id: int, model_version: int = 0, chunk_size: int | None = None, + max_prompt_length: int | None = None, + max_response_length: int | None = None, + temperature: float = 1.0, ): """Initializes the worker. @@ -86,6 +116,11 @@ def __init__( eos_id: End-of-sequence token id. model_version: Version tag for the hosted weights; constant while frozen. chunk_size: Optional maximum batch size for scoring to reduce peak memory. + max_prompt_length: Optional fixed prompt length for role-oriented + per-token log-prob requests. + max_response_length: Optional fixed completion length for role-oriented + per-token log-prob requests. + temperature: Default sampling temperature used for reference log-probs. """ self._worker_id = worker_id self._core = core @@ -93,6 +128,9 @@ def __init__( self._eos_id = eos_id self._model_version = model_version self._chunk_size = chunk_size + self._max_prompt_length = max_prompt_length + self._max_response_length = max_response_length + self._temperature = temperature def info(self) -> datatypes.WorkerInfo: return datatypes.WorkerInfo( @@ -166,6 +204,82 @@ def _score(prompt_chunk: np.ndarray, completion_chunk: np.ndarray) -> Any: ), ) + def per_token_logps( + self, + items: Sequence[datatypes.TrajectoryItem], + **kwargs: Any, + ) -> np.ndarray: + """Scores reference log-probs for trajectory items. + + TrajectoryItem token arrays are produced by rollout and can be ragged. This + method pads them into the fixed-shape LogprobsRequest expected by + compute_logps(). + """ + item_list = list(items) + if not item_list: + response_len = kwargs.get( + "max_response_length", self._max_response_length or 0 + ) + return np.zeros((0, response_len), dtype=np.float32) + + max_prompt_length = kwargs.get( + "max_prompt_length", self._max_prompt_length + ) + if max_prompt_length is None: + max_prompt_length = max( + 1, + max( + len(item.prompt_tokens) + if item.prompt_tokens is not None + else 0 + for item in item_list + ), + ) + + max_response_length = kwargs.get( + "max_response_length", self._max_response_length + ) + if max_response_length is None: + max_response_length = max( + 1, + max( + len(item.completion_tokens) + if item.completion_tokens is not None + else 0 + for item in item_list + ), + ) + + req = datatypes.LogprobsRequest( + request_id="reference_logps", + prompt_tokens=np.stack([ + _left_pad_tokens( + item.prompt_tokens if item.prompt_tokens is not None else [], + int(max_prompt_length), + pad_id=self._pad_id, + ) + for item in item_list + ]), + completion_tokens=np.stack([ + _right_pad_tokens( + ( + item.completion_tokens + if item.completion_tokens is not None + else [] + ), + int(max_response_length), + pad_id=self._pad_id, + ) + for item in item_list + ]), + temperature=float(kwargs.get("temperature", self._temperature)), + model_role="reference", + ) + resp = self.compute_logps(req) + if resp.error is not None: + raise RuntimeError(resp.error.message) + return np.asarray(resp.per_token_logps, dtype=np.float32) + def score(self, req: datatypes.ScoreRequest) -> datatypes.ScoreResponse: """Scores one scalar per row under a hosted (frozen) reward model.""" try: diff --git a/tunix/experimental/worker/remote_execution.py b/tunix/experimental/worker/remote_execution.py index 60b0c4207..8460473c2 100644 --- a/tunix/experimental/worker/remote_execution.py +++ b/tunix/experimental/worker/remote_execution.py @@ -471,10 +471,17 @@ class ActorHandle(abc.ABC): """Stateful 1-to-1 routing handle targeting a specific remote worker instance.""" @classmethod - def from_address(cls, target_address: str) -> "ActorHandle": + def from_address( + cls, + target_address: str, + *, + rpc_timeout_s: Optional[float] = RPC_TIMEOUT_S, + ) -> "ActorHandle": """Instantiates a remote actor handle targeting the specified string URI.""" if target_address.startswith("grpc://") and _GRPC_AVAILABLE: - return GrpcRemoteActorHandle(target_address=target_address) + return GrpcRemoteActorHandle( + target_address=target_address, rpc_timeout_s=rpc_timeout_s + ) return RemoteActorHandle(target_address=target_address) @abc.abstractmethod diff --git a/tunix/experimental/worker/rollout_worker.py b/tunix/experimental/worker/rollout_worker.py index 4386e4961..036a15a99 100644 --- a/tunix/experimental/worker/rollout_worker.py +++ b/tunix/experimental/worker/rollout_worker.py @@ -51,8 +51,6 @@ class RolloutConfig(base_rollout.RolloutConfig): WorkerState = datatypes.WorkerState -WorkerState = datatypes.WorkerState - class RolloutWorker(abstract_worker.Worker): """Worker wrapper for rollout collection. @@ -75,6 +73,7 @@ def __init__( super().__init__() self.worker_id = worker_id self.config = config + self._policy_version = 0 if tokenizer is None or chat_parser is None: raise ValueError( "RolloutWorker requires valid tokenizer and chat_parser arguments" @@ -91,7 +90,7 @@ def __init__( ) @property - def sampler(self) -> Any: + def sampler(self) -> sampler_lib.Sampler: return self.manager.sampler def get_worker_id(self) -> str: @@ -100,17 +99,31 @@ def get_worker_id(self) -> str: def info(self) -> datatypes.WorkerInfo: return datatypes.WorkerInfo( - worker_id=self.worker_id, roles=frozenset({"rollout"}) + worker_id=self.worker_id, + roles=frozenset({"rollout"}), + resources={ + "sampler": type(self.sampler).__name__, + "policy_version": self._policy_version, + }, ) def initialize(self) -> datatypes.Response: self.state = WorkerState.INITIALIZING + self.sampler.initialize() try: - return datatypes.Response() + return datatypes.Response( + metadata={ + "worker_id": self.worker_id, + "state": self.state.value, + "policy_version": self._policy_version, + } + ) finally: self.state = WorkerState.READY def compile(self, dummy_data: Any) -> datatypes.Response: + if self.state == WorkerState.PENDING: + self.initialize() self.state = WorkerState.COMPILING try: return datatypes.Response() @@ -118,7 +131,15 @@ def compile(self, dummy_data: Any) -> datatypes.Response: self.state = WorkerState.READY def start(self) -> datatypes.Response: - return datatypes.Response() + if self.state == WorkerState.PENDING: + self.initialize() + return datatypes.Response( + metadata={ + "worker_id": self.worker_id, + "state": self.state.value, + "policy_version": self._policy_version, + } + ) def stop(self) -> datatypes.Response: self.state = WorkerState.STOPPED @@ -140,7 +161,110 @@ def _compile_with_shapes(self, abstract_state: Any) -> None: pass def heartbeat(self) -> datatypes.HealthReport: - return datatypes.HealthReport(state=self.state) + return datatypes.HealthReport( + state=self.state, + policy_version=self._policy_version, + inflight=len(self.manager._active_tasks), # pylint: disable=protected-access + queue_depth=self.manager._completed_queue.qsize(), # pylint: disable=protected-access + ) + + def _left_pad_prompt_token_ids( + self, prompt_token_ids: Sequence[np.ndarray] + ) -> np.ndarray: + pad_id = getattr(self.manager.tokenizer, "pad_token_id", None) + if pad_id is None: + pad_id = getattr(self.manager.tokenizer, "eos_token_id", 0) or 0 + configured_len = ( + getattr(self.config, "max_prompt_length", 0) if self.config else 0 + ) + max_len = max([1, configured_len] + [len(ids) for ids in prompt_token_ids]) + padded = np.full((len(prompt_token_ids), max_len), pad_id, dtype=np.int32) + for i, ids in enumerate(prompt_token_ids): + if ids.size: + padded[i, -min(ids.size, max_len) :] = ids[-max_len:] + return padded + + def _as_sampling_response_list( + self, responses: Any + ) -> list[sampler_lib.SamplingResponse]: + if isinstance(responses, (list, tuple)): + return list(responses) + return [responses] + + async def sample_prompts( + self, + prompts: str | Sequence[str], + *, + max_generation_steps: int | None = None, + temperature: float | None = None, + top_p: float | None = None, + top_k: int | None = None, + seed: int | None = None, + return_logprobs: bool = True, + ) -> base_rollout.RolloutOutput: + """Direct single-turn prompt sampling path using the worker's Sampler.""" + if self.state == WorkerState.PENDING: + self.initialize() + prompt_list = [prompts] if isinstance(prompts, str) else list(prompts) + if not prompt_list: + return base_rollout.RolloutOutput( + text=[], + logits=None, + tokens=[], + left_padded_prompt_tokens=np.zeros((0, 1), dtype=np.int32), + logprobs=[] if return_logprobs else None, + ) + + config = self.config or base_rollout.RolloutConfig() + sampling_params = sampler_lib.SamplingParams( + max_tokens=( + max_generation_steps + if max_generation_steps is not None + else config.max_tokens_to_generate + ), + temperature=temperature if temperature is not None else config.temperature, + top_p=top_p if top_p is not None else config.top_p, + top_k=top_k if top_k is not None else config.top_k, + seed=seed if seed is not None else config.seed, # pyrefly: ignore[bad-argument-type] + return_logprobs=return_logprobs, + ) + requests = [ + sampler_lib.SamplingRequest( + request_id=f"{self.worker_id}_sample_{i}", + prompt=prompt, + sampling_params=sampling_params, + ) + for i, prompt in enumerate(prompt_list) + ] + responses = self._as_sampling_response_list( + await self.sampler.sample(requests) + ) + if len(responses) != len(prompt_list): + raise RuntimeError( + f"Sampler returned {len(responses)} responses for" + f" {len(prompt_list)} prompts." + ) + prompt_token_ids = [ + np.asarray(response.prompt_token_ids, dtype=np.int32).reshape(-1) + for response in responses + ] + + logprobs: list[np.ndarray] | None = None + if return_logprobs: + logprobs = [] + for response in responses: + assert response.logprobs is not None + logprobs.append(response.logprobs) + + return base_rollout.RolloutOutput( + text=[response.text for response in responses], + logits=None, + tokens=[response.token_ids for response in responses], + left_padded_prompt_tokens=self._left_pad_prompt_token_ids( + prompt_token_ids + ), + logprobs=logprobs, + ) def _to_rollout_response( self, @@ -180,14 +304,131 @@ def _to_rollout_response( ) return item + def _sampling_to_rollout_response( + self, + request: datatypes.RolloutRequest, + text: str, + prompt_tokens: Any, + token_ids: Any, + logprobs: Any | None, + ) -> datatypes.RolloutResponse: + """Builds the v2 rollout DTO for the direct single-turn sampler path.""" + completion_tokens = np.asarray(token_ids, dtype=np.int32).reshape(-1) + completion_logps = ( + np.asarray(logprobs, dtype=np.float32).reshape(-1) + if logprobs is not None + else None + ) + if ( + completion_logps is not None + and completion_logps.shape != completion_tokens.shape + ): + completion_logps = None + prompt_token_arr = np.asarray(prompt_tokens, dtype=np.int32).reshape(-1) + if prompt_token_arr.size == 0: + raise RuntimeError( + "Sampler response is missing prompt_token_ids for " + f"{request.request_id or request.traj_id}." + ) + metadata = dict(request.metadata or {}) + metadata.setdefault("text", text) + return datatypes.RolloutResponse( + request_id=request.request_id or request.traj_id, + prompt_id=request.prompt_id, + status="COMPLETED", + prompt_tokens=prompt_token_arr, + segments=[ + datatypes.TokenSegment( + source="assistant", + tokens=completion_tokens, + loss_mask=np.ones(completion_tokens.shape, dtype=np.float32), + logps=completion_logps, + ) + ], + env_reward=0.0, + policy_version=self._policy_version, + metadata=metadata, + ) + + async def _generate_rollout_requests_direct( + self, + requests: Sequence[datatypes.RolloutRequest], + **generation_kwargs, + ) -> list[datatypes.RolloutResponse]: + """Runs RolloutRequest batches through the direct string sampler.""" + config = self.config or base_rollout.RolloutConfig() + sampling_requests = [] + for req in requests: + sample_kwargs = dict(req.generation_kwargs) + sample_kwargs.update(generation_kwargs) + sampling_requests.append( + sampler_lib.SamplingRequest( + request_id=req.request_id or req.traj_id, + prompt=req.prompt, + metadata=sample_kwargs, + sampling_params=sampler_lib.SamplingParams( + max_tokens=sample_kwargs.get( + "max_generation_steps", config.max_tokens_to_generate + ), + temperature=sample_kwargs.get("temperature", config.temperature), + top_p=sample_kwargs.get("top_p", config.top_p), + top_k=sample_kwargs.get("top_k", config.top_k), + seed=sample_kwargs.get("seed", config.seed), + return_logprobs=sample_kwargs.get("return_logprobs", True), + ), + ) + ) + responses = self._as_sampling_response_list( + await self.sampler.sample(sampling_requests) + ) + if len(responses) != len(requests): + raise RuntimeError( + f"Sampler returned {len(responses)} responses for" + f" {len(requests)} rollout requests." + ) + return [ + self._sampling_to_rollout_response( + request=req, + text=responses[i].text, + prompt_tokens=responses[i].prompt_token_ids, + token_ids=responses[i].token_ids, + logprobs=responses[i].logprobs, + ) + for i, req in enumerate(requests) + ] + async def generate( self, requests: ( datatypes.RolloutRequest | Sequence[datatypes.RolloutRequest] | Any - ), + ) = None, on_complete: Optional[Callable[[datatypes.RolloutResponse], None]] = None, + prompts: Any = None, + **generation_kwargs, ) -> datatypes.RolloutResponse | List[datatypes.RolloutResponse] | Any: """Coroutine method for single or batched generate requests.""" + if requests is None: + requests = prompts + if requests is None: + raise ValueError("generate requires `requests` or v2 `prompts`.") + if isinstance(requests, str) or ( + isinstance(requests, (list, tuple)) + and all(isinstance(req, str) for req in requests) + ): + return await self.sample_prompts(requests, **generation_kwargs) # pyrefly: ignore[bad-argument-type] + # if isinstance(requests, datatypes.RolloutRequest): + # return ( + # await self._generate_rollout_requests_direct( + # [requests], **generation_kwargs + # ) + # )[0] + # if isinstance(requests, (list, tuple)) and all( + # isinstance(req, datatypes.RolloutRequest) for req in requests + # ): + # return await self._generate_rollout_requests_direct( + # list(requests), **generation_kwargs + # ) + cb = None if on_complete is not None: cb = lambda item: on_complete(self._to_rollout_response(item)) @@ -210,6 +451,8 @@ async def as_completed_stream( async def pre_weight_sync(self, sync_request: Any = None, **kwargs) -> Any: """Prepares the worker for an upcoming weight synchronization step.""" + if self.state == WorkerState.PENDING: + self.initialize() self.state = WorkerState.SYNCING try: return await self.manager.pre_weight_sync(sync_request, **kwargs) @@ -218,9 +461,15 @@ async def pre_weight_sync(self, sync_request: Any = None, **kwargs) -> Any: async def weight_sync(self, sync_request: Any = None, **kwargs) -> Any: """Synchronizes the worker's internal model weights.""" + if self.state == WorkerState.PENDING: + self.initialize() self.state = WorkerState.SYNCING try: - return await self.manager.weight_sync(sync_request, **kwargs) + metadata = kwargs.pop("metadata", None) + request = sync_request if sync_request is not None else metadata + result = await self.manager.weight_sync(request, **kwargs) + self._policy_version += 1 + return result finally: self.state = WorkerState.READY diff --git a/tunix/experimental/worker/trainer_worker.py b/tunix/experimental/worker/trainer_worker.py index bd9519df9..c53f97513 100644 --- a/tunix/experimental/worker/trainer_worker.py +++ b/tunix/experimental/worker/trainer_worker.py @@ -14,7 +14,8 @@ """TrainerWorker implementation for role-based isolation.""" -from typing import Any, Callable +import contextlib +from typing import Any, Callable, ContextManager, cast from tunix.experimental.common import datatypes from tunix.experimental.train import abstract_trainer @@ -46,93 +47,209 @@ def __init__( self._trainer = trainer_factory() self._is_running = False self._worker_id = worker_id + self._state = WorkerState.PENDING + self._last_error: str | None = None + + def _policy_version(self) -> int: + return int(getattr(self._trainer, "policy_version", 0)) + + def _response(self, **metadata: Any) -> datatypes.Response: + return datatypes.Response( + metadata={ + "worker_id": self._worker_id, + "state": self.state.value, + "policy_version": self._policy_version(), + **metadata, + } + ) + + def _ensure_ready(self) -> None: + if self.state == WorkerState.PENDING: + self.initialize() + if self.state != WorkerState.READY: + raise RuntimeError(f"TrainerWorker is not ready: {self.state.value}.") def initialize(self) -> datatypes.Response: """Initializes the worker and the underlying trainer.""" + if self.state == WorkerState.READY: + return self._response(initialized=True, already_ready=True) self.state = WorkerState.INITIALIZING try: - return datatypes.Response() + return self._response(initialized=True) finally: self.state = WorkerState.READY - def compile(self, dummy_data: Any) -> datatypes.Response: + def compile(self, dummy_data: Any = None) -> datatypes.Response: """Triggers JIT compilation using the provided dummy_data.""" + if self.state == WorkerState.PENDING: + self.initialize() self.state = WorkerState.COMPILING try: self._trainer.compile(dummy_data) - return datatypes.Response() + return self._response(compiled=True) + except Exception as exc: + self._last_error = str(exc) + self.state = WorkerState.ERROR + raise finally: - self.state = WorkerState.READY + if self.state == WorkerState.COMPILING: + self.state = WorkerState.READY def start(self) -> datatypes.Response: """Starts the worker's main loop.""" + if self.state == WorkerState.PENDING: + self.initialize() + if self.state != WorkerState.READY: + raise RuntimeError(f"Cannot start TrainerWorker from {self.state.value}.") self._is_running = True - return datatypes.Response() + return self._response(started=True) def stop(self) -> datatypes.Response: """Gracefully stops the worker.""" + if self.state == WorkerState.STOPPED: + return self._response(stopped=True, already_stopped=True) self._is_running = False - self.state = WorkerState.STOPPED + if self.state == WorkerState.READY: + self.state = WorkerState.DRAINING self._trainer.close() - return datatypes.Response() + self.state = WorkerState.STOPPED + return self._response(stopped=True) def info(self) -> datatypes.WorkerInfo: return datatypes.WorkerInfo( - worker_id=self._worker_id, roles=frozenset({"trainer"}) + worker_id=self._worker_id, + roles=frozenset({"trainer", "weight_sync"}), + resources={ + "trainer": type(self._trainer).__name__, + "policy_version": self._policy_version(), + }, ) def heartbeat(self) -> datatypes.HealthReport: - return datatypes.HealthReport(state=self.state) + return datatypes.HealthReport( + state=self.state, + policy_version=self._policy_version(), + last_error=self._last_error, + ) def with_loss_fn( self, loss_fn: Callable[..., Any], has_aux: bool = False - ) -> "TrainerWorker": + ) -> datatypes.Response: """Sets the loss function used by `fwd_bwd` (and evaluation).""" self._trainer.with_loss_fn(loss_fn, has_aux) - return self + return self._response(loss_fn_configured=True) def with_gen_model_input_fn( self, gen_model_input_fn: Callable[[Any], dict[str, Any]] - ) -> "TrainerWorker": + ) -> datatypes.Response: """Sets the last-mile adapter mapping a payload to the loss fn's kwargs.""" self._trainer.with_gen_model_input_fn(gen_model_input_fn) - return self + return self._response(gen_model_input_fn_configured=True) def fwd_bwd( - self, payload: datatypes.TrainerPayload, **kwargs + self, + payload: datatypes.TrainerPayload, + **kwargs: Any, ) -> datatypes.Response: - """Executes forward and backward passes.""" - self._trainer.fwd_bwd(payload, **kwargs) - return datatypes.Response() + """Executes one forward/backward pass.""" + self._ensure_ready() + kwargs.pop("skip_jit", None) + try: + self._trainer.fwd_bwd(payload, **kwargs) + self._last_error = None + return self._response(queued=True) + except Exception as exc: + self._last_error = str(exc) + self.state = WorkerState.ERROR + raise def update(self, **kwargs) -> int: """Applies the accumulated (mean) gradients as one optimizer update.""" - return self._trainer.update(**kwargs) + self._ensure_ready() + try: + train_step = self._trainer.update(**kwargs) + self._last_error = None + return train_step + except Exception as exc: + self._last_error = str(exc) + self.state = WorkerState.ERROR + raise def eval_step( self, payload: datatypes.TrainerPayload, **kwargs ) -> datatypes.Response: """Executes one evaluation step on the given payload.""" - self._trainer.eval_step(payload, **kwargs) - return datatypes.Response() + self._ensure_ready() + try: + self._trainer.eval_step(payload, **kwargs) + self._last_error = None + return self._response(evaluated=True) + except Exception as exc: + self._last_error = str(exc) + self.state = WorkerState.ERROR + raise + + def run_eval(self, eval_ds: Any, **kwargs) -> datatypes.Response: + """Runs an explicit evaluation phase over eval micro-batches.""" + self._ensure_ready() + if eval_ds is None: + return self._response(evaluated=True, eval_batches=0) + try: + run_eval = getattr(self._trainer, "run_eval", None) + if callable(run_eval): + run_eval(eval_ds, **kwargs) + self._last_error = None + return self._response(evaluated=True) + + eval_context = getattr(self._trainer, "eval_context", None) + context = ( + cast(ContextManager[Any], eval_context()) + if callable(eval_context) + else contextlib.nullcontext() + ) + eval_batches = 0 + with context: + for payload in eval_ds: + self._trainer.eval_step(payload, **kwargs) + eval_batches += 1 + self._last_error = None + return self._response(evaluated=True, eval_batches=eval_batches) + except Exception as exc: + self._last_error = str(exc) + self.state = WorkerState.ERROR + raise def save_checkpoint(self, metadata: Any, **kwargs) -> datatypes.Response: """Force the trainer to serialize its state (model + optimizer).""" - self._trainer.save_checkpoint(metadata, **kwargs) - return datatypes.Response() + self._ensure_ready() + try: + self._trainer.save_checkpoint(metadata, **kwargs) + self._last_error = None + return self._response(checkpoint_saved=True) + except Exception as exc: + self._last_error = str(exc) + self.state = WorkerState.ERROR + raise def restore_checkpoint(self, **kwargs) -> Any: """Restore state from latest checkpoint and return the metadata pytree.""" return self._trainer.restore_checkpoint(**kwargs) - def prepare_weight_sync(self, **kwargs) -> datatypes.Response: - """Stages weights for transfer and returns coordinates/metadata for Rollouts to pull.""" + def prepare_weight_sync(self, **kwargs) -> Any: + """Stages weights for transfer and returns coordinates/metadata.""" + self._ensure_ready() self.state = WorkerState.SYNCING try: - self._trainer.prepare_weight_sync(**kwargs) - return datatypes.Response() - finally: + metadata = self._trainer.prepare_weight_sync(**kwargs) self.state = WorkerState.READY + self._last_error = None + if metadata is not None: + return metadata + return self._response(weight_sync_ready=True) + except Exception as exc: + self._last_error = str(exc) + self.state = WorkerState.ERROR + raise def get_metrics(self) -> Any: """Returns and clears the recently collected step metric records.""" diff --git a/tunix/rl/rl_cluster.py b/tunix/rl/rl_cluster.py index 32c6175b1..649381332 100644 --- a/tunix/rl/rl_cluster.py +++ b/tunix/rl/rl_cluster.py @@ -880,6 +880,27 @@ def update_actor(self, train_ds, eval_ds, skip_jit=False): self.actor_trainer.train(train_ds, eval_ds, skip_jit) self._maybe_offload_model_to_cpu(self.actor_trainer.model, Role.ACTOR) + def eval_actor(self, eval_ds: Any) -> Any: + """Runs an explicit actor evaluation phase.""" + if eval_ds is None: + return None + with self._get_mesh_and_logical_axis_rules_cm(Role.ACTOR): + self._maybe_load_model_from_cpu(self.actor_trainer.model, Role.ACTOR) + run_eval = getattr(self.actor_trainer, "_run_eval", None) + if callable(run_eval): + res = run_eval(eval_ds) + else: + eval_step = getattr(self.actor_trainer, "eval_step", None) + if not callable(eval_step): + raise TypeError( + "actor_trainer must expose _run_eval(...) or eval_step(...)." + ) + res = None + for chunk in eval_ds: + eval_step(chunk) + self._maybe_offload_model_to_cpu(self.actor_trainer.model, Role.ACTOR) + return res + def update_critic(self, train_ds, eval_ds, skip_jit=False): with self._get_mesh_and_logical_axis_rules_cm(Role.CRITIC): self._maybe_load_model_from_cpu(self.critic_trainer.model, Role.CRITIC) diff --git a/tunix/tests/test_common.py b/tunix/tests/test_common.py index bef6b27c2..767d5e054 100644 --- a/tunix/tests/test_common.py +++ b/tunix/tests/test_common.py @@ -450,3 +450,19 @@ def is_running_in_colab() -> bool: return hasattr(sys.modules['IPython'].get_ipython(), 'kernel') except (NameError, KeyError, AttributeError): return False + + +def safe_set_n_cpu_devices(n: int) -> None: + """Safely set the number of CPU devices for JAX.""" + import chex # pylint: disable=g-import-not-at-top + try: + chex.set_n_cpu_devices(n) + except RuntimeError as e: + # If JAX is already initialized, check if we have enough devices. + devices = jax.local_devices() + if len(devices) < n: + raise RuntimeError( + f"JAX already initialized with {len(devices)} CPU devices, " + f"but {n} are required." + ) from e +