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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 5 additions & 3 deletions tests/rl/conftest.py
Original file line number Diff line number Diff line change
Expand Up @@ -140,6 +140,8 @@ def _install_rl_trainer_worker_stubs() -> None:
class TrainingController:
pass

TrainingLogInfo = dict

class WorkerConfig(BaseModel):
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)

Expand All @@ -148,8 +150,8 @@ class TrainingWorker:

controller_mod = _new_module("xtuner.v1.rl.trainer.controller")
controller_mod.TrainingController = TrainingController
controller_mod.ColateItem = object
controller_mod.__all__ = ["TrainingController", "ColateItem"]
controller_mod.TrainingLogInfo = TrainingLogInfo
controller_mod.__all__ = ["TrainingController", "TrainingLogInfo"]

worker_mod = _new_module("xtuner.v1.rl.trainer.worker")
worker_mod.TrainingWorker = TrainingWorker
Expand All @@ -166,12 +168,12 @@ class TrainingWorker:
]

trainer_pkg.TrainingController = TrainingController
trainer_pkg.TrainingLogInfo = TrainingLogInfo
trainer_pkg.WorkerConfig = WorkerConfig
trainer_pkg.TrainingWorker = TrainingWorker
trainer_pkg.WorkerInputItem = dict
trainer_pkg.WorkerLogItem = dict
trainer_pkg.WorkerTrainLogItem = dict
trainer_pkg.ColateItem = object

sys.modules.setdefault("xtuner.v1.rl.trainer", trainer_pkg)
sys.modules.setdefault("xtuner.v1.rl.trainer.controller", controller_mod)
Expand Down
115 changes: 115 additions & 0 deletions tests/rl/test_pack.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,115 @@
import random
import unittest
from types import SimpleNamespace

import torch

from xtuner.v1.data_proto.sequence_context import SequenceContext
from xtuner.v1.rl.trainer.pack import RLDataPacker
from xtuner.v1.rl.trainer.worker import TrainingWorker


class TestDataBatchPacker(unittest.TestCase):
def setUp(self):
self.pack_max_length = 3072

def _run_strategy_test(
self,
strategy,
world_size,
optimizer_steps,
lengths,
pack_max_length,
expected_padding=None,
):
packer = RLDataPacker(
pack_max_length=pack_max_length,
world_size=world_size,
data_replicate_size=1,
optimizer_steps=optimizer_steps,
pack_strategy=strategy,
)

packed_indices, padding_tokens = packer.pack(lengths)

all_packs = [
pack_indices
for rank_indices in packed_indices
for step_indices in rank_indices
for pack_indices in step_indices
]
seen_indices = [data_index for pack_indices in all_packs for data_index in pack_indices]
self.assertEqual(sorted(seen_indices), list(range(len(lengths))))
pack_token_counts = [sum(lengths[data_index] for data_index in pack_indices) for pack_indices in all_packs]
self.assertTrue(all(pack_tokens <= pack_max_length for pack_tokens in pack_token_counts))
self.assertEqual(len(all_packs) * pack_max_length, sum(lengths) + padding_tokens)

if strategy == "balance":
rank_token_counts = [
sum(
lengths[data_index]
for step_indices in rank_indices
for pack_indices in step_indices
for data_index in pack_indices
)
for rank_indices in packed_indices
]
self.assertLessEqual(max(rank_token_counts) - min(rank_token_counts), max(lengths))

if expected_padding is not None:
self.assertEqual(padding_tokens, expected_padding)

def test_variable_packs(self):
lengths = [1500, 1000, 2800, 3000, 1500, 2000, 2100, 1000, 800]
self._run_strategy_test("native", 2, 2, lengths, self.pack_max_length, 15020)
self._run_strategy_test("balance", 2, 2, lengths, self.pack_max_length, 8876)
self._run_strategy_test("greedy", 2, 2, lengths, self.pack_max_length, 8876)

def test_imbalance_dp_size(self):
lengths = [500]
for strategy in ["native", "balance", "greedy"]:
self._run_strategy_test(strategy, 2, 1, lengths, self.pack_max_length, 5644)

