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
79 changes: 79 additions & 0 deletions tests/models/kimi_k3/test_latent_moe_tail.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,8 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from types import SimpleNamespace

import pytest
import ray
import torch
Expand All @@ -13,8 +15,13 @@
multi_process_parallel,
)
from vllm.distributed import get_tp_group
from vllm.model_executor.layers.fused_moe.experts.trtllm_mxfp4_moe import (
TrtLlmMxfp4ExpertsMonolithic,
)
from vllm.model_executor.layers.fused_moe.moe_output import UnfinalizedMoEOutput
from vllm.model_executor.layers.fused_moe.runner.moe_runner import MoERunner
from vllm.model_executor.warmup.cutedsl_warmup import cutedsl_warmup
from vllm.models.kimi_k3.nvidia import latent_moe_runner
from vllm.models.kimi_k3.nvidia.ops.latent_moe_tail import KimiK3LatentMoETailOp
from vllm.platforms import current_platform

Expand All @@ -24,6 +31,78 @@
TOP_K = 8


def test_deferred_finalize_enabled_before_moe_kernel_setup(
monkeypatch: pytest.MonkeyPatch,
) -> None:
class FakeMoEConfig:
tp_size = 8
dp_size = 1
ep_size = 1
pcp_size = 1
is_sequence_parallel = False
hidden_dim = LATENT_SIZE
hidden_dim_unpadded = LATENT_SIZE
experts_per_token = 16
defer_moe_finalize = False
defer_moe_finalize_max_num_tokens = -1

@property
def use_deferred_moe_finalize(self) -> bool:
return self.defer_moe_finalize

moe_config = FakeMoEConfig()
quant_method = SimpleNamespace(
experts_cls=TrtLlmMxfp4ExpertsMonolithic,
moe_kernel=None,
)
norm_weight = torch.empty(LATENT_SIZE, dtype=torch.bfloat16)
transform = SimpleNamespace(
norm=SimpleNamespace(weight=norm_weight, variance_epsilon=EPS),
up_proj=SimpleNamespace(
weight=SimpleNamespace(shape=(HIDDEN_SIZE, LATENT_SIZE))
),
)

def fake_runner_init(runner, *args, **kwargs) -> None:
runner.moe_config = moe_config
runner.routed_experts = SimpleNamespace(quant_method=quant_method)
runner._shared_experts = object()
runner.routed_output_transform = transform

initialized_with: dict[str, object] = {}
tail_op = SimpleNamespace(contract=SimpleNamespace(max_num_tokens=128))

def fake_tail_initialize(**kwargs):
initialized_with.update(kwargs)
return tail_op

monkeypatch.setattr(MoERunner, "__init__", fake_runner_init)
monkeypatch.setattr(latent_moe_runner.torch.cuda, "Event", lambda: object())
monkeypatch.setattr(
latent_moe_runner,
"current_platform",
SimpleNamespace(
is_cuda=lambda: True,
is_device_capability_family=lambda capability: capability == 100,
),
)
monkeypatch.setattr(
latent_moe_runner,
"get_current_vllm_config",
lambda: SimpleNamespace(
parallel_config=SimpleNamespace(use_ubatching=False),
model_config=SimpleNamespace(enable_sleep_mode=False),
),
)
monkeypatch.setattr(KimiK3LatentMoETailOp, "initialize", fake_tail_initialize)

latent_moe_runner.LatentMoERunner()

assert moe_config.defer_moe_finalize
assert moe_config.defer_moe_finalize_max_num_tokens == 128
assert initialized_with["experts_per_token"] == 16


def _make_deferred_routed_output(
num_tokens: int,
device: torch.device,
Expand Down
20 changes: 11 additions & 9 deletions vllm/models/kimi_k3/nvidia/latent_moe_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -113,16 +113,12 @@ def __init__(
)

experts_cls = getattr(self._quant_method, "experts_cls", None)
moe_kernel = self._quant_method.moe_kernel
# The kernel instance is built after weight loading, while the tail
# must register its CuTeDSL warmup units during runner construction.
self.moe_config.defer_moe_finalize = (
(
experts_cls is TrtLlmMxfp4ExpertsMonolithic
or experts_cls is TrtLlmNvFp4ExpertsMonolithic
)
and self.moe_config.hidden_dim == self.moe_config.hidden_dim_unpadded
and moe_kernel is not None
and moe_kernel.supports_deferred_moe_finalize()
)
experts_cls is TrtLlmMxfp4ExpertsMonolithic
or experts_cls is TrtLlmNvFp4ExpertsMonolithic
) and self.moe_config.hidden_dim == self.moe_config.hidden_dim_unpadded
from vllm.models.kimi_k3.nvidia.ops.latent_moe_tail import (
KimiK3LatentMoETailOp,
)
Expand All @@ -145,6 +141,12 @@ def __init__(
if self.moe_config.use_deferred_moe_finalize
else -1
)
if self.moe_config.use_deferred_moe_finalize:
logger.info_once(
"K3 latent-MoE tail fusion with deferred top-k finalization "
"is enabled for up to %d tokens.",
op.contract.max_num_tokens,
)
else:
self.moe_config.defer_moe_finalize = False
self.moe_config.defer_moe_finalize_max_num_tokens = -1
Expand Down
Loading