Skip to content

Commit e9d1398

Browse files
zyongyecodex
andauthored
[Bugfix][Kimi K3] Enable deferred MoE finalization before weight loading (#53327)
Signed-off-by: Yongye Zhu <zyy1102000@gmail.com> Co-authored-by: OpenAI Codex <codex@openai.com>
1 parent b2db227 commit e9d1398

2 files changed

Lines changed: 90 additions & 9 deletions

File tree

tests/models/kimi_k3/test_latent_moe_tail.py

Lines changed: 79 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,8 @@
11
# SPDX-License-Identifier: Apache-2.0
22
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
33

4+
from types import SimpleNamespace
5+
46
import pytest
57
import ray
68
import torch
@@ -13,8 +15,13 @@
1315
multi_process_parallel,
1416
)
1517
from vllm.distributed import get_tp_group
18+
from vllm.model_executor.layers.fused_moe.experts.trtllm_mxfp4_moe import (
19+
TrtLlmMxfp4ExpertsMonolithic,
20+
)
1621
from vllm.model_executor.layers.fused_moe.moe_output import UnfinalizedMoEOutput
22+
from vllm.model_executor.layers.fused_moe.runner.moe_runner import MoERunner
1723
from vllm.model_executor.warmup.cutedsl_warmup import cutedsl_warmup
24+
from vllm.models.kimi_k3.nvidia import latent_moe_runner
1825
from vllm.models.kimi_k3.nvidia.ops.latent_moe_tail import KimiK3LatentMoETailOp
1926
from vllm.platforms import current_platform
2027

@@ -24,6 +31,78 @@
2431
TOP_K = 8
2532

2633

34+
def test_deferred_finalize_enabled_before_moe_kernel_setup(
35+
monkeypatch: pytest.MonkeyPatch,
36+
) -> None:
37+
class FakeMoEConfig:
38+
tp_size = 8
39+
dp_size = 1
40+
ep_size = 1
41+
pcp_size = 1
42+
is_sequence_parallel = False
43+
hidden_dim = LATENT_SIZE
44+
hidden_dim_unpadded = LATENT_SIZE
45+
experts_per_token = 16
46+
defer_moe_finalize = False
47+
defer_moe_finalize_max_num_tokens = -1
48+
49+
@property
50+
def use_deferred_moe_finalize(self) -> bool:
51+
return self.defer_moe_finalize
52+
53+
moe_config = FakeMoEConfig()
54+
quant_method = SimpleNamespace(
55+
experts_cls=TrtLlmMxfp4ExpertsMonolithic,
56+
moe_kernel=None,
57+
)
58+
norm_weight = torch.empty(LATENT_SIZE, dtype=torch.bfloat16)
59+
transform = SimpleNamespace(
60+
norm=SimpleNamespace(weight=norm_weight, variance_epsilon=EPS),
61+
up_proj=SimpleNamespace(
62+
weight=SimpleNamespace(shape=(HIDDEN_SIZE, LATENT_SIZE))
63+
),
64+
)
65+
66+
def fake_runner_init(runner, *args, **kwargs) -> None:
67+
runner.moe_config = moe_config
68+
runner.routed_experts = SimpleNamespace(quant_method=quant_method)
69+
runner._shared_experts = object()
70+
runner.routed_output_transform = transform
71+
72+
initialized_with: dict[str, object] = {}
73+
tail_op = SimpleNamespace(contract=SimpleNamespace(max_num_tokens=128))
74+
75+
def fake_tail_initialize(**kwargs):
76+
initialized_with.update(kwargs)
77+
return tail_op
78+
79+
monkeypatch.setattr(MoERunner, "__init__", fake_runner_init)
80+
monkeypatch.setattr(latent_moe_runner.torch.cuda, "Event", lambda: object())
81+
monkeypatch.setattr(
82+
latent_moe_runner,
83+
"current_platform",
84+
SimpleNamespace(
85+
is_cuda=lambda: True,
86+
is_device_capability_family=lambda capability: capability == 100,
87+
),
88+
)
89+
monkeypatch.setattr(
90+
latent_moe_runner,
91+
"get_current_vllm_config",
92+
lambda: SimpleNamespace(
93+
parallel_config=SimpleNamespace(use_ubatching=False),
94+
model_config=SimpleNamespace(enable_sleep_mode=False),
95+
),
96+
)
97+
monkeypatch.setattr(KimiK3LatentMoETailOp, "initialize", fake_tail_initialize)
98+
99+
latent_moe_runner.LatentMoERunner()
100+
101+
assert moe_config.defer_moe_finalize
102+
assert moe_config.defer_moe_finalize_max_num_tokens == 128
103+
assert initialized_with["experts_per_token"] == 16
104+
105+
27106
def _make_deferred_routed_output(
28107
num_tokens: int,
29108
device: torch.device,

vllm/models/kimi_k3/nvidia/latent_moe_runner.py

Lines changed: 11 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -113,16 +113,12 @@ def __init__(
113113
)
114114

115115
experts_cls = getattr(self._quant_method, "experts_cls", None)
116-
moe_kernel = self._quant_method.moe_kernel
116+
# The kernel instance is built after weight loading, while the tail
117+
# must register its CuTeDSL warmup units during runner construction.
117118
self.moe_config.defer_moe_finalize = (
118-
(
119-
experts_cls is TrtLlmMxfp4ExpertsMonolithic
120-
or experts_cls is TrtLlmNvFp4ExpertsMonolithic
121-
)
122-
and self.moe_config.hidden_dim == self.moe_config.hidden_dim_unpadded
123-
and moe_kernel is not None
124-
and moe_kernel.supports_deferred_moe_finalize()
125-
)
119+
experts_cls is TrtLlmMxfp4ExpertsMonolithic
120+
or experts_cls is TrtLlmNvFp4ExpertsMonolithic
121+
) and self.moe_config.hidden_dim == self.moe_config.hidden_dim_unpadded
126122
from vllm.models.kimi_k3.nvidia.ops.latent_moe_tail import (
127123
KimiK3LatentMoETailOp,
128124
)
@@ -145,6 +141,12 @@ def __init__(
145141
if self.moe_config.use_deferred_moe_finalize
146142
else -1
147143
)
144+
if self.moe_config.use_deferred_moe_finalize:
145+
logger.info_once(
146+
"K3 latent-MoE tail fusion with deferred top-k finalization "
147+
"is enabled for up to %d tokens.",
148+
op.contract.max_num_tokens,
149+
)
148150
else:
149151
self.moe_config.defer_moe_finalize = False
150152
self.moe_config.defer_moe_finalize_max_num_tokens = -1

0 commit comments

Comments
 (0)