From dab9e7b9fc3f7833fbee41dc0e3a3dd0748bd294 Mon Sep 17 00:00:00 2001 From: matrix72c Date: Wed, 26 Aug 2026 00:12:56 +0800 Subject: [PATCH 1/2] fix(rl): store routed experts as uint16 --- tests/rl/test_routed_experts_dtype.py | 116 ++++++++++++++++++ .../module/decoder_layer/moe_decoder_layer.py | 20 ++- xtuner/v1/rl/trainer/worker.py | 35 +++++- 3 files changed, 167 insertions(+), 4 deletions(-) create mode 100644 tests/rl/test_routed_experts_dtype.py diff --git a/tests/rl/test_routed_experts_dtype.py b/tests/rl/test_routed_experts_dtype.py new file mode 100644 index 0000000000..28aee3b350 --- /dev/null +++ b/tests/rl/test_routed_experts_dtype.py @@ -0,0 +1,116 @@ +from types import SimpleNamespace +from unittest.mock import patch + +import numpy as np +import pytest +import ray +import torch + +from xtuner.v1.module.decoder_layer.moe_decoder_layer import _prepare_rollout_routed_experts_for_router +from xtuner.v1.rl.trainer.worker import ( + TrainingWorker, + _as_rollout_routed_experts_tensor, + _rollout_routed_experts_storage_dtype, +) + + +@pytest.mark.parametrize("n_routed_experts", [256, 257, 65536]) +def test_rollout_routed_experts_use_uint16_storage(n_routed_experts: int): + assert _rollout_routed_experts_storage_dtype(n_routed_experts) == torch.uint16 + + +def test_rollout_routed_experts_fall_back_to_long_above_uint16_capacity(): + assert _rollout_routed_experts_storage_dtype(65537) == torch.long + + +@pytest.mark.parametrize("expert_ids", [np.array([-1], dtype=np.int64), np.array([65536], dtype=np.int64)]) +def test_rollout_routed_experts_reject_uint16_overflow(expert_ids: np.ndarray): + with pytest.raises(ValueError, match="cannot be represented as uint16"): + _as_rollout_routed_experts_tensor(expert_ids, n_routed_experts=65536) + + +def _fake_worker(*, n_routed_experts: int, pack_max_length: int = 2): + language_cfg = SimpleNamespace( + n_routed_experts=n_routed_experts, + num_hidden_layers=2, + num_experts_per_tok=2, + ) + config = SimpleNamespace( + model_cfg=language_cfg, + free_rollout_routed_experts_in_worker=False, + pack_max_length=pack_max_length, + ) + return SimpleNamespace(config=config, sp_mesh=None) + + +def test_training_worker_keeps_ray_routes_and_padding_in_uint16(): + worker = _fake_worker(n_routed_experts=65536) + route_ref = ray.ObjectRef(bytes(28)) + routed_experts = np.array([[[0, 255], [256, 65535]]], dtype=np.uint16) + padding_marker = torch.empty(1) + seq_ctx = SimpleNamespace( + input_ids=torch.zeros((1, 2), dtype=torch.long), + rollout_routed_experts=[route_ref, padding_marker], + ) + + with patch("xtuner.v1.rl.trainer.worker.ray.get", return_value=routed_experts): + TrainingWorker._add_rollout_routed_experts(worker, seq_ctx, seq_ctx.rollout_routed_experts) + + assert seq_ctx.rollout_routed_experts.dtype == torch.uint16 + torch.testing.assert_close( + seq_ctx.rollout_routed_experts[0].long(), + torch.tensor([[0, 255], [256, 65535]], dtype=torch.long), + ) + + +def test_training_worker_uses_uint16_for_full_padding_batch(): + worker = _fake_worker(n_routed_experts=257) + seq_ctx = SimpleNamespace( + input_ids=torch.zeros((1, 2), dtype=torch.long), + rollout_routed_experts=torch.empty(0), + ) + + TrainingWorker._add_rollout_routed_experts(worker, seq_ctx, seq_ctx.rollout_routed_experts) + + assert seq_ctx.rollout_routed_experts.dtype == torch.uint16 + assert seq_ctx.rollout_routed_experts.shape == (2, 2, 2) + + +def test_training_worker_preserves_long_fallback(): + worker = _fake_worker(n_routed_experts=65537, pack_max_length=1) + route_ref = ray.ObjectRef(bytes(28)) + routed_experts = np.array([[[65536, 0], [1, 2]]], dtype=np.int64) + seq_ctx = SimpleNamespace( + input_ids=torch.zeros((1, 1), dtype=torch.long), + rollout_routed_experts=[route_ref], + ) + + with patch("xtuner.v1.rl.trainer.worker.ray.get", return_value=routed_experts): + TrainingWorker._add_rollout_routed_experts(worker, seq_ctx, seq_ctx.rollout_routed_experts) + + assert seq_ctx.rollout_routed_experts.dtype == torch.long + assert seq_ctx.rollout_routed_experts[0, 0, 0].item() == 65536 + + +@pytest.mark.parametrize("offload_rollout_routed_experts", [False, True]) +def test_uint16_layer_slice_is_converted_to_router_index_dtype(offload_rollout_routed_experts: bool): + stored_routes = torch.tensor( + [ + [[0, 255], [256, 65535]], + [[1, 2], [3, 4]], + ], + dtype=torch.uint16, + ) + layer_slice = stored_routes[:, 1, :] + assert not layer_slice.is_contiguous() + + router_routes = _prepare_rollout_routed_experts_for_router( + layer_slice, + torch.zeros((2, 4)), + offload_rollout_routed_experts=offload_rollout_routed_experts, + ) + + assert router_routes.dtype == torch.long + torch.testing.assert_close(router_routes, torch.tensor([[256, 65535], [3, 4]], dtype=torch.long)) + routing_weights = torch.zeros((2, 65536)) + routing_weights.gather(dim=1, index=router_routes) diff --git a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py index 64dbb2f37a..0ee3cab47b 100644 --- a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py @@ -47,6 +47,18 @@ HiddenStates: TypeAlias = torch.Tensor +def _prepare_rollout_routed_experts_for_router( + rollout_routed_experts: torch.Tensor, + hidden_states: torch.Tensor, + *, + offload_rollout_routed_experts: bool, +) -> torch.Tensor: + """Move one layer's routed-expert IDs to the router index dtype.""" + if offload_rollout_routed_experts and rollout_routed_experts.device != hidden_states.device: + rollout_routed_experts = rollout_routed_experts.contiguous() + return rollout_routed_experts.to(device=hidden_states.device, dtype=torch.long) + + class MoEDecoderLayerOutput(TypedDict): """Per-micro-batch outputs of one :class:`MoEDecoderLayer` forward.""" @@ -755,9 +767,11 @@ def _pre_moe_forward( rollout_routed_experts = seq_ctx.rollout_routed_experts[:, self.layer_idx, :] # seq_l, expert # TODO: pin_memory() + to(device, non_blocking=True) on a CUDA stream would allow overlapping the transfer # with prior-layer compute - if seq_ctx.offload_rollout_routed_experts and rollout_routed_experts.device != hidden_states.device: - rollout_routed_experts = rollout_routed_experts.contiguous() - rollout_routed_experts = rollout_routed_experts.to(hidden_states.device) + rollout_routed_experts = _prepare_rollout_routed_experts_for_router( + rollout_routed_experts, + hidden_states, + offload_rollout_routed_experts=seq_ctx.offload_rollout_routed_experts, + ) else: rollout_routed_experts = None router_results: RouterResults = self.gate(hidden_states, rollout_routed_experts) diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index 1cb2cbe65d..5bb4183e0b 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -66,6 +66,33 @@ DEVICE = get_device() DEVICE_MODULE = get_torch_device_module() +_UINT16_CAPACITY = 1 << 16 + + +def _rollout_routed_experts_storage_dtype(n_routed_experts: int) -> torch.dtype: + """Choose a compact lossless dtype for stored routed-expert IDs.""" + if n_routed_experts <= _UINT16_CAPACITY: + return torch.uint16 + return torch.long + + +def _as_rollout_routed_experts_tensor( + rollout_routed_experts: np.ndarray, + *, + n_routed_experts: int, +) -> torch.Tensor: + """Convert finalized routed-expert IDs to their CPU storage dtype.""" + storage_dtype = _rollout_routed_experts_storage_dtype(n_routed_experts) + if storage_dtype == torch.uint16 and rollout_routed_experts.dtype != np.uint16: + if rollout_routed_experts.size > 0: + min_expert_id = rollout_routed_experts.min() + max_expert_id = rollout_routed_experts.max() + if min_expert_id < 0 or max_expert_id >= _UINT16_CAPACITY: + raise ValueError( + "Routed-expert IDs cannot be represented as uint16: " + f"min={min_expert_id}, max={max_expert_id}" + ) + return torch.as_tensor(rollout_routed_experts, dtype=storage_dtype) def calculate_entropy( @@ -522,6 +549,7 @@ def _add_rollout_routed_experts( if isinstance(self.config.model_cfg, BaseComposeConfig) else self.config.model_cfg ) + storage_dtype = _rollout_routed_experts_storage_dtype(language_cfg.n_routed_experts) to_free_routed_expert_refs: list[ray.ObjectRef] = [] if isinstance(rollout_routed_experts, list): @@ -537,6 +565,7 @@ def _add_rollout_routed_experts( language_cfg.num_hidden_layers, language_cfg.num_experts_per_tok, ), + dtype=storage_dtype, ) out_rollout_routed_expert.append(rollout_routed_experts_tensor) else: @@ -566,7 +595,10 @@ def _add_rollout_routed_experts( if self.sp_mesh.get_local_rank() == 0: # only free once of sp mesh to_free_routed_expert_refs.append(rollout_routed_expert_refs) - rollout_routed_expert = torch.as_tensor(rollout_routed_expert, dtype=torch.long) + rollout_routed_expert = _as_rollout_routed_experts_tensor( + rollout_routed_expert, + n_routed_experts=language_cfg.n_routed_experts, + ) rollout_routed_expert = rollout_routed_expert.reshape( -1, language_cfg.num_hidden_layers, language_cfg.num_experts_per_tok ) @@ -585,6 +617,7 @@ def _add_rollout_routed_experts( language_cfg.num_hidden_layers, language_cfg.num_experts_per_tok, ), + dtype=storage_dtype, ) seq_ctx.rollout_routed_experts = rollout_routed_experts_tensor From 7d6a8d7ac91324d0cea6e5ea586ec2a8555d45fd Mon Sep 17 00:00:00 2001 From: matrix72c Date: Wed, 26 Aug 2026 09:31:05 +0800 Subject: [PATCH 2/2] style: format routed expert error message --- xtuner/v1/rl/trainer/worker.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/xtuner/v1/rl/trainer/worker.py b/xtuner/v1/rl/trainer/worker.py index 5bb4183e0b..7b69bd42f5 100644 --- a/xtuner/v1/rl/trainer/worker.py +++ b/xtuner/v1/rl/trainer/worker.py @@ -89,8 +89,7 @@ def _as_rollout_routed_experts_tensor( max_expert_id = rollout_routed_experts.max() if min_expert_id < 0 or max_expert_id >= _UINT16_CAPACITY: raise ValueError( - "Routed-expert IDs cannot be represented as uint16: " - f"min={min_expert_id}, max={max_expert_id}" + f"Routed-expert IDs cannot be represented as uint16: min={min_expert_id}, max={max_expert_id}" ) return torch.as_tensor(rollout_routed_experts, dtype=storage_dtype)