Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
149 changes: 16 additions & 133 deletions tests/experimental/orchestrator/batch_assembly_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,6 @@
import numpy as np
from tunix.experimental.common import datatypes
from tunix.experimental.orchestrator import batch_assembly
from tunix.rl import common as rl_common


class HelperFunctionsTest(absltest.TestCase):
Expand Down Expand Up @@ -104,8 +103,8 @@ def test_completion_aligned_slices_full_sequence(self):

class WithRefPerTokenLogpsTest(absltest.TestCase):

def _make_train_example(self, b=2, p=3, c=4):
return rl_common.TrainExample(
def _make_payload(self, b=2, p=3, c=4):
return datatypes.RLTrainerPayload(
prompt_ids=np.ones((b, p), dtype=np.int32),
prompt_mask=np.ones((b, p), dtype=np.float32),
completion_ids=np.ones((b, c), dtype=np.int32),
Expand All @@ -116,31 +115,33 @@ def _make_train_example(self, b=2, p=3, c=4):
)

def test_success_with_ndarray(self):
batch = self._make_train_example(b=2, p=3, c=4)
batch = self._make_payload(b=2, p=3, c=4)
ref_logps = np.full((2, 4), -0.5, dtype=np.float32)
updated = batch_assembly.with_ref_per_token_logps(batch, ref_logps)

self.assertIsInstance(updated, datatypes.RLTrainerPayload)
self.assertIsNotNone(updated.ref_per_token_logps)
self.assertEqual(updated.ref_per_token_logps.shape, (2, 4))
np.testing.assert_allclose(updated.ref_per_token_logps, ref_logps)
np.testing.assert_array_equal(updated.prompt_ids, batch.prompt_ids)
np.testing.assert_array_equal(updated.completion_ids, batch.completion_ids)

def test_success_with_logprobs_response(self):
batch = self._make_train_example(b=2, p=3, c=4)
batch = self._make_payload(b=2, p=3, c=4)
resp = datatypes.LogprobsResponse(
per_token_logps=np.full((2, 4), -0.8, dtype=np.float32)
)
updated = batch_assembly.with_ref_per_token_logps(batch, resp)

self.assertIsInstance(updated, datatypes.RLTrainerPayload)
self.assertIsNotNone(updated.ref_per_token_logps)
self.assertEqual(updated.ref_per_token_logps.shape, (2, 4))
np.testing.assert_allclose(
updated.ref_per_token_logps, resp.per_token_logps
)

def test_error_in_logprobs_response_raises_runtime_error(self):
batch = self._make_train_example(b=2, p=3, c=4)
batch = self._make_payload(b=2, p=3, c=4)
resp = datatypes.LogprobsResponse(
per_token_logps=None,
error=datatypes.ErrorInfo(
Expand All @@ -150,18 +151,14 @@ def test_error_in_logprobs_response_raises_runtime_error(self):
with self.assertRaisesRegex(RuntimeError, "inference worker failed"):
batch_assembly.with_ref_per_token_logps(batch, resp)

def test_rejects_non_train_example(self):
payload = datatypes.RLTrainerPayload(
token_ids=np.array([1, 2], dtype=np.int32),
token_mask=np.array([1, 1], dtype=np.float32),
loss_mask=np.array([0, 1], dtype=np.float32),
advantages=np.array([1.0, 1.0], dtype=np.float32),
)
with self.assertRaisesRegex(TypeError, "expects a padded TrainExample"):
batch_assembly.with_ref_per_token_logps(payload, np.zeros((2, 2)))
def test_rejects_unsupported_type(self):
with self.assertRaisesRegex(TypeError, "expects a padded RLTrainerPayload"):
batch_assembly.with_ref_per_token_logps(
{"raw": "batch"}, np.zeros((2, 2))
)

def test_mismatched_shape_raises_value_error(self):
batch = self._make_train_example(b=2, p=3, c=4)
batch = self._make_payload(b=2, p=3, c=4)
bad_shape_logps = np.zeros((2, 3), dtype=np.float32)
with self.assertRaisesRegex(
ValueError,
Expand Down Expand Up @@ -277,121 +274,6 @@ def test_sequence_packed_assembler_multiple_bins(self):
self.assertEqual(payloads[1].token_ids.shape, (1, 12))


class GRPOTrainExampleAssemblerTest(absltest.TestCase):

def test_rejects_non_positive_batch_size(self):
with self.assertRaisesRegex(ValueError, "batch size must be positive"):
batch_assembly.GRPOTrainExampleAssembler(
batch_size=0,
max_prompt_length=4,
max_response_length=5,
pad_id=0,
)

def test_empty_input_returns_empty_list(self):
assembler = batch_assembly.GRPOTrainExampleAssembler(
batch_size=2,
max_prompt_length=4,
max_response_length=5,
pad_id=0,
)
self.assertEmpty(assembler.pack([]))

def test_grpo_train_example_assembler_basic(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])
)

def test_grpo_assembler_optional_fields_propagation(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, 1], dtype=np.float32),
action_mask=np.array([0, 0, 1, 1, 1], dtype=np.float32),
advantages=np.array([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.ones(3, dtype=np.float32),
ref_per_token_logps=np.array([-0.3, -0.4, -0.5], dtype=np.float32),
old_per_token_logps=np.array([-0.1, -0.2, -0.3], 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.assertIsNotNone(train_example.ref_per_token_logps)
self.assertIsNotNone(train_example.old_per_token_logps)

self.assertEqual(train_example.ref_per_token_logps.shape, (2, 5))
self.assertEqual(train_example.old_per_token_logps.shape, (2, 5))

np.testing.assert_allclose(
train_example.ref_per_token_logps[0], [-0.3, -0.4, -0.5, 0.0, 0.0]
)
np.testing.assert_allclose(
train_example.old_per_token_logps[0], [-0.1, -0.2, -0.3, 0.0, 0.0]
)

def test_grpo_assembler_chunks_multiple_microbatches(self):
payload = datatypes.RLTrainerPayload(
token_ids=np.array([1, 2, 3], dtype=np.int32),
token_mask=np.ones(3, dtype=np.float32),
loss_mask=np.array([0, 1, 1], dtype=np.float32),
advantages=np.array([1.0, 1.0], dtype=np.float32),
prompt_ids=np.array([1], dtype=np.int32),
completion_ids=np.array([2, 3], dtype=np.int32),
)

assembler = batch_assembly.GRPOTrainExampleAssembler(
batch_size=2,
max_prompt_length=3,
max_response_length=4,
pad_id=0,
)
train_examples = assembler.pack([payload, payload, payload])

self.assertLen(train_examples, 2)
self.assertEqual(train_examples[0].prompt_ids.shape, (2, 3))
self.assertEqual(train_examples[1].prompt_ids.shape, (2, 3))


def _make_payload(
prompt_len: int,
completion_len: int,
Expand Down Expand Up @@ -585,8 +467,9 @@ def test_scalar_advantage_broadcasts_over_completion(self):
np.testing.assert_allclose(payload.advantages[0], [2.5, 2.5, 2.5, 0, 0])

def test_sequence_aligned_advantage_is_sliced_to_completion(self):
item = _make_payload(2, 3)
item.advantages = np.array([0, 0, 2, 2, 2], dtype=np.float32)
item = _make_payload(
2, 3, advantage=np.array([0, 0, 2, 2, 2], dtype=np.float32)
)
payload = self._assembler().pack([item])[0]

np.testing.assert_allclose(payload.advantages[0], [2, 2, 2, 0, 0])
Expand Down
9 changes: 4 additions & 5 deletions tests/experimental/orchestrator/rl_program_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,6 @@
from tunix.experimental.orchestrator import distributed_rl_engine
from tunix.experimental.orchestrator import rl_program
from tunix.experimental.worker import remote_execution
from tunix.rl import common as rl_common
from tunix.sft import metrics_logger as metrics_logger_lib
from tunix.sft import utils as sft_utils

Expand Down Expand Up @@ -726,7 +725,7 @@ async def _run():
def test_reference_kl_logprobs_scoring_in_train_stage(self):
async def _run():
self.mock_algo.requires_reference_kl = True
mock_train_example = rl_common.TrainExample(
mock_payload = datatypes.RLTrainerPayload(
prompt_ids=np.array([[1, 2]], dtype=np.int32),
prompt_mask=np.ones((1, 2), dtype=np.float32),
completion_ids=np.array([[3, 4]], dtype=np.int32),
Expand All @@ -735,7 +734,7 @@ async def _run():
ref_per_token_logps=None,
old_per_token_logps=None,
)
self.assembler.pack = mock.MagicMock(return_value=[mock_train_example])
self.assembler.pack = mock.MagicMock(return_value=[mock_payload])
self.mock_engine.per_token_logps = mock.AsyncMock(
return_value=np.array([[-0.1, -0.2]], dtype=np.float32)
)
Expand All @@ -746,7 +745,7 @@ async def _run():
await program.run_async(self.mock_engine)

self.mock_engine.per_token_logps.assert_called_once_with(
datatypes.Role.REFERENCE, items=mock_train_example
datatypes.Role.REFERENCE, items=mock_payload
)
self.assertEqual(program.step, 1)

Expand All @@ -755,7 +754,7 @@ async def _run():
def test_reference_kl_raises_type_error_for_invalid_microbatch(self):
async def _run():
self.mock_algo.requires_reference_kl = True
# Returning a raw dict instead of TrainExample
# Returning a raw dict instead of RLTrainerPayload
self.assembler.pack = mock.MagicMock(return_value=[{"raw": "batch"}])
_set_mock_poll_batches(self.mock_engine, _make_trajectory_group())
program = self._create_program(dataset=["prompt_0"])
Expand Down
3 changes: 1 addition & 2 deletions tests/experimental/worker/inference_worker_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,6 @@
from tunix.experimental.common import datatypes
from tunix.experimental.common import rpc_utils
from tunix.experimental.worker import inference_worker as inference_lib
from tunix.rl import common as rl_common

WorkerState = datatypes.WorkerState

Expand Down Expand Up @@ -107,7 +106,7 @@ def test_chunking_matches_single_pass(self):

def test_per_token_logps_uses_padded_batch_without_repadding(self):
core = _StubCore()
batch = rl_common.TrainExample(
batch = datatypes.RLTrainerPayload(
prompt_ids=np.array([[0, 0, 5, 6], [0, 7, 8, 9]], dtype=np.int32),
prompt_mask=np.array([[0, 0, 1, 1], [0, 1, 1, 1]], dtype=np.float32),
completion_ids=np.array([[10, 11, 0], [12, 0, 0]], dtype=np.int32),
Expand Down
15 changes: 9 additions & 6 deletions tunix/experimental/common/datatypes.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
import enum
import time
from typing import Any, Dict
import flax
import uuid
from jax.typing import ArrayLike # pylint: disable=g-importing-member
import numpy as np
Expand Down Expand Up @@ -519,7 +520,7 @@ def __post_init__(self):
##### Training DTOs #####


@dataclasses.dataclass(kw_only=True)
@flax.struct.dataclass(frozen=True, kw_only=True)
class TrainerPayload:
"""Base class for generic trainer payloads.

Expand All @@ -540,7 +541,7 @@ class TrainerPayload:
segment_positions: ArrayLike | None = None


@dataclasses.dataclass(kw_only=True)
@flax.struct.dataclass(frozen=True, kw_only=True)
class SFTTrainerPayload(TrainerPayload):
"""Supervised Fine-Tuning (SFT) trainer payload.

Expand All @@ -560,7 +561,7 @@ class SFTTrainerPayload(TrainerPayload):

# TODO(tunix-dev): Introduce PPOTrainerPayload to replace generic
# RLTrainerPayload when PPO specific fields are needed.
@dataclasses.dataclass(kw_only=True)
@flax.struct.dataclass(frozen=True, kw_only=True)
class RLTrainerPayload(TrainerPayload):
"""RL training payload.

Expand All @@ -582,8 +583,8 @@ class RLTrainerPayload(TrainerPayload):
metadata: Extra payload metadata dictionary.
"""

advantages: ArrayLike
loss_mask: ArrayLike
advantages: ArrayLike | None = None
loss_mask: ArrayLike | None = None
action_mask: ArrayLike | None = None
# TODO(tunix-dev): make prompt_ids/mask and completion_ids/mask required after
# SequencePackedBatchAssembler refactor is done.
Expand All @@ -596,7 +597,9 @@ class RLTrainerPayload(TrainerPayload):
sampler_is_weights: ArrayLike | None = None
returns: ArrayLike | None = None
old_values: ArrayLike | None = None
metadata: dict[str, Any] = dataclasses.field(default_factory=dict)
metadata: dict[str, Any] = flax.struct.field(
default_factory=dict, pytree_node=False
)
# TODO(tunix-dev): add ppo specific fields in PPORLTrainerPayload.


Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -263,7 +263,7 @@ def _grpo_model_input(
pad_id: int,
eos_id: int,
) -> dict[str, Any]:
"""Maps a TrainExample microbatch to algo_core.grpo_loss_fn kwargs."""
"""Maps an RLTrainerPayload microbatch to algo_core.grpo_loss_fn kwargs."""
return {
"train_example": train_example,
"algo_config": algo_config,
Expand Down Expand Up @@ -592,7 +592,7 @@ def accept_worker(hostname: str, _: int, metadata: bytes) -> None:
dataset=_iter_prompt_items(args),
max_steps=args.max_steps,
reward_fns=reward_fns,
assembler=batch_assembly.GRPOTrainExampleAssembler(
assembler=batch_assembly.PaddedBatchAssembler(
batch_size=args.train_micro_batch_size,
max_prompt_length=args.max_prompt_length,
max_response_length=args.max_response_length,
Expand Down
Loading
Loading