def test_imbalanced_steps(self):
lengths = [100, 200, 2500, 3000, 50, 400, 1000, 1500]
self._run_strategy_test("native", 2, 4, lengths, self.pack_max_length, 15826)
self._run_strategy_test("balance", 2, 4, lengths, self.pack_max_length, 15826)
self._run_strategy_test("greedy", 2, 4, lengths, self.pack_max_length, 3538)

def test_random_lengths(self):
lengths = [random.randint(1, 32768) for _ in range(1024)]
for strategy in ["native", "balance", "greedy"]:
self._run_strategy_test(strategy, 8, 16, lengths, 32768)

def test_native_supports_pack_length_below_split_size(self):
self._run_strategy_test("native", 2, 1, [128], 256, 384)


class TestTrainingWorkerPackMaterialization(unittest.TestCase):
@staticmethod
def _create_dummy_item(length: int, value: int):
input_ids = torch.full((1, length), value, dtype=torch.long)
seq_ctx = SequenceContext.from_input_ids((input_ids,), device="cpu")
return {
"seq_ctx": seq_ctx,
"shifted_labels": torch.full((1, length), value, dtype=torch.long),
"advantages": torch.full((1, length), float(value), dtype=torch.float32),
"rollout_logprobs": torch.full((1, length), float(value), dtype=torch.float32),
}

def test_worker_selects_indices_and_materializes_packs(self):
worker = TrainingWorker.__new__(TrainingWorker)
worker.config = SimpleNamespace(pack_max_length=8, model_cfg=None)
data_batches = [self._create_dummy_item(3, 1), self._create_dummy_item(2, 2)]

packed_data = worker._materialize_packs(data_batches, [[[0, 1]], [[]]])

self.assertEqual(len(packed_data), 2)
self.assertEqual(packed_data[0][0]["seq_ctx"].input_ids.numel(), 8)
self.assertEqual(packed_data[0][0]["seq_ctx"].num_padding, 3)
self.assertEqual(packed_data[1][0]["seq_ctx"].num_padding, 8)
self.assertEqual(packed_data[0][0]["seq_ctx"].input_ids[0, :5].tolist(), [1, 1, 1, 2, 2])


if __name__ == "__main__":
unittest.main()
11 changes: 7 additions & 4 deletions tests/rl/test_rl_colocate_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -167,16 +167,19 @@ def _make_trainer(self, agent_loop_manager, *, total_train_steps: int = 1, sync_
offload=MagicMock(return_value="train_offloaded"),
weight_update=MagicMock(return_value="weights_updated"),
fit=MagicMock(
return_value=[
{
return_value={
"worker_log_infos": [{
"rollout_is_metrics": {},
"mismatch_metrics": {},
"rollout_entropy": 0.0,
"train_entropy": 0.0,
"train_metrics": [],
"sft_train_metrics": {},
}
]
}],
"padding_tokens": 0,
"pack_time": 0.0,
"train_time": 0.0,
}
),
)
return trainer
Expand Down
9 changes: 4 additions & 5 deletions tests/rl/test_rl_colocate_trainer_integration.py
Original file line number Diff line number Diff line change
Expand Up @@ -263,14 +263,13 @@ def test_rl_train_with_sft(self):
shifted_labels_tensor = torch.tensor(shifted_labels, dtype=torch.int64).unsqueeze(0)

adv_val = advantages[i].item()
# Controller._packing expects `advantage` as a list and will flatten it.
# Keep the length consistent with shifted_labels/input_ids.
advantage_list = [adv_val] * (len(prompt_ids) - 1) + [adv_val] * len(response_ids)

data_batches.append(dict(
seq_ctx=SequenceContext.from_input_ids((input_ids_tensor,), device="cpu"),
shifted_labels=shifted_labels_tensor,
advantage=advantage_list,
advantages=torch.tensor(advantage_list, dtype=torch.float32),
rollout_logprobs=None,
))

# RLColocateTrainer initializes by offloading train workers to CPU.
Expand All @@ -286,7 +285,7 @@ def test_rl_train_with_sft(self):
train_controller.onload(target="all")
log_infos = train_controller.fit(data_batches, pack_max_length=1024, rollout_idx=1)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude: [测试] 该测试的 train_worker_cfg 用的是 pack_max_length=2048(本文件 L135),但 L280/286/316 仍传 pack_max_length=1024,会被 controller 新增的一致性校验(controller.py:81-84)直接抛 ValueError,测试必然失败。

