diff --git a/tests/module/test_greedy_router.py b/tests/module/test_greedy_router.py new file mode 100644 index 0000000000..1d9ccd451d --- /dev/null +++ b/tests/module/test_greedy_router.py @@ -0,0 +1,52 @@ +import torch + +from xtuner.v1.module.router.greedy import GreedyRouterConfig + + +def test_force_load_balance_is_opt_in(): + config = GreedyRouterConfig( + scoring_func="sigmoid", + router_scaling_factor=1.0, + norm_topk_prob=True, + ) + + assert config.force_load_balance is False + assert config.build(n_routed_experts=32, num_experts_per_tok=4).force_load_balance is False + + +def test_force_load_balance_randomizes_routing_and_preserves_gradient(): + torch.manual_seed(0) + logits = torch.zeros(16384, 32, requires_grad=True) + router = GreedyRouterConfig( + scoring_func="sigmoid", + router_scaling_factor=1.0, + norm_topk_prob=True, + force_load_balance=True, + ).build(n_routed_experts=32, num_experts_per_tok=4) + + router_results = router(logits) + counts = torch.bincount(router_results["topk_ids"].flatten(), minlength=32).float() + + assert not torch.equal(router_results["logits"], logits) + assert counts.max() / counts.mean() < 1.1 + router_results["router_weights"].sum().backward() + assert logits.grad is not None + assert torch.isfinite(logits.grad).all() + assert logits.grad.norm() > 0 + + +def test_force_load_balance_random_logits_is_torch_compile_compatible(): + torch.manual_seed(0) + logits = torch.zeros(8, 4) + router = GreedyRouterConfig( + scoring_func="sigmoid", + router_scaling_factor=1.0, + norm_topk_prob=True, + force_load_balance=True, + ).build(n_routed_experts=4, num_experts_per_tok=2) + compiled_router = torch.compile(router, backend="eager", fullgraph=True) + + router_results = compiled_router(logits) + + assert router_results["logits"].shape == logits.shape + assert not torch.equal(router_results["logits"], logits) diff --git a/tests/module/test_ultra_ep.py b/tests/module/test_ultra_ep.py new file mode 100644 index 0000000000..1067e887de --- /dev/null +++ b/tests/module/test_ultra_ep.py @@ -0,0 +1,553 @@ +import pytest +import torch + +from xtuner.v1.config import FSDPConfig +from xtuner.v1.float8.config import Float8Config, ScalingGranularity +from xtuner.v1.model.moe.moe import MoE +from xtuner.v1.model.moe.qwen3 import Qwen3MoE235BA22Config +from xtuner.v1.module.decoder_layer.moe_decoder_layer import ( + _UltraEPGradReduceJoin, + _UltraEPGradReduceStart, + _UltraEPWeightSyncForBackward, +) +from xtuner.v1.module.grouped_linear.moe_group_linear import GroupedLinear +from xtuner.v1.module.mtp import MTPConfig +from xtuner.v1.module.ultraep import UltraEPConfig +from xtuner.v1.module.ultraep import runtime as ultraep_runtime +from xtuner.v1.module.ultraep.runtime import UltraEPLayerRuntime, UltraEPManagerProvider +from xtuner.v1.ops.moe.cuda import group_gemm as group_gemm_module + + +class FakeGroup: + def __init__(self, size: int = 8): + self._size = size + + def size(self): + return self._size + + +class FakeGroupedLinear: + def __init__(self, shape): + self.weight = torch.nn.Parameter(torch.zeros(shape)) + self.configure_calls = [] + + def configure_ultra_ep_buffers(self, replica_weight, replica_grad): + self.configure_calls.append((replica_weight, replica_grad)) + + +class FakeEvent: + def __init__(self, calls, virtual_layer_id): + self.calls = calls + self.virtual_layer_id = virtual_layer_id + + def current_stream_wait(self): + self.calls.append(("wait", self.virtual_layer_id)) + + +class FakeLayerManager: + def __init__(self, *, num_master_experts=2, redundant=1, hidden_size=4, intermediate_size=3): + self.num_local_redundant_experts = redundant + self.local_replica_fc1_weight_buffer = torch.empty(redundant * 2 * intermediate_size * hidden_size) + self.local_replica_fc2_weight_buffer = torch.empty(redundant * hidden_size * intermediate_size) + self.local_replica_fc1_grad_buffer = torch.empty( + redundant * 2 * intermediate_size * hidden_size, + dtype=torch.float32, + ) + self.local_replica_fc2_grad_buffer = torch.empty( + redundant * hidden_size * intermediate_size, + dtype=torch.float32, + ) + self.master_fc1_grad_staging = torch.empty(num_master_experts, 2 * intermediate_size, hidden_size) + self.master_fc2_grad_staging = torch.empty(num_master_experts, hidden_size, intermediate_size) + self.register_calls = [] + self.refresh_calls = [] + self.weight_sync_calls = [] + self.stage_calls = [] + self.grad_reduce_calls = [] + self.restore_calls = [] + self.event_calls = [] + + def register_master_pointers(self, **kwargs): + self.register_calls.append(kwargs) + + def refresh_master_weight_pointers(self, **kwargs): + self.refresh_calls.append(kwargs) + + def weight_sync(self, layer_id, *, async_finish): + self.weight_sync_calls.append((layer_id, async_finish)) + return FakeEvent(self.event_calls, layer_id) + + def stage_master_gradients(self, *, virtual_layer_id, fc1_grad, fc2_grad): + self.stage_calls.append(virtual_layer_id) + self.master_fc1_grad_staging.copy_(fc1_grad) + self.master_fc2_grad_staging.copy_(fc2_grad) + + def grad_reduce(self, layer_id, *, async_finish): + self.grad_reduce_calls.append((layer_id, async_finish)) + self.master_fc1_grad_staging.add_(1.0) + self.master_fc2_grad_staging.add_(2.0) + return FakeEvent(self.event_calls, layer_id) + + def restore_master_gradients(self, *, virtual_layer_id, fc1_grad, fc2_grad): + self.restore_calls.append(virtual_layer_id) + fc1_grad.copy_(self.master_fc1_grad_staging) + fc2_grad.copy_(self.master_fc2_grad_staging) + + +class FakeManagerProvider: + num_model_layers = 4 + num_logical_experts = 16 + hidden_size = 4 + expert_intermediate_size = 3 + num_redundant_experts_per_rank = 1 + max_microbatches = 1 + + def __init__(self, manager): + self.manager = manager + self.get_manager_calls = 0 + + @property + def num_dispatch_experts(self): + return self.num_logical_experts + 8 * self.num_redundant_experts_per_rank + + def get_manager(self): + self.get_manager_calls += 1 + return self.manager + + +def make_fake_layer_runtime(): + manager = FakeLayerManager() + provider = FakeManagerProvider(manager) + fused_w1w3 = FakeGroupedLinear((2, 6, 4)) + fused_w2 = FakeGroupedLinear((2, 4, 3)) + runtime = UltraEPLayerRuntime( + layer_id=2, + manager_provider=provider, # type: ignore[arg-type] + fused_w1w3=fused_w1w3, + fused_w2=fused_w2, + ) + return runtime, manager, fused_w1w3, fused_w2 + + +def test_ultra_ep_is_opt_in(): + config = Qwen3MoE235BA22Config() + + assert config.ultraep_cfg is None + + +def test_ultra_ep_requires_redundant_experts(): + with pytest.raises(ValueError, match="num_redundant_experts_per_rank"): + UltraEPConfig(num_redundant_experts_per_rank=0) + + +@pytest.mark.parametrize( + ("overrides", "match"), + [ + ({"n_routed_experts": 33}, "n_routed_experts"), + ({"ep_size": 1}, "ep_size"), + ({"dispatcher": "all2all"}, "dispatcher='deepep'"), + ( + {"float8_cfg": Float8Config(scaling_granularity_grouped_gemm=ScalingGranularity.TILEWISE)}, + "BF16", + ), + ({"moe_bias": True}, "expert bias"), + ({"expert_tp_size": 2}, "expert_tp_size == 1"), + ({"mtp_config": MTPConfig(num_layers=1)}, "MTP expert layers"), + ], +) +def test_ultra_ep_rejects_unsupported_model_config(overrides, match): + kwargs = { + "ep_size": 8, + "n_routed_experts": 32, + "dispatcher": "deepep", + "ultraep_cfg": UltraEPConfig(num_redundant_experts_per_rank=1), + } + kwargs.update(overrides) + config = Qwen3MoE235BA22Config(**kwargs) + + with pytest.raises(ValueError, match=match): + config.build() + + +def test_ultra_ep_rejects_activation_recompute(): + model = object.__new__(MoE) + model.config = Qwen3MoE235BA22Config( + ep_size=8, + ultraep_cfg=UltraEPConfig(num_redundant_experts_per_rank=1), + ) + + with pytest.raises(ValueError, match="activation recompute"): + MoE.fully_shard(model, FSDPConfig(ep_size=8, recompute_ratio=0.5)) + + +def test_ultra_ep_manager_provider_derives_shape_from_xtuner_config(): + config = Qwen3MoE235BA22Config( + ultraep_cfg=UltraEPConfig( + num_redundant_experts_per_rank=2, + ) + ) + + provider = UltraEPManagerProvider.from_xtuner_config( + group=FakeGroup(), # type: ignore[arg-type] + config=config, + ) + + assert provider.num_model_layers == config.num_hidden_layers + assert provider.num_logical_experts == config.n_routed_experts + assert provider.hidden_size == config.hidden_size + assert provider.expert_intermediate_size == config.moe_intermediate_size + assert provider.num_redundant_experts_per_rank == config.ultraep_cfg.num_redundant_experts_per_rank + assert provider.max_microbatches == 1 + assert provider.num_dispatch_experts == config.n_routed_experts + 8 * config.ultraep_cfg.num_redundant_experts_per_rank + assert provider._manager is None + + runtime = UltraEPLayerRuntime( + layer_id=config.num_hidden_layers - 1, + manager_provider=provider, + fused_w1w3=object(), # type: ignore[arg-type] + fused_w2=object(), # type: ignore[arg-type] + ) + assert runtime.manager_provider is provider + assert provider._manager is None + + with pytest.raises(ValueError, match="layer_id"): + UltraEPLayerRuntime( + layer_id=config.num_hidden_layers, + manager_provider=provider, + fused_w1w3=object(), # type: ignore[arg-type] + fused_w2=object(), # type: ignore[arg-type] + ) + + +def test_ultra_ep_manager_registry_reuses_group_and_rejects_shape_mismatch(monkeypatch): + created = [] + + class FakeUltraEPManager: + def __init__(self, **kwargs): + self.kwargs = kwargs + created.append(self) + + monkeypatch.setattr(ultraep_runtime, "_MANAGERS", {}) + monkeypatch.setattr(ultraep_runtime, "UltraEPManager", FakeUltraEPManager) + group = FakeGroup() + + manager = ultraep_runtime.get_or_create_ultra_ep_manager( + group=group, # type: ignore[arg-type] + num_layers=20, + num_local_master_experts=4, + num_local_redundant_experts=1, + expert_fc1_numel=24, + expert_fc2_numel=12, + max_microbatches=1, + ) + same_manager = ultraep_runtime.get_or_create_ultra_ep_manager( + group=group, # type: ignore[arg-type] + num_layers=20, + num_local_master_experts=4, + num_local_redundant_experts=1, + expert_fc1_numel=24, + expert_fc2_numel=12, + max_microbatches=1, + ) + + assert same_manager is manager + assert created == [manager] + + with pytest.raises(RuntimeError, match="same UltraEP shape/configuration"): + ultraep_runtime.get_or_create_ultra_ep_manager( + group=group, # type: ignore[arg-type] + num_layers=21, + num_local_master_experts=4, + num_local_redundant_experts=1, + expert_fc1_numel=24, + expert_fc2_numel=12, + max_microbatches=1, + ) + + +def test_ultra_ep_buffers_are_not_parameters_or_state_dict_entries(): + linear = GroupedLinear(4, 6, 2) + parameter_names_before = tuple(name for name, _ in linear.named_parameters()) + state_dict_names_before = tuple(linear.state_dict()) + + linear.configure_ultra_ep_buffers( + torch.empty(1, 6, 4, dtype=torch.bfloat16), + torch.empty(1, 6, 4, dtype=torch.float32), + ) + + assert tuple(name for name, _ in linear.named_parameters()) == parameter_names_before + assert tuple(linear.state_dict()) == state_dict_names_before + assert parameter_names_before == ("weight",) + + +@pytest.mark.parametrize( + ("replica_weight_shape", "replica_grad_shape", "replica_grad_dtype", "match"), + [ + ((1, 5, 4), (1, 5, 4), torch.float32, "Unexpected UltraEP replica weight shape"), + ((1, 6, 4), (2, 6, 4), torch.float32, "FP32 tensor matching replica weight shape"), + ((1, 6, 4), (1, 6, 4), torch.bfloat16, "FP32 tensor matching replica weight shape"), + ], +) +def test_ultra_ep_rejects_invalid_replica_buffers( + replica_weight_shape, + replica_grad_shape, + replica_grad_dtype, + match, +): + linear = GroupedLinear(4, 6, 2) + + with pytest.raises(ValueError, match=match): + linear.configure_ultra_ep_buffers( + torch.empty(replica_weight_shape, dtype=torch.bfloat16), + torch.empty(replica_grad_shape, dtype=replica_grad_dtype), + ) + + +def test_ultra_ep_rejects_replica_buffers_when_expert_bias_is_enabled(): + linear = GroupedLinear(4, 6, 2, moe_bias=True) + + with pytest.raises(NotImplementedError, match="expert bias"): + linear.configure_ultra_ep_buffers( + torch.empty(1, 6, 4, dtype=torch.bfloat16), + torch.empty(1, 6, 4, dtype=torch.float32), + ) + + +def test_ultra_ep_layer_runtime_configures_buffers_and_refreshes_weight_pointers(): + runtime, manager, fused_w1w3, fused_w2 = make_fake_layer_runtime() + + runtime.sync_weights(7, async_finish=True) + assert len(fused_w1w3.configure_calls) == 1 + assert len(fused_w2.configure_calls) == 1 + assert fused_w1w3.configure_calls[0][0].shape == (1, 6, 4) + assert fused_w1w3.configure_calls[0][1].dtype == torch.float32 + assert fused_w2.configure_calls[0][0].shape == (1, 4, 3) + assert fused_w2.configure_calls[0][1].dtype == torch.float32 + assert len(manager.register_calls) == 1 + assert manager.register_calls[0]["layer_id"] == 2 + assert manager.refresh_calls[0]["fc1_weight"] is fused_w1w3.weight + assert manager.refresh_calls[0]["fc2_weight"] is fused_w2.weight + + runtime.sync_weights(8, async_finish=False) + assert len(fused_w1w3.configure_calls) == 1 + assert len(fused_w2.configure_calls) == 1 + assert len(manager.register_calls) == 1 + assert [call["layer_id"] for call in manager.refresh_calls] == [2, 2] + assert manager.weight_sync_calls == [(7, True), (8, False)] + + +def test_ultra_ep_grad_reduce_lifecycle_stages_reduces_and_restores(): + runtime, manager, fused_w1w3, fused_w2 = make_fake_layer_runtime() + + with pytest.raises(RuntimeError, match="not started"): + runtime.finish_grad_reduce(3) + + with pytest.raises(RuntimeError, match="master gradients are unavailable"): + runtime.start_grad_reduce(3) + + fused_w1w3.weight.grad = torch.full_like(fused_w1w3.weight, 1.0) + fused_w2.weight.grad = torch.full_like(fused_w2.weight, 3.0) + + runtime.start_grad_reduce(3) + assert manager.stage_calls == [3] + assert manager.grad_reduce_calls == [(3, True)] + + runtime.finish_grad_reduce(3) + assert manager.event_calls == [("wait", 3)] + assert manager.restore_calls == [3] + torch.testing.assert_close(fused_w1w3.weight.grad, torch.full_like(fused_w1w3.weight, 2.0)) + torch.testing.assert_close(fused_w2.weight.grad, torch.full_like(fused_w2.weight, 5.0)) + assert runtime._grad_reduce_events == {} + + +def test_ultra_ep_grad_reduce_autograd_nodes_start_before_join(): + runtime = object.__new__(UltraEPLayerRuntime) + calls = [] + + def start_grad_reduce(virtual_layer_id): + calls.append(("start", virtual_layer_id)) + + def finish_grad_reduce(virtual_layer_id): + calls.append(("finish", virtual_layer_id)) + + runtime.start_grad_reduce = start_grad_reduce + runtime.finish_grad_reduce = finish_grad_reduce + + x = torch.ones(2, requires_grad=True) + joined = _UltraEPGradReduceJoin.apply(x, runtime, 11) + output = _UltraEPGradReduceStart.apply(joined, runtime, 11) + output.sum().backward() + + assert calls == [("start", 11), ("finish", 11)] + torch.testing.assert_close(x.grad, torch.ones_like(x)) + + +def test_ultra_ep_output_wrapper_restores_replica_weight_before_group_gemm_backward(monkeypatch): + dual_gemm_calls = [] + + def fake_m_grouped_gemm_dual_weight(x, master_weight, replica_weight, tokens_per_expert, *, trans_b): + dual_gemm_calls.append( + (master_weight.shape[0], replica_weight.shape[0], trans_b, replica_weight.data_ptr()) + ) + weight = torch.cat((master_weight, replica_weight), dim=0) + chunks = [] + offset = 0 + for expert_idx, count in enumerate(tokens_per_expert.tolist()): + x_chunk = x[offset : offset + count] + chunks.append(x_chunk @ (weight[expert_idx].T if trans_b else weight[expert_idx])) + offset += count + return torch.cat(chunks) + + def fake_k_grouped_gemm(grad_output, x, tokens_per_expert): + chunks = [] + offset = 0 + for count in tokens_per_expert.tolist(): + grad_chunk = grad_output[offset : offset + count] + x_chunk = x[offset : offset + count] + chunks.append(grad_chunk.T @ x_chunk) + offset += count + return torch.stack(chunks) + + monkeypatch.setattr(group_gemm_module, "m_grouped_gemm_dual_weight", fake_m_grouped_gemm_dual_weight) + monkeypatch.setattr(group_gemm_module, "k_grouped_gemm", fake_k_grouped_gemm) + + x = torch.tensor([[1.0, 1.0], [2.0, 1.0]], requires_grad=True) + master_weight = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]], requires_grad=True) + replica_weight = torch.tensor([[[5.0, 6.0], [7.0, 8.0]]]) + original_replica_weight = replica_weight.clone() + replica_grad = torch.empty_like(replica_weight, dtype=torch.float32) + tokens_per_expert = torch.tensor([1, 1]) + + # The production wrapper invokes Manager.weight_sync during backward. + # A bare runtime instance is sufficient to validate autograd ordering. + runtime = object.__new__(UltraEPLayerRuntime) + sync_calls = [] + + def fake_sync_weights(virtual_layer_id, *, async_finish): + sync_calls.append((virtual_layer_id, async_finish)) + replica_weight.copy_(original_replica_weight) + + runtime.sync_weights = fake_sync_weights + + output = group_gemm_module.ultra_ep_group_gemm( + x, + master_weight, + replica_weight, + replica_grad, + tokens_per_expert, + ) + # UltraEP reuses its persistent slots for the next layer before this + # layer's backward. The output-side node restores them before DGrad. + replica_weight.fill_(100.0) + _UltraEPWeightSyncForBackward.apply(output, runtime, 7).sum().backward() + + assert sync_calls == [(7, False)] + torch.testing.assert_close(x.grad, torch.tensor([[4.0, 6.0], [12.0, 14.0]])) + torch.testing.assert_close(master_weight.grad, torch.tensor([[[1.0, 1.0], [1.0, 1.0]]])) + torch.testing.assert_close(replica_grad, torch.tensor([[[2.0, 1.0], [2.0, 1.0]]])) + assert dual_gemm_calls == [ + (1, 1, True, replica_weight.data_ptr()), + (1, 1, False, replica_weight.data_ptr()), + ] + + +def test_ultra_ep_group_gemm_supports_empty_local_dispatch(): + x = torch.empty(0, 2, requires_grad=True) + master_weight = torch.ones(1, 3, 2, requires_grad=True) + replica_weight = torch.ones(1, 3, 2) + replica_grad = torch.full_like(replica_weight, torch.nan, dtype=torch.float32) + + output = group_gemm_module.ultra_ep_group_gemm( + x, + master_weight, + replica_weight, + replica_grad, + torch.tensor([0, 0]), + ) + output.sum().backward() + + assert output.shape == (0, 3) + torch.testing.assert_close(x.grad, torch.empty_like(x)) + torch.testing.assert_close(master_weight.grad, torch.zeros_like(master_weight)) + torch.testing.assert_close(replica_grad, torch.zeros_like(replica_grad)) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA/Triton") +@pytest.mark.parametrize("trans_b", [True, False]) +def test_dual_weight_group_gemm_matches_contiguous_reference(trans_b): + torch.manual_seed(0) + device = torch.device("cuda") + counts = torch.tensor([128, 128, 128], device=device, dtype=torch.int64) + num_master, num_replica, n, k = 2, 1, 256, 256 + x = torch.randn(int(counts.sum()), k, device=device, dtype=torch.bfloat16) + if trans_b: + master = torch.randn(num_master, n, k, device=device, dtype=torch.bfloat16) + replica = torch.randn(num_replica, n, k, device=device, dtype=torch.bfloat16) + else: + master = torch.randn(num_master, k, n, device=device, dtype=torch.bfloat16) + replica = torch.randn(num_replica, k, n, device=device, dtype=torch.bfloat16) + + actual = group_gemm_module.m_grouped_gemm_dual_weight( + x, + master, + replica, + counts, + trans_b=trans_b, + ) + expected = group_gemm_module.m_grouped_gemm( + x, + torch.cat((master, replica), dim=0), + counts, + trans_b=trans_b, + ) + + torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA/Triton") +@pytest.mark.parametrize("counts_values", [(64, 64, 64), (0, 17, 33)]) +@pytest.mark.parametrize("shape", [(64, 64), (96, 64)]) +def test_dual_weight_group_gemm_backward_matches_contiguous_reference(counts_values, shape): + torch.manual_seed(0) + device = torch.device("cuda") + counts = torch.tensor(counts_values, device=device, dtype=torch.int64) + out_features, in_features = shape + num_master, num_replica = 2, 1 + + x = torch.randn(int(counts.sum()), in_features, device=device, dtype=torch.bfloat16, requires_grad=True) + master = torch.randn( + num_master, + out_features, + in_features, + device=device, + dtype=torch.bfloat16, + requires_grad=True, + ) + replica = torch.randn(num_replica, out_features, in_features, device=device, dtype=torch.bfloat16) + replica_grad = torch.full_like(replica, torch.nan, dtype=torch.float32) + + actual = group_gemm_module.ultra_ep_group_gemm( + x, + master, + replica, + replica_grad, + counts, + ) + + ref_x = x.detach().clone().requires_grad_(True) + ref_master = master.detach().clone().requires_grad_(True) + ref_replica = replica.detach().clone().requires_grad_(True) + expected = group_gemm_module.triton_group_gemm( + ref_x, + torch.cat((ref_master, ref_replica), dim=0), + counts, + ) + + grad_output = torch.randn_like(actual) + actual.backward(grad_output) + expected.backward(grad_output) + + torch.testing.assert_close(actual, expected, rtol=2e-2, atol=2e-2) + torch.testing.assert_close(x.grad, ref_x.grad, rtol=2e-2, atol=2e-2) + torch.testing.assert_close(master.grad, ref_master.grad, rtol=2e-2, atol=2e-2) + torch.testing.assert_close(replica_grad, ref_replica.grad.float(), rtol=2e-2, atol=2e-2) diff --git a/xtuner/v1/model/moe/moe.py b/xtuner/v1/model/moe/moe.py index b0111afa73..0774bc227a 100644 --- a/xtuner/v1/model/moe/moe.py +++ b/xtuner/v1/model/moe/moe.py @@ -64,6 +64,8 @@ from xtuner.v1.module.decoder_layer.dense_decoder_layer import DenseDecoderLayer from xtuner.v1.module.decoder_layer.moe_decoder_layer import MoEActFnConfig, MoEBlock, MoEDecoderLayer, MoEGate from xtuner.v1.module.mtp import MTPBlock, MTPConfig, MTPLayer +from xtuner.v1.module.ultraep import UltraEPConfig +from xtuner.v1.module.ultraep.runtime import UltraEPManagerProvider from xtuner.v1.utils import ( get_device, get_logger, @@ -166,6 +168,9 @@ class MoEConfig(TransformerConfig): # Compose models call `self.embed_tokens` multiple times per step, so default to # keeping it unsharded after forward to avoid repeated all-gathers. embed_reshard_after_forward: bool = True + # ``None`` keeps the baseline route byte-for-byte independent of UltraEP. + # When configured, replicas remain runtime-owned rather than model state. + ultraep_cfg: UltraEPConfig | None = None def build(self) -> "MoE": from xtuner.v1.model.moe.moe import MoE @@ -178,6 +183,25 @@ def use_moe_ep_compile_cfg(config: MoEConfig) -> bool: return config.ep_size > 1 or config.expert_tp_size > 1 +def _validate_ultraep_model_config(config: MoEConfig, *, ep_size: int, expert_tp_size: int) -> None: + if config.ultraep_cfg is None: + return + if expert_tp_size != 1: + raise ValueError("Xtuner UltraEP requires expert_tp_size == 1") + if ep_size <= 1: + raise ValueError("UltraEP requires ep_size > 1") + if config.n_routed_experts % ep_size != 0: + raise ValueError("UltraEP requires n_routed_experts to be divisible by ep_size") + if config.dispatcher != "deepep": + raise ValueError("Xtuner UltraEP currently supports dispatcher='deepep' only") + if config.float8_cfg is not None: + raise ValueError("Xtuner UltraEP currently supports BF16 grouped experts only") + if config.moe_bias: + raise ValueError("Xtuner UltraEP does not currently support expert bias") + if config.mtp_config is not None: + raise ValueError("Xtuner UltraEP does not currently support MTP expert layers") + + class MoE(BaseModel): """Transformer decoder consisting of *config.num_hidden_layers* layers. Each layer is a [`InternLM3DecoderLayer`] @@ -190,6 +214,7 @@ class MoE(BaseModel): ep_mesh: DeviceMesh | None = None expert_tp_mesh: DeviceMesh | None = None ep_tp_mesh: DeviceMesh | None = None + ultraep_manager_provider: UltraEPManagerProvider | None = None def __init__(self, config: MoEConfig): # Concrete MoE configs override build(), so validate dispatcher support @@ -205,6 +230,7 @@ def __init__(self, config: MoEConfig): super().__init__(config) ep_size = config.ep_size if config.ep_size is not None else 1 expert_tp_size = config.expert_tp_size if config.expert_tp_size > 1 else 1 + _validate_ultraep_model_config(config, ep_size=ep_size, expert_tp_size=expert_tp_size) if ep_size > 1 or expert_tp_size > 1: world_size = dist.get_world_size() fsdp_size = world_size // (ep_size * expert_tp_size) @@ -239,6 +265,15 @@ def __init__(self, config: MoEConfig): self.expert_tp_mesh = None self.ep_tp_mesh = None + if config.ultraep_cfg is not None: + assert self.ep_mesh is not None + self.ultraep_manager_provider = UltraEPManagerProvider.from_xtuner_config( + group=self.ep_mesh.get_group(), + config=config, + ) + else: + self.ultraep_manager_provider = None + self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps, type=config.rms_norm_type) self.lm_head = LMHead(config.hidden_size, config.vocab_size, bias=False) @@ -1040,6 +1075,7 @@ def build_layers(self, config: MoEConfig) -> nn.ModuleDict: ep_mesh=self.ep_mesh, expert_tp_mesh=self.expert_tp_mesh, ep_tp_mesh=self.ep_tp_mesh, + ultraep_manager_provider=self.ultraep_manager_provider, ) if self.config.freeze_routers: layers[str(layer_idx)].gate.requires_grad_(False) @@ -1148,6 +1184,11 @@ def fully_shard( if fsdp_config.hsdp_sharding_size is not None and self.config.expert_tp_size > 1: raise NotImplementedError("HSDP with ExpertTP is not supported") + if self.config.ultraep_cfg is not None and fsdp_config.recompute_ratio > 0: + raise ValueError( + "Xtuner UltraEP does not support FSDP activation recompute yet; " + "set fsdp_config.recompute_ratio=0 or disable ultraep" + ) self.fsdp_config = fsdp_config assert self.fsdp_config.ep_size == self.config.ep_size self.mp_policy = MixedPrecisionPolicy( diff --git a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py index 862d09c030..d8de1fc546 100644 --- a/xtuner/v1/module/decoder_layer/moe_decoder_layer.py +++ b/xtuner/v1/module/decoder_layer/moe_decoder_layer.py @@ -35,6 +35,7 @@ ) from xtuner.v1.module.grouped_linear.moe_group_linear import build_grouped_linear from xtuner.v1.module.rope import RopeScalingConfig +from xtuner.v1.module.ultraep.runtime import UltraEPLayerRuntime, UltraEPManagerProvider from xtuner.v1.ops.act_fn import get_act_fn from xtuner.v1.utils import ForwardState @@ -47,6 +48,63 @@ HiddenStates: TypeAlias = torch.Tensor +class _UltraEPGradReduceStart(Function): + """Start replica-gradient reduction after expert and dispatch backward.""" + + @staticmethod + def forward(ctx, hidden_states: torch.Tensor, runtime: UltraEPLayerRuntime, virtual_layer_id: int): + ctx.runtime = runtime + ctx.virtual_layer_id = virtual_layer_id + return hidden_states + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): # type: ignore[override] + ctx.runtime.start_grad_reduce(ctx.virtual_layer_id) + return grad_output, None, None + + +class _UltraEPGradReduceJoin(Function): + """Join replica-gradient reduction after attention backward.""" + + @staticmethod + def forward(ctx, hidden_states: torch.Tensor, runtime: UltraEPLayerRuntime, virtual_layer_id: int): + ctx.runtime = runtime + ctx.virtual_layer_id = virtual_layer_id + return hidden_states + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): # type: ignore[override] + ctx.runtime.finish_grad_reduce(ctx.virtual_layer_id) + return grad_output, None, None + + +class _UltraEPWeightSyncForBackward(Function): + """Restore Manager-owned replica slots before expert DGrad runs. + + Replica weights are reusable communication buffers, not model parameters. A later virtual layer can overwrite them + after this layer's forward. This identity node sits immediately after expert compute, so its backward runs after + combine backward but before the grouped-GEMM backward that needs the replica weights for DGrad. + """ + + @staticmethod + def forward( + ctx, + expert_output: torch.Tensor, + runtime: UltraEPLayerRuntime, + virtual_layer_id: int, + ) -> torch.Tensor: + ctx.runtime = runtime + ctx.virtual_layer_id = virtual_layer_id + return expert_output + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): # type: ignore[override] + # The blocking form establishes the compute-stream dependency before + # UltraEPGroupedGemm reads the mutable slots for DGrad. + ctx.runtime.sync_weights(ctx.virtual_layer_id, async_finish=False) + return grad_output, None, None + + class MoEActFnProtocol(Protocol): def __call__(self, fused_x: torch.Tensor, split_dim: int = -1) -> torch.Tensor: ... @@ -233,6 +291,7 @@ def __init__( ep_mesh: DeviceMesh | None = None, expert_tp_mesh: DeviceMesh | None = None, ep_tp_mesh: DeviceMesh | None = None, + ultraep_manager_provider: UltraEPManagerProvider | None = None, ): super().__init__() self.ep_mesh = ep_mesh @@ -241,6 +300,7 @@ def __init__( self.n_routed_experts = n_routed_experts self.n_shared_experts = n_shared_experts self.hidden_factor = hidden_factor + self._ultraep: UltraEPLayerRuntime | None = None self.self_attn: MultiHeadAttention | MultiLatentAttention | GatedDeltaNet = attention_config.build( hidden_size=hidden_size, @@ -293,13 +353,23 @@ def __init__( moe_act_fn_cfg=moe_act_fn_cfg, ep_tp_mesh=ep_tp_mesh, ) + if ultraep_manager_provider is not None: + if ep_mesh is None: + raise ValueError("UltraEP requires an EP device mesh") + self._ultraep = UltraEPLayerRuntime( + layer_id=layer_idx, + manager_provider=ultraep_manager_provider, + fused_w1w3=self.experts.fused_w1w3, + fused_w2=self.experts.fused_w2, + ) # TODO: (yehaochen) Maybe should be replaced by build_dispatcher process_group = ep_mesh.get_group() if ep_mesh is not None else None tp_group = expert_tp_mesh.get_group() if expert_tp_mesh is not None else None ep_tp_group = ep_tp_mesh._flatten().get_group() if ep_tp_mesh is not None else None + num_dispatch_experts = self._ultraep.num_dispatch_experts if self._ultraep is not None else n_routed_experts self.dispatcher = build_dispatcher( dispatcher=dispatcher, - n_routed_experts=n_routed_experts, + n_routed_experts=num_dispatch_experts, ep_group=process_group, tp_group=tp_group, ep_tp_group=ep_tp_group, @@ -395,6 +465,12 @@ def _forward( seq_ctx: SequenceContext, position_embeddings: tuple[torch.Tensor, torch.Tensor], ) -> tuple[HiddenStates, RouterLogits, RouterWeights, RouterTopKIds]: + ultraep = self._ultraep + virtual_layer_id: int | None = None + if ultraep is not None: + virtual_layer_id = ultraep.allocate_virtual_layer_id() + hidden_states = _UltraEPGradReduceJoin.apply(hidden_states, ultraep, virtual_layer_id) + residual, hidden_states, router_results = self._pre_moe_forward( hidden_states=hidden_states, seq_ctx=seq_ctx, @@ -408,9 +484,19 @@ def _forward( # ProberList.before_dispatch( # self.layer_idx, hidden_states, router_results["topk_ids"], router_results["topk_weights"] # ) + dispatch_hidden_states = hidden_states + dispatch_topk_ids = router_results["topk_ids"] + weight_sync_event = None + if ultraep is not None: + assert virtual_layer_id is not None + ultraep.update_placement(dispatch_topk_ids, virtual_layer_id) + weight_sync_event = ultraep.sync_weights(virtual_layer_id, async_finish=True) + dispatch_topk_ids = ultraep.reroute(dispatch_topk_ids, virtual_layer_id) + dispatch_hidden_states = _UltraEPGradReduceStart.apply(hidden_states, ultraep, virtual_layer_id) + pre_dispatched = self.dispatcher.dispatch_preprocess( - hidden_states=hidden_states.view(-1, hidden_states.shape[-1]), - topk_ids=router_results["topk_ids"], + hidden_states=dispatch_hidden_states.view(-1, dispatch_hidden_states.shape[-1]), + topk_ids=dispatch_topk_ids, topk_weights=router_results["topk_weights"], ) dispatched = self.dispatcher.dispatch( @@ -429,11 +515,18 @@ def _forward( # post_dispatched.get("row_ids_map"), # type: ignore[arg-type] # dispatched["topk_weights"], # ) + if weight_sync_event is not None: + # Replica slots are first used by expert GEMMs. Deferring this + # wait overlaps their async refresh with DeepEP dispatch work. + weight_sync_event.current_stream_wait() experts_out = self.experts( post_dispatched["hidden_states"], post_dispatched["tokens_per_expert"], decoding=False, ) + if ultraep is not None: + assert virtual_layer_id is not None + experts_out = _UltraEPWeightSyncForBackward.apply(experts_out, ultraep, virtual_layer_id) # ProberList.before_combine( # self.layer_idx, # experts_out, @@ -498,6 +591,11 @@ def _micro_batch_forward( "All hidden states should have the same shape" ) intra_layer_micro_batch = len(hidden_states_list) + if self._ultraep is not None: + raise RuntimeError( + "Xtuner UltraEP currently supports intra_layer_micro_batch == 1 only; " + "per-microbatch replica weight/grad slots are not implemented yet" + ) residual_list: list[torch.Tensor] = [] router_results_list: list[RouterResults] = [] diff --git a/xtuner/v1/module/grouped_linear/moe_group_linear.py b/xtuner/v1/module/grouped_linear/moe_group_linear.py index e00a129e34..e98d89a9c9 100644 --- a/xtuner/v1/module/grouped_linear/moe_group_linear.py +++ b/xtuner/v1/module/grouped_linear/moe_group_linear.py @@ -8,6 +8,7 @@ from xtuner.v1.float8.config import Float8Config, ScalingGranularity from xtuner.v1.float8.float8_gmm_tile_wise import TileWiseFloat8GroupedLinear from xtuner.v1.ops import group_gemm +from xtuner.v1.ops.moe.cuda.group_gemm import ultra_ep_group_gemm from xtuner.v1.utils.interleaved_shard import InterleavedShard @@ -114,6 +115,8 @@ def __init__( self.weight = nn.Parameter(weight) self.moe_bias = moe_bias + self._ultra_ep_replica_weight: torch.Tensor | None = None + self._ultra_ep_replica_grad: torch.Tensor | None = None if self.moe_bias: if self.parallel_style == "column": # Keep column-parallel bias flattened like the weight's output dimension. This lets EP and Expert TP @@ -159,10 +162,41 @@ def __init__( else: self.bias = nn.Parameter(bias) + def configure_ultra_ep_buffers( + self, + replica_weight: torch.Tensor, + replica_grad: torch.Tensor, + ) -> None: + """Attach Manager-owned replica buffers without registering + parameters.""" + if self.moe_bias: + raise NotImplementedError("UltraEP does not currently support grouped expert bias") + if replica_weight.ndim != 3 or replica_weight.shape[1:] != (self.out_features, self.in_features): + raise ValueError( + "Unexpected UltraEP replica weight shape: " + f"expected [R, {self.out_features}, {self.in_features}], got {tuple(replica_weight.shape)}" + ) + if replica_grad.shape != replica_weight.shape or replica_grad.dtype != torch.float32: + raise ValueError("UltraEP replica grad must be an FP32 tensor matching replica weight shape") + # Tensor attributes are intentionally plain references: Manager owns + # their storage and they must stay out of state_dict()/parameters(). + self._ultra_ep_replica_weight = replica_weight + self._ultra_ep_replica_grad = replica_grad + def forward(self, x: torch.Tensor, tokens_per_expert: torch.Tensor, decoding: bool = False): weight = self.weight.to_local() if isinstance(self.weight, DTensor) else self.weight weight = weight.view(-1, self.local_out_features, self.local_in_features) - out = group_gemm(x, weight, tokens_per_expert) + if self._ultra_ep_replica_weight is None: + out = group_gemm(x, weight, tokens_per_expert) + else: + assert self._ultra_ep_replica_grad is not None + out = ultra_ep_group_gemm( + x, + weight, + self._ultra_ep_replica_weight, + self._ultra_ep_replica_grad, + tokens_per_expert, + ) if self.moe_bias: bias = self.bias.to_local() if isinstance(self.bias, DTensor) else self.bias diff --git a/xtuner/v1/module/router/greedy.py b/xtuner/v1/module/router/greedy.py index b1a34b1de0..99406b400f 100644 --- a/xtuner/v1/module/router/greedy.py +++ b/xtuner/v1/module/router/greedy.py @@ -11,6 +11,12 @@ from .protocol import RouterProtocol, RouterResults +def _apply_random_logits(logits: torch.Tensor) -> torch.Tensor: + """Apply the force-balanced benchmark proxy to router logits.""" + random_logits = torch.randn_like(logits) + return logits + (random_logits - logits).detach() + + class GreedyRouterConfig(BaseModel): model_config = ConfigDict(extra="forbid") scoring_func: Annotated[Literal["sigmoid", "softmax"], Parameter(group="router")] @@ -18,6 +24,7 @@ class GreedyRouterConfig(BaseModel): norm_topk_prob: Annotated[bool, Parameter(group="router")] use_grouped_router: bool = False router_n_groups: int | None = None + force_load_balance: bool = False def build( self, @@ -53,6 +60,7 @@ def __init__( norm_topk_prob: bool = True, scoring_func: Literal["sigmoid", "softmax"] = "softmax", router_scaling_factor: float = 1.0, + force_load_balance: bool = False, ): super().__init__() self.n_routed_experts = n_routed_experts @@ -60,12 +68,16 @@ def __init__( self.norm_topk_prob = norm_topk_prob self.scoring_func = scoring_func self.router_scaling_factor = router_scaling_factor + self.force_load_balance = force_load_balance def forward(self, logits: torch.Tensor, rollout_routed_experts: torch.Tensor | None = None) -> RouterResults: if os.getenv("XTUNER_ROUTER_DEBUG") == "true": noise = torch.randn_like(logits) * 50 logits = logits + noise + if self.force_load_balance: + logits = _apply_random_logits(logits) + # TODO: (yehaochen) Support sigmoid if self.scoring_func == "sigmoid": routing_weights = logits.sigmoid() @@ -87,7 +99,7 @@ def forward(self, logits: torch.Tensor, rollout_routed_experts: torch.Tensor | N # moe forward # (e, ) - tokens_per_expert = torch.histc(topk_ids, bins=self.n_routed_experts, min=0, max=self.n_routed_experts) + tokens_per_expert = torch.histc(topk_ids.float(), bins=self.n_routed_experts, min=0, max=self.n_routed_experts) return { "logits": logits, @@ -108,6 +120,7 @@ def __init__( norm_topk_prob: bool = True, scoring_func: Literal["sigmoid", "softmax"] = "softmax", router_scaling_factor: float = 1.0, + force_load_balance: bool = False, ): super().__init__( n_routed_experts=n_routed_experts, @@ -115,6 +128,7 @@ def __init__( norm_topk_prob=norm_topk_prob, scoring_func=scoring_func, router_scaling_factor=router_scaling_factor, + force_load_balance=force_load_balance, ) self.router_n_groups = router_n_groups @@ -123,6 +137,9 @@ def forward(self, logits: torch.Tensor, rollout_routed_experts: torch.Tensor | N noise = torch.randn_like(logits) * 50 logits = logits + noise + if self.force_load_balance: + logits = _apply_random_logits(logits) + # TODO: (yehaochen) Support sigmoid if self.scoring_func == "sigmoid": routing_weights = logits.sigmoid() @@ -163,7 +180,7 @@ def forward(self, logits: torch.Tensor, rollout_routed_experts: torch.Tensor | N # moe forward # (e, ) - tokens_per_expert = torch.histc(topk_ids, bins=self.n_routed_experts, min=0, max=self.n_routed_experts) + tokens_per_expert = torch.histc(topk_ids.float(), bins=self.n_routed_experts, min=0, max=self.n_routed_experts) return { "logits": logits, diff --git a/xtuner/v1/module/ultraep/__init__.py b/xtuner/v1/module/ultraep/__init__.py new file mode 100644 index 0000000000..883a660774 --- /dev/null +++ b/xtuner/v1/module/ultraep/__init__.py @@ -0,0 +1,8 @@ +"""UltraEP runtime integration for MoE token balancing.""" + +from .config import UltraEPConfig + + +__all__ = [ + "UltraEPConfig", +] diff --git a/xtuner/v1/module/ultraep/config.py b/xtuner/v1/module/ultraep/config.py new file mode 100644 index 0000000000..f3b810cb32 --- /dev/null +++ b/xtuner/v1/module/ultraep/config.py @@ -0,0 +1,29 @@ +"""Configuration contract for Xtuner's optional UltraEP execution path.""" + +from __future__ import annotations + +from typing import Annotated + +from cyclopts import Parameter +from pydantic import BaseModel, ConfigDict, model_validator + + +class UltraEPConfig(BaseModel): + """Runtime-only redundant-expert configuration. + + ``MoEConfig.ultraep_cfg is None`` is the only disabled state. Replica weights + and gradients remain owned by the UltraEP runtime; this configuration never + changes model parameters, optimizer state, or checkpoints. + """ + + model_config = ConfigDict(extra="forbid") + + num_redundant_experts_per_rank: Annotated[ + int, Parameter(help="UltraEP redundant-expert slots reserved on each EP rank") + ] + + @model_validator(mode="after") + def validate_redundant_experts(self) -> UltraEPConfig: + if self.num_redundant_experts_per_rank <= 0: + raise ValueError("num_redundant_experts_per_rank must be > 0") + return self diff --git a/xtuner/v1/module/ultraep/runtime.py b/xtuner/v1/module/ultraep/runtime.py new file mode 100644 index 0000000000..6a879683b6 --- /dev/null +++ b/xtuner/v1/module/ultraep/runtime.py @@ -0,0 +1,510 @@ +"""Optional UltraEP runtime integration for Xtuner MoE layers.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING, Protocol + +import torch +import torch.distributed as dist +from torch.distributed.tensor import DTensor + + +if TYPE_CHECKING: + from ultra_ep import EventHandle, Manager + + from xtuner.v1.model.moe.moe import MoEConfig + + +class UltraEPGroupedLinear(Protocol): + """The small grouped-linear surface owned by one UltraEP layer binding.""" + + weight: torch.Tensor + + def configure_ultra_ep_buffers(self, replica_weight: torch.Tensor, replica_grad: torch.Tensor) -> None: ... + + +class UltraEPManager: + """One UltraEP Manager and its shared replica slots per EP group.""" + + def __init__( + self, + *, + group: dist.ProcessGroup, + num_layers: int, + num_local_master_experts: int, + num_local_redundant_experts: int, + expert_fc1_numel: int, + expert_fc2_numel: int, + max_microbatches: int, + ) -> None: + try: + import ultra_ep + except ImportError as exc: + raise ImportError( + "UltraEP is enabled but its Python package/CUDA extension is unavailable. " + "Build UltraEP outside the Xtuner environment and prepend its build/lib.* directory to PYTHONPATH." + ) from exc + + if dist.get_world_size() != group.size(): + raise NotImplementedError( + "Xtuner UltraEP currently requires DP=1 (the EP group must cover the world); " + "FSDP gradient-reduction ordering for DP>1 is not implemented yet." + ) + + self.group = group + self.num_layers = num_layers + self.num_local_master_experts = num_local_master_experts + self.num_local_redundant_experts = num_local_redundant_experts + self.expert_fc1_numel = expert_fc1_numel + self.expert_fc2_numel = expert_fc2_numel + self.max_microbatches = max_microbatches + self.runtime: Manager = ultra_ep.Manager( + group=group, + num_layers=num_layers, + num_local_master_experts=num_local_master_experts, + num_local_redundant_experts=num_local_redundant_experts, + expert_fc1_numel=expert_fc1_numel, + expert_fc2_numel=expert_fc2_numel, + is_train=True, + explicitly_destroy=False, + max_microbatches=max_microbatches, + weight_data_dtype=torch.bfloat16, + grad_dtype=torch.float32, + ) + # UltraEP's native grad-reduce kernel is FP32-only, while Xtuner FSDP + # keeps BF16 gradients beside BF16 expert parameters. One staging + # pair is sufficient: backward joins each layer's async reduction + # before the preceding layer can start its own reduction. Reusing the + # pair avoids allocating FP32 master-grad copies for every layer. + device = torch.device("cuda", torch.cuda.current_device()) + self.master_fc1_grad_staging = torch.empty( + num_local_master_experts, + expert_fc1_numel, + dtype=torch.float32, + device=device, + ) + self.master_fc2_grad_staging = torch.empty( + num_local_master_experts, + expert_fc2_numel, + dtype=torch.float32, + device=device, + ) + self._staging_owner: int | None = None + self._master_weight_ptr_hosts: dict[int, tuple[torch.Tensor, torch.Tensor]] = {} + + @property + def local_replica_fc1_weight_buffer(self) -> torch.Tensor: + return self.runtime.local_replica_fc1_weight_buffer + + @property + def local_replica_fc2_weight_buffer(self) -> torch.Tensor: + return self.runtime.local_replica_fc2_weight_buffer + + @property + def local_replica_fc1_grad_buffer(self) -> torch.Tensor: + return self.runtime.local_replica_fc1_grad_buffer + + @property + def local_replica_fc2_grad_buffer(self) -> torch.Tensor: + return self.runtime.local_replica_fc2_grad_buffer + + def allocate_microbatch_slot(self, layer_id: int) -> int: + return self.runtime.allocate_microbatch_slot(layer_id) + + def update_placement_sparse(self, layer_id: int, logical_topk_ids: torch.Tensor) -> None: + self.runtime.update_placement_sparse(layer_id, logical_topk_ids) + + def reroute_sparse(self, layer_id: int, physical_topk_ids: torch.Tensor) -> None: + self.runtime.reroute_sparse(layer_id, physical_topk_ids) + + @staticmethod + def _local(tensor: torch.Tensor) -> torch.Tensor: + return tensor.to_local() if isinstance(tensor, DTensor) else tensor + + def stage_master_gradients( + self, + *, + virtual_layer_id: int, + fc1_grad: torch.Tensor, + fc2_grad: torch.Tensor, + ) -> tuple[torch.Tensor, torch.Tensor]: + """Cast Xtuner's local BF16 master grads into shared FP32 staging.""" + if self._staging_owner is not None: + raise RuntimeError( + "UltraEP FP32 master-grad staging is still owned by virtual layer " + f"{self._staging_owner}; attempted to start {virtual_layer_id}" + ) + local_fc1 = self._local(fc1_grad).view(self.num_local_master_experts, -1) + local_fc2 = self._local(fc2_grad).view(self.num_local_master_experts, -1) + if ( + local_fc1.numel() != self.master_fc1_grad_staging.numel() + or local_fc2.numel() != self.master_fc2_grad_staging.numel() + ): + raise ValueError("Master expert gradient shapes do not match UltraEP FP32 staging") + self.master_fc1_grad_staging.copy_(local_fc1) + self.master_fc2_grad_staging.copy_(local_fc2) + self._staging_owner = virtual_layer_id + return self.master_fc1_grad_staging, self.master_fc2_grad_staging + + def restore_master_gradients( + self, + *, + virtual_layer_id: int, + fc1_grad: torch.Tensor, + fc2_grad: torch.Tensor, + ) -> None: + """Cast the reduced FP32 staging tensors back to Xtuner FSDP grads.""" + if self._staging_owner != virtual_layer_id: + raise RuntimeError(f"UltraEP staging owner is {self._staging_owner}, not {virtual_layer_id}") + local_fc1 = self._local(fc1_grad).view(self.num_local_master_experts, -1) + local_fc2 = self._local(fc2_grad).view(self.num_local_master_experts, -1) + local_fc1.copy_(self.master_fc1_grad_staging) + local_fc2.copy_(self.master_fc2_grad_staging) + self._staging_owner = None + + def register_master_pointers( + self, + *, + layer_id: int, + fc1_weight: torch.Tensor, + fc2_weight: torch.Tensor, + fc1_grad: torch.Tensor, + fc2_grad: torch.Tensor, + ) -> None: + fc1_weight = self._local(fc1_weight).view(self.num_local_master_experts, -1) + fc2_weight = self._local(fc2_weight).view(self.num_local_master_experts, -1) + if fc1_weight.shape[1] != self.expert_fc1_numel or fc2_weight.shape[1] != self.expert_fc2_numel: + raise ValueError("Master expert weight shapes do not match the UltraEP Manager configuration") + + fc1_grad = self._local(fc1_grad).view(self.num_local_master_experts, -1) + fc2_grad = self._local(fc2_grad).view(self.num_local_master_experts, -1) + if fc1_grad.dtype != torch.float32 or fc2_grad.dtype != torch.float32: + raise TypeError(f"UltraEP requires FP32 master grads, got {fc1_grad.dtype} and {fc2_grad.dtype}") + fc1_grads = list(fc1_grad.unbind(0)) + fc2_grads = list(fc2_grad.unbind(0)) + + self.runtime.construct_local_master_ptr_pool( + layer_id=layer_id, + fc1_weights=list(fc1_weight.unbind(0)), + fc2_weights=list(fc2_weight.unbind(0)), + fc1_grads=fc1_grads, + fc2_grads=fc2_grads, + ) + # Xtuner FSDP may replace the local parameter storage between gradient + # accumulation microbatches even with reshard_after_forward=False. Keep + # reusable pinned host arrays so each weight_sync can refresh only the + # two device pointer arrays without rebuilding all weight/grad pools. + self._master_weight_ptr_hosts[layer_id] = ( + torch.empty( + self.num_local_master_experts, + dtype=torch.int64, + device="cpu", + pin_memory=True, + ), + torch.empty( + self.num_local_master_experts, + dtype=torch.int64, + device="cpu", + pin_memory=True, + ), + ) + + def refresh_master_weight_pointers( + self, + *, + layer_id: int, + fc1_weight: torch.Tensor, + fc2_weight: torch.Tensor, + ) -> None: + """Refresh only FSDP-movable master-weight addresses in cached + pools.""" + hosts = self._master_weight_ptr_hosts.get(layer_id) + if hosts is None: + raise RuntimeError(f"UltraEP master pointer pool for layer {layer_id} is not registered") + fc1_weight = self._local(fc1_weight).view(self.num_local_master_experts, -1) + fc2_weight = self._local(fc2_weight).view(self.num_local_master_experts, -1) + fc1_host, fc2_host = hosts + for expert_idx in range(self.num_local_master_experts): + fc1_host[expert_idx] = fc1_weight[expert_idx].data_ptr() + fc2_host[expert_idx] = fc2_weight[expert_idx].data_ptr() + + fc1_device = self.runtime.local_master_fc1_weight_ptr_pool[layer_id] + fc2_device = self.runtime.local_master_fc2_weight_ptr_pool[layer_id] + if fc1_device is None or fc2_device is None: + raise RuntimeError(f"UltraEP device pointer pool for layer {layer_id} is unavailable") + fc1_device.copy_(fc1_host, non_blocking=True) + fc2_device.copy_(fc2_host, non_blocking=True) + + def weight_sync(self, layer_id: int, *, async_finish: bool) -> EventHandle: + return self.runtime.weight_sync(layer_id=layer_id, async_finish=async_finish) + + def grad_reduce(self, layer_id: int, *, async_finish: bool) -> EventHandle: + return self.runtime.grad_reduce(layer_id=layer_id, async_finish=async_finish) + + +class UltraEPManagerProvider: + """Lazy process-group-level owner of one shared UltraEP Manager.""" + + def __init__( + self, + *, + group: dist.ProcessGroup, + num_model_layers: int, + num_logical_experts: int, + hidden_size: int, + expert_intermediate_size: int, + num_redundant_experts_per_rank: int, + max_microbatches: int, + ) -> None: + if group.size() <= 1: + raise ValueError("UltraEP requires an EP process group with size > 1") + if num_logical_experts % group.size() != 0: + raise ValueError("UltraEP requires logical experts to be evenly sharded by EP") + if num_model_layers <= 0: + raise ValueError("UltraEP requires num_model_layers > 0") + + self.group = group + self.num_model_layers = num_model_layers + self.num_logical_experts = num_logical_experts + self.hidden_size = hidden_size + self.expert_intermediate_size = expert_intermediate_size + self.num_redundant_experts_per_rank = num_redundant_experts_per_rank + self.max_microbatches = max_microbatches + self._manager: UltraEPManager | None = None + + @classmethod + def from_xtuner_config( + cls, + *, + group: dist.ProcessGroup, + config: MoEConfig, + ) -> UltraEPManagerProvider: + """Build a provider from Xtuner's model config for the single- + microbatch path.""" + ultraep_cfg = config.ultraep_cfg + if ultraep_cfg is None: + raise ValueError("UltraEP manager provider requires config.ultraep_cfg") + + return cls( + group=group, + num_model_layers=config.num_hidden_layers, + num_logical_experts=config.n_routed_experts, + hidden_size=config.hidden_size, + expert_intermediate_size=config.moe_intermediate_size, + num_redundant_experts_per_rank=ultraep_cfg.num_redundant_experts_per_rank, + max_microbatches=1, + ) + + @property + def num_dispatch_experts(self) -> int: + """Global physical-expert count expected by the dispatcher.""" + return self.num_logical_experts + self.group.size() * self.num_redundant_experts_per_rank + + def get_manager(self) -> UltraEPManager: + if self._manager is None: + self._manager = get_or_create_ultra_ep_manager( + group=self.group, + num_layers=self.num_model_layers, + num_local_master_experts=self.num_logical_experts // self.group.size(), + num_local_redundant_experts=self.num_redundant_experts_per_rank, + expert_fc1_numel=2 * self.expert_intermediate_size * self.hidden_size, + expert_fc2_numel=self.hidden_size * self.expert_intermediate_size, + max_microbatches=self.max_microbatches, + ) + return self._manager + + +class UltraEPLayerRuntime: + """Runtime-only UltraEP binding for one MoE layer. + + The decoder owns ordinary model modules and the autograd graph boundaries. This object owns every interaction with + the process-group-level UltraEP manager and never registers a tensor as model state. + """ + + def __init__( + self, + *, + layer_id: int, + manager_provider: UltraEPManagerProvider, + fused_w1w3: UltraEPGroupedLinear, + fused_w2: UltraEPGroupedLinear, + ) -> None: + if layer_id < 0 or layer_id >= manager_provider.num_model_layers: + raise ValueError(f"UltraEP layer_id must be in [0, {manager_provider.num_model_layers}), got {layer_id}") + + self.layer_id = layer_id + self.manager_provider = manager_provider + self.num_logical_experts = manager_provider.num_logical_experts + self.hidden_size = manager_provider.hidden_size + self.expert_intermediate_size = manager_provider.expert_intermediate_size + self.num_redundant_experts_per_rank = manager_provider.num_redundant_experts_per_rank + self.max_microbatches = manager_provider.max_microbatches + self.fused_w1w3 = fused_w1w3 + self.fused_w2 = fused_w2 + + self._buffers_configured = False + self._master_pointers_registered = False + self._grad_reduce_events: dict[int, tuple[object, torch.Tensor, torch.Tensor]] = {} + + @property + def num_dispatch_experts(self) -> int: + """Global physical-expert count expected by the dispatcher.""" + return self.manager_provider.num_dispatch_experts + + def validate_microbatch_capacity(self, requested_microbatches: int) -> None: + """Fail before allocation rather than silently reusing a virtual + slot.""" + if requested_microbatches > self.max_microbatches: + raise ValueError( + "UltraEP virtual-layer capacity is too small for this layer call: " + f"requested={requested_microbatches}, max_microbatches={self.max_microbatches}. " + "UltraEP capacity is resolved from Trainer/TrainEngine.intra_layer_micro_batch." + ) + + def allocate_virtual_layer_id(self) -> int: + """Allocate the UltraEP virtual-layer slot for this forward + microbatch.""" + return self._ensure_manager().allocate_microbatch_slot(self.layer_id) + + def update_placement( + self, + logical_topk_ids: torch.Tensor, + virtual_layer_id: int, + ) -> None: + """Build the replication placement from logical expert IDs.""" + self._ensure_manager().update_placement_sparse(virtual_layer_id, logical_topk_ids) + + def reroute(self, logical_topk_ids: torch.Tensor, virtual_layer_id: int) -> torch.Tensor: + """Return a dispatcher-only copy rewritten into physical expert IDs.""" + physical_topk_ids = logical_topk_ids.clone() + self._ensure_manager().reroute_sparse(virtual_layer_id, physical_topk_ids) + return physical_topk_ids + + def sync_weights(self, virtual_layer_id: int, *, async_finish: bool): + manager = self._ensure_manager() + manager.refresh_master_weight_pointers( + layer_id=self.layer_id, + fc1_weight=self.fused_w1w3.weight, + fc2_weight=self.fused_w2.weight, + ) + return manager.weight_sync(virtual_layer_id, async_finish=async_finish) + + def start_grad_reduce(self, virtual_layer_id: int) -> None: + if virtual_layer_id in self._grad_reduce_events: + raise RuntimeError(f"UltraEP virtual layer slot {virtual_layer_id} is still in use") + fc1_grad = self.fused_w1w3.weight.grad + fc2_grad = self.fused_w2.weight.grad + if fc1_grad is None or fc2_grad is None: + raise RuntimeError( + f"UltraEP master gradients are unavailable at layer {self.layer_id}; " + "the FSDP/autograd hook ordering is incompatible with replica grad-reduce" + ) + manager = self._ensure_manager() + manager.stage_master_gradients( + virtual_layer_id=virtual_layer_id, + fc1_grad=fc1_grad, + fc2_grad=fc2_grad, + ) + event = manager.grad_reduce(virtual_layer_id, async_finish=True) + self._grad_reduce_events[virtual_layer_id] = (event, fc1_grad, fc2_grad) + + def finish_grad_reduce(self, virtual_layer_id: int) -> None: + state = self._grad_reduce_events.pop(virtual_layer_id, None) + if state is None: + raise RuntimeError(f"UltraEP grad-reduce event for virtual layer {virtual_layer_id} was not started") + event, fc1_grad, fc2_grad = state + event.current_stream_wait() # type: ignore[attr-defined] + self._ensure_manager().restore_master_gradients( + virtual_layer_id=virtual_layer_id, + fc1_grad=fc1_grad, + fc2_grad=fc2_grad, + ) + + def _ensure_manager(self) -> UltraEPManager: + manager = self.manager_provider.get_manager() + if not self._buffers_configured: + redundant = manager.num_local_redundant_experts + self.fused_w1w3.configure_ultra_ep_buffers( + manager.local_replica_fc1_weight_buffer.view( + redundant, + 2 * self.expert_intermediate_size, + self.hidden_size, + ), + manager.local_replica_fc1_grad_buffer.view( + redundant, + 2 * self.expert_intermediate_size, + self.hidden_size, + ), + ) + self.fused_w2.configure_ultra_ep_buffers( + manager.local_replica_fc2_weight_buffer.view( + redundant, + self.hidden_size, + self.expert_intermediate_size, + ), + manager.local_replica_fc2_grad_buffer.view( + redundant, + self.hidden_size, + self.expert_intermediate_size, + ), + ) + self._buffers_configured = True + + if not self._master_pointers_registered: + manager.register_master_pointers( + layer_id=self.layer_id, + fc1_weight=self.fused_w1w3.weight, + fc2_weight=self.fused_w2.weight, + fc1_grad=manager.master_fc1_grad_staging, + fc2_grad=manager.master_fc2_grad_staging, + ) + self._master_pointers_registered = True + return manager + + +_MANAGERS: dict[int, tuple[tuple[int, ...], UltraEPManager]] = {} + + +def get_or_create_ultra_ep_manager( + *, + group: dist.ProcessGroup, + num_layers: int, + num_local_master_experts: int, + num_local_redundant_experts: int, + expert_fc1_numel: int, + expert_fc2_numel: int, + max_microbatches: int, +) -> UltraEPManager: + """Return the single Manager associated with this process-local EP + group.""" + signature = ( + group.size(), + num_layers, + num_local_master_experts, + num_local_redundant_experts, + expert_fc1_numel, + expert_fc2_numel, + max_microbatches, + ) + key = id(group) + cached = _MANAGERS.get(key) + if cached is not None: + cached_signature, manager = cached + if cached_signature != signature: + raise RuntimeError( + "All MoE layers sharing an EP group must use the same UltraEP shape/configuration: " + f"existing={cached_signature}, requested={signature}" + ) + return manager + + manager = UltraEPManager( + group=group, + num_layers=num_layers, + num_local_master_experts=num_local_master_experts, + num_local_redundant_experts=num_local_redundant_experts, + expert_fc1_numel=expert_fc1_numel, + expert_fc2_numel=expert_fc2_numel, + max_microbatches=max_microbatches, + ) + _MANAGERS[key] = (signature, manager) + return manager diff --git a/xtuner/v1/ops/moe/cuda/group_gemm.py b/xtuner/v1/ops/moe/cuda/group_gemm.py index bcd5313904..b3430f9ccd 100644 --- a/xtuner/v1/ops/moe/cuda/group_gemm.py +++ b/xtuner/v1/ops/moe/cuda/group_gemm.py @@ -2,7 +2,7 @@ import torch -from .triton_kernels import k_grouped_gemm, m_grouped_gemm +from .triton_kernels import k_grouped_gemm, m_grouped_gemm, m_grouped_gemm_dual_weight class GroupedGemm(torch.autograd.Function): @@ -35,3 +35,94 @@ def triton_group_gemm(x, w, tokens_per_expert): # put x and w to the pytorch graph return torch.matmul(x, w[0].T) return GroupedGemm.apply(x, w, tokens_per_expert) + + +class UltraEPGroupedGemm(torch.autograd.Function): + """Grouped GEMM over inherent and UltraEP replica experts in one launch. + + Replica weights and gradients are runtime-owned, cross-layer buffers. They + must not become model parameters or optimizer state. This bridge therefore + returns only the inherent-expert weight gradient to autograd and writes the + replica Wgrad into UltraEP's FP32 buffer as a side effect. + + Xtuner's FSDP-managed master weights and UltraEP's shared replica slots live + in separate allocations. A dual-base Triton kernel selects the right weight + allocation per physical expert while retaining one persistent GMM launch. + Replica slots are shared by all MoE layers and can be overwritten by a + later layer after this forward has finished. The MoE decoder must place + ``_UltraEPWeightSyncForBackward`` on the expert output: its backward + restores this virtual layer's replica slots before this Function performs + DGrad. The replica tensor intentionally stays outside + ``save_for_backward`` so that refresh is not treated as an invalid in-place + mutation, and so no full replica-buffer snapshot is required per forward. + """ + + @staticmethod + def forward( + ctx, + x: torch.Tensor, + master_weight: torch.Tensor, + replica_weight: torch.Tensor, + replica_grad: torch.Tensor, + tokens_per_expert: torch.Tensor, + ) -> torch.Tensor: + empty_input = x.shape[0] == 0 + if empty_input: + out = torch.matmul(x, master_weight[0].T) + else: + out = m_grouped_gemm_dual_weight( + x, + master_weight, + replica_weight, + tokens_per_expert, + trans_b=True, + ) + ctx.save_for_backward(x, master_weight, tokens_per_expert) + # This is a Manager-owned mutable buffer. The output-side autograd + # node restores it before DGrad reads it in backward. + ctx.replica_weight = replica_weight + ctx.num_master_experts = master_weight.shape[0] + ctx.replica_grad = replica_grad + ctx.empty_input = empty_input + return out + + @staticmethod + def backward(ctx, grad_output: torch.Tensor): # type: ignore[override] + x, master_weight, tokens_per_expert = ctx.saved_tensors + replica_weight = ctx.replica_weight + num_master_experts = ctx.num_master_experts + replica_grad = ctx.replica_grad + + if ctx.empty_input: + dx = torch.matmul(grad_output, master_weight[0]) + master_dw = torch.zeros_like(master_weight) + replica_grad.zero_() + else: + dx = m_grouped_gemm_dual_weight( + grad_output, + master_weight, + replica_weight, + tokens_per_expert, + trans_b=False, + ) + physical_dw = k_grouped_gemm(grad_output, x, tokens_per_expert) + master_dw = physical_dw[:num_master_experts] + replica_grad.copy_(physical_dw[num_master_experts:].float()) + return dx, master_dw, None, None, None + + +def ultra_ep_group_gemm( + x: torch.Tensor, + master_weight: torch.Tensor, + replica_weight: torch.Tensor, + replica_grad: torch.Tensor, + tokens_per_expert: torch.Tensor, +) -> torch.Tensor: + """Run one grouped GEMM over inherent and redundant physical experts.""" + return UltraEPGroupedGemm.apply( + x, + master_weight, + replica_weight, + replica_grad, + tokens_per_expert, + ) diff --git a/xtuner/v1/ops/moe/cuda/triton_kernels/__init__.py b/xtuner/v1/ops/moe/cuda/triton_kernels/__init__.py index 2214f0eddd..a629c571e5 100644 --- a/xtuner/v1/ops/moe/cuda/triton_kernels/__init__.py +++ b/xtuner/v1/ops/moe/cuda/triton_kernels/__init__.py @@ -12,19 +12,23 @@ if triton.__version__ >= "3.4.0": from .k_grouped_gemm_TMA_triton3_4 import k_grouped_gemm - from .m_grouped_gemm_TMA_triton3_4 import m_grouped_gemm + from .m_grouped_gemm_TMA_triton3_4 import m_grouped_gemm, m_grouped_gemm_dual_weight elif triton.__version__ >= "3.2.0": from .k_grouped_gemm_TMA import k_grouped_gemm from .m_grouped_gemm_TMA import m_grouped_gemm + + m_grouped_gemm_dual_weight = get_env_not_available_func(["triton>=3.4"]) else: env_not_available_func = get_env_not_available_func(["torch.accelerator", "triton"]) k_grouped_gemm = env_not_available_func m_grouped_gemm = env_not_available_func + m_grouped_gemm_dual_weight = env_not_available_func else: env_not_available_func = get_env_not_available_func(["torch.accelerator", "triton"]) k_grouped_gemm = env_not_available_func m_grouped_gemm = env_not_available_func + m_grouped_gemm_dual_weight = env_not_available_func -__all__ = ["k_grouped_gemm", "m_grouped_gemm"] +__all__ = ["k_grouped_gemm", "m_grouped_gemm", "m_grouped_gemm_dual_weight"] diff --git a/xtuner/v1/ops/moe/cuda/triton_kernels/m_grouped_gemm_TMA_triton3_4.py b/xtuner/v1/ops/moe/cuda/triton_kernels/m_grouped_gemm_TMA_triton3_4.py index f8d891a4ea..07edd00d74 100644 --- a/xtuner/v1/ops/moe/cuda/triton_kernels/m_grouped_gemm_TMA_triton3_4.py +++ b/xtuner/v1/ops/moe/cuda/triton_kernels/m_grouped_gemm_TMA_triton3_4.py @@ -243,6 +243,205 @@ def m_grouped_gemm_bNmajor_kernel( tl.store(c_ptrs, c, mask=mask) +@triton.autotune(configs=get_cuda_autotune_config(), key=["N", "K"]) +@triton.jit +def m_grouped_gemm_dual_bKmajor_kernel( + A, + B_master, + B_replica, + C, + pad_starts, + pad_ends, + group_starts, + group_ends, + m_indices_pad, + M_pad_ptr, + M, + B_MASTER_ROWS, + B_REPLICA_ROWS, + NUM_MASTER_GROUPS: tl.constexpr, + N: tl.constexpr, + K: tl.constexpr, + dtype_a: tl.constexpr, + dtype_b: tl.constexpr, + dtype_c: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr, +): + """B-K-major GMM over two disjoint expert-weight allocations.""" + dtypeA = tl.bfloat16 if dtype_a == 0 else tl.float16 + dtypeB = tl.bfloat16 if dtype_b == 0 else tl.float16 + dtypeC = tl.bfloat16 if dtype_c == 0 else tl.float16 + BLOCKS = tl.num_programs(axis=0) + start_pid = tl.program_id(axis=0) + M_pad = tl.load(M_pad_ptr) + num_pid_m = tl.cdiv(M_pad, BLOCK_M) + num_pid_n = tl.cdiv(N, BLOCK_N) + num_tiles = num_pid_m * num_pid_n + + a_ptr = A.to(tl.pointer_type(dtypeA)) + master_ptr = B_master.to(tl.pointer_type(dtypeB)) + replica_ptr = B_replica.to(tl.pointer_type(dtypeB)) + c_ptr = C.to(tl.pointer_type(dtypeC)) + + a_desc = tl.make_tensor_descriptor( + a_ptr, + shape=[M, K], + strides=[K, 1], + block_shape=[BLOCK_M, BLOCK_K], + ) + master_desc = tl.make_tensor_descriptor( + master_ptr, + shape=[B_MASTER_ROWS, K], + strides=[K, 1], + block_shape=[BLOCK_N, BLOCK_K], + ) + replica_desc = tl.make_tensor_descriptor( + replica_ptr, + shape=[B_REPLICA_ROWS, K], + strides=[K, 1], + block_shape=[BLOCK_N, BLOCK_K], + ) + + for tile_id in tl.range(start_pid, num_tiles, BLOCKS): + pid_m, pid_n = grouped_launch(tile_id, M_pad, N, BLOCK_M, BLOCK_N, GROUP_M) + + group = tl.load(m_indices_pad + pid_m).to(tl.int32) + pad_off = tl.load(pad_starts + group).to(tl.int32) + group_start = (tl.load(group_starts + group) + (pid_m * BLOCK_M - pad_off)).to(tl.int32) + group_end = tl.load(group_ends + group).to(tl.int32) + + offs_bn = (pid_n * BLOCK_N).to(tl.int32) + accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + + # ``group`` is uniform for the whole program tile. Branch once around + # the K loop, then use the matching TMA descriptor without materializing + # a concatenated physical-expert tensor. + if group < NUM_MASTER_GROUPS: + offs_k = 0 + for _ in tl.range(0, tl.cdiv(K, BLOCK_K)): + a = a_desc.load([group_start, offs_k]) + b = master_desc.load([group * N + offs_bn, offs_k]) + accumulator = tl.dot(a, b.T, acc=accumulator, input_precision="tf32x3") + offs_k += BLOCK_K + else: + replica_group = group - NUM_MASTER_GROUPS + offs_k = 0 + for _ in tl.range(0, tl.cdiv(K, BLOCK_K)): + a = a_desc.load([group_start, offs_k]) + b = replica_desc.load([replica_group * N + offs_bn, offs_k]) + accumulator = tl.dot(a, b.T, acc=accumulator, input_precision="tf32x3") + offs_k += BLOCK_K + + offs_m_range = group_start + tl.arange(0, BLOCK_M) + offs_n_range = (pid_n * BLOCK_N).to(tl.int32) + tl.arange(0, BLOCK_N) + mask = (offs_m_range[:, None] < group_end) & (offs_n_range[None, :] < N) + c_ptrs = c_ptr + offs_m_range[:, None].to(tl.int64) * N + offs_n_range[None, :].to(tl.int64) + tl.store(c_ptrs, accumulator.to(dtypeC), mask=mask) + + +@triton.autotune(configs=get_cuda_autotune_config(), key=["N", "K"]) +@triton.jit +def m_grouped_gemm_dual_bNmajor_kernel( + A, + B_master, + B_replica, + C, + pad_starts, + pad_ends, + group_starts, + group_ends, + m_indices_pad, + M_pad_ptr, + M, + B_MASTER_ROWS, + B_REPLICA_ROWS, + NUM_MASTER_GROUPS: tl.constexpr, + N: tl.constexpr, + K: tl.constexpr, + dtype_a: tl.constexpr, + dtype_b: tl.constexpr, + dtype_c: tl.constexpr, + BLOCK_M: tl.constexpr, + BLOCK_N: tl.constexpr, + BLOCK_K: tl.constexpr, + GROUP_M: tl.constexpr, +): + """B-N-major GMM over two disjoint expert-weight allocations.""" + dtypeA = tl.bfloat16 if dtype_a == 0 else tl.float16 + dtypeB = tl.bfloat16 if dtype_b == 0 else tl.float16 + dtypeC = tl.bfloat16 if dtype_c == 0 else tl.float16 + BLOCKS = tl.num_programs(axis=0) + start_pid = tl.program_id(axis=0) + M_pad = tl.load(M_pad_ptr) + num_pid_m = tl.cdiv(M_pad, BLOCK_M) + num_pid_n = tl.cdiv(N, BLOCK_N) + num_tiles = num_pid_m * num_pid_n + + a_ptr = A.to(tl.pointer_type(dtypeA)) + master_ptr = B_master.to(tl.pointer_type(dtypeB)) + replica_ptr = B_replica.to(tl.pointer_type(dtypeB)) + c_ptr = C.to(tl.pointer_type(dtypeC)) + + a_desc = tl.make_tensor_descriptor( + a_ptr, + shape=[M, K], + strides=[K, 1], + block_shape=[BLOCK_M, BLOCK_K], + ) + master_desc = tl.make_tensor_descriptor( + master_ptr, + shape=[B_MASTER_ROWS, N], + strides=[N, 1], + block_shape=[BLOCK_K, BLOCK_N], + ) + replica_desc = tl.make_tensor_descriptor( + replica_ptr, + shape=[B_REPLICA_ROWS, N], + strides=[N, 1], + block_shape=[BLOCK_K, BLOCK_N], + ) + + for tile_id in tl.range(start_pid, num_tiles, BLOCKS): + pid_m, pid_n = grouped_launch(tile_id, M_pad, N, BLOCK_M, BLOCK_N, GROUP_M) + + group = tl.load(m_indices_pad + pid_m).to(tl.int32) + pad_off = tl.load(pad_starts + group).to(tl.int32) + group_start = (tl.load(group_starts + group) + (pid_m * BLOCK_M - pad_off)).to(tl.int32) + group_end = tl.load(group_ends + group).to(tl.int32) + + offs_bn = (pid_n * BLOCK_N).to(tl.int32) + accumulator = tl.zeros((BLOCK_M, BLOCK_N), dtype=tl.float32) + + if group < NUM_MASTER_GROUPS: + offs_k = 0 + offs_bk = (group * K).to(tl.int32) + for _ in tl.range(0, tl.cdiv(K, BLOCK_K)): + a = a_desc.load([group_start, offs_k]) + b = master_desc.load([offs_bk, offs_bn]) + accumulator = tl.dot(a, b, acc=accumulator, input_precision="tf32x3") + offs_k += BLOCK_K + offs_bk += BLOCK_K + else: + replica_group = group - NUM_MASTER_GROUPS + offs_k = 0 + offs_bk = (replica_group * K).to(tl.int32) + for _ in tl.range(0, tl.cdiv(K, BLOCK_K)): + a = a_desc.load([group_start, offs_k]) + b = replica_desc.load([offs_bk, offs_bn]) + accumulator = tl.dot(a, b, acc=accumulator, input_precision="tf32x3") + offs_k += BLOCK_K + offs_bk += BLOCK_K + + offs_m_range = group_start + tl.arange(0, BLOCK_M) + offs_n_range = (pid_n * BLOCK_N).to(tl.int32) + tl.arange(0, BLOCK_N) + mask = (offs_m_range[:, None] < group_end) & (offs_n_range[None, :] < N) + c_ptrs = c_ptr + offs_m_range[:, None].to(tl.int64) * N + offs_n_range[None, :].to(tl.int64) + tl.store(c_ptrs, accumulator.to(dtypeC), mask=mask) + + @triton.jit def repeat_interleave_kernel( group_ptr, @@ -363,6 +562,116 @@ def _(A: Tensor, B: Tensor, size_per_group: torch.Tensor, trans_b: bool = False) return C +@torch.library.custom_op("moe::m_grouped_gemm_dual_weight", mutates_args=()) +def m_grouped_gemm_dual_weight( + A: Tensor, + B_master: Tensor, + B_replica: Tensor, + size_per_group: torch.Tensor, + trans_b: bool = False, +) -> Tensor: + """Grouped GEMM whose expert weights live in two contiguous allocations. + + Logical groups are ordered as ``[master experts, replica experts]``. The + persistent Triton kernel selects the appropriate TMA descriptor per group, + so all physical experts are still computed in one launch. + """ + assert A.dim() == 2 + assert B_master.dim() == 3 + assert B_replica.dim() == 3 + assert A.stride(-1) == 1, "Please make sure A is K-major" + assert B_master.is_contiguous(), "Master expert weights must be contiguous" + assert B_replica.is_contiguous(), "Replica expert weights must be contiguous" + assert B_master.dtype == B_replica.dtype + assert B_master.device == B_replica.device == A.device + + M, K = A.shape + if trans_b: + num_master, N, BK = B_master.shape + num_replica, replica_N, replica_BK = B_replica.shape + else: + num_master, BK, N = B_master.shape + num_replica, replica_BK, replica_N = B_replica.shape + + assert num_master > 0 and num_replica > 0 + assert BK == K and replica_BK == K, "K of A should be equal to K of both B tensors" + assert replica_N == N, "Master and replica expert output dimensions must match" + num_groups = num_master + num_replica + assert size_per_group.numel() == num_groups + C = A.new_empty(M, N) + + BLOCK_M = 128 + m_per_group_padding = triton.cdiv(size_per_group, BLOCK_M) * BLOCK_M + M_pad = m_per_group_padding.sum() + repeats = (m_per_group_padding // BLOCK_M).to(torch.int32) + m_indices_pad = torch.empty(M // BLOCK_M + num_groups, device=size_per_group.device, dtype=torch.int64) + repeat_interleave( + torch.arange(num_groups, device=size_per_group.device, dtype=torch.int32), + repeats, + repeats.cumsum(0), + m_indices_pad, + ) + + pad_start = m_per_group_padding.cumsum(0) - m_per_group_padding + pad_end = m_per_group_padding.cumsum(0) + group_end = size_per_group.cumsum(0) + group_start = group_end - size_per_group + NUM_SMS = torch.cuda.get_device_properties(A.device).multi_processor_count - SM_MARGIN + + dtype_mapping = {torch.bfloat16: 0, torch.float16: 1} + dtype_a = dtype_mapping.get(A.dtype, -1) + dtype_b = dtype_mapping.get(B_master.dtype, -1) + dtype_c = dtype_mapping.get(C.dtype, -1) + assert dtype_a >= 0 and dtype_b >= 0 and dtype_c >= 0, "Only BF16 and FP16 are supported" + + def grid(META): + return (NUM_SMS,) + + def alloc_fn(size: int, alignment: int, stream: Optional[int]): + return torch.empty(size, device=A.device, dtype=torch.int8) + + triton.set_allocator(alloc_fn) + kernel = m_grouped_gemm_dual_bKmajor_kernel if trans_b else m_grouped_gemm_dual_bNmajor_kernel + master_rows = num_master * (N if trans_b else K) + replica_rows = num_replica * (N if trans_b else K) + kernel[grid]( + A, + B_master, + B_replica, + C, + pad_start, + pad_end, + group_start, + group_end, + m_indices_pad, + M_pad, + M, + master_rows, + replica_rows, + num_master, + N, + K, + dtype_a, + dtype_b, + dtype_c, + BLOCK_M=BLOCK_M, + ) + return C + + +@m_grouped_gemm_dual_weight.register_fake +def _( + A: Tensor, + B_master: Tensor, + B_replica: Tensor, + size_per_group: torch.Tensor, + trans_b: bool = False, +) -> Tensor: + M, _ = A.shape + N = B_master.shape[1] if trans_b else B_master.shape[2] + return A.new_empty(M, N) + + if __name__ == "__main__": from torch.profiler import ProfilerActivity, profile from utils import generate_random_list, row_max_normalization