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
116 changes: 116 additions & 0 deletions tests/rl/test_routed_experts_dtype.py
Original file line number Diff line number Diff line change
@@ -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)
20 changes: 17 additions & 3 deletions xtuner/v1/module/decoder_layer/moe_decoder_layer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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."""

Expand Down Expand Up @@ -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)
Expand Down
34 changes: 33 additions & 1 deletion xtuner/v1/rl/trainer/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,32 @@

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(
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)


def calculate_entropy(
Expand Down Expand Up @@ -522,6 +548,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):
Expand All @@ -537,6 +564,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:
Expand Down Expand Up @@ -566,7 +594,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
)
Expand All @@ -585,6 +616,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

Expand Down
Loading