efficient_attn_ratio_list = []
for log_info in log_infos:
for log_info in log_infos["worker_log_infos"]:
efficient_attn_ratio_list.append(log_info['sft_train_metrics']['efficient_attn_ratio'])
self.assertTrue(all([ratio > 0 for ratio in efficient_attn_ratio_list]))

Expand Down Expand Up @@ -316,7 +315,7 @@ def test_rl_train_with_sft(self):
train_controller.onload(target="all")
log_infos = train_controller.fit(data_batches, pack_max_length=1024, rollout_idx=1)
new_efficient_attn_ratio_list = []
for log_info in log_infos:
for log_info in log_infos["worker_log_infos"]:
new_efficient_attn_ratio_list.append(log_info['sft_train_metrics']['efficient_attn_ratio'])

efficient_attn_ratio_list.sort()
Expand Down
9 changes: 8 additions & 1 deletion tests/rl/test_rl_disaggregated_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,7 +142,14 @@ def _make_trainer(self, agent_loop_manager):
trainer._maybe_save_hf = MagicMock()
trainer._checkpoint_no_save_replay_buffer = False
trainer.train_controller = SimpleNamespace(
fit=MagicMock(return_value=[{"train_metrics": [], "sft_train_metrics": {}}]),
fit=MagicMock(
return_value={
"worker_log_infos": [{"train_metrics": [], "sft_train_metrics": {}}],
"padding_tokens": 0,
"pack_time": 0.0,
"train_time": 0.0,
}
),
onload=MagicMock(return_value="onload"),
offload=MagicMock(return_value="offload"),
weight_update=MagicMock(return_value="update"),
Expand Down
11 changes: 7 additions & 4 deletions tests/rl/test_rl_trainer_checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -152,16 +152,19 @@ def weight_update(self):

def fit(self, data_batches, pack_max_length: int, rollout_idx: int):
self.fit_steps.append(rollout_idx)
return [
{
return {
"worker_log_infos": [{
"rollout_is_metrics": {},
"mismatch_metrics": {},
"rollout_entropy": 0.0,
"train_entropy": 0.0,
"train_metrics": [],
"sft_train_metrics": {},
}
]
}],
"padding_tokens": 0,
"pack_time": 0.0,
"train_time": 0.0,
}

def save(self, checkpoint_path: str, no_save_optimizer: bool):
path = Path(checkpoint_path)
Expand Down
37 changes: 37 additions & 0 deletions tests/rl/test_training_worker_rank.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,37 @@
from types import SimpleNamespace

import pytest

from xtuner.v1.rl.trainer.worker import TrainingWorker


class TestTrainingWorkerDataParallelRank:
@pytest.mark.parametrize(
("rank", "tp_size", "sp_size", "expected_dp_rank"),
[
(0, 1, 1, 0),
(3, 1, 1, 3),
(0, 2, 1, 0),
(1, 2, 1, 0),
(2, 2, 1, 1),
(3, 2, 1, 1),
(0, 2, 2, 0),
(3, 2, 2, 0),
(4, 2, 2, 1),
(7, 2, 2, 1),
],
)
def test_get_dp_rank_accounts_for_all_data_replicas(
self,
rank: int,
tp_size: int,
sp_size: int,
expected_dp_rank: int,
) -> None:
worker = SimpleNamespace(
rank=rank,
_engine=SimpleNamespace(data_replicate_size=tp_size),
sp_mesh=SimpleNamespace(size=lambda: sp_size),
)

assert TrainingWorker.get_dp_rank(worker) == expected_dp_rank
4 changes: 2 additions & 2 deletions xtuner/v1/rl/trainer/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,13 @@
compute_rollout_importance_weights,
merge_rollout_is_metrics,
)
from .controller import ColateItem, TrainingController
from .controller import TrainingController, TrainingLogInfo
from .worker import TrainingWorker, WorkerConfig, WorkerInputItem, WorkerLogItem, WorkerTrainLogItem


__all__ = [
"ColateItem",
"TrainingController",
"TrainingLogInfo",
"RolloutImportanceSampling",
"compute_rollout_importance_weights",
"compute_is_metrics",
Expand Down
Loading
Loading