Compare commits

...
Author SHA1 Message Date
Woosuk Kwon 58b0c78a42 [MRV2] Support expert index capture
Signed-off-by: Woosuk Kwon <woosuk@inferact.ai>
2026-03-25 23:48:32 +00:00
5 changed files with 231 additions and 0 deletions
@@ -12,6 +12,8 @@ steps:
- tests/v1/engine/test_llm_engine.py
- tests/v1/e2e/
- tests/entrypoints/llm/test_struct_output_generate.py
- tests/model_executor/test_routed_experts_capture.py
- tests/v1/worker/test_gpu_model_runner_v2_eplb.py
commands:
- set -x
- export VLLM_USE_V2_MODEL_RUNNER=1
@@ -23,6 +25,8 @@ steps:
- pytest -v -s v1/e2e/general/test_min_tokens.py
# Temporary hack filter to exclude ngram spec decoding based tests.
- pytest -v -s entrypoints/llm/test_struct_output_generate.py -k "xgrammar and not speculative_config6 and not speculative_config7 and not speculative_config8 and not speculative_config0"
- pytest -v -s model_executor/test_routed_experts_capture.py -k "gpu_model_runner_v2_binds_router_capture"
- pytest -v -s v1/worker/test_gpu_model_runner_v2_eplb.py -k "test_v2_sample_tokens_saves_routed_experts"
- label: Model Runner V2 Examples
timeout_in_minutes: 45
@@ -158,3 +158,43 @@ def test_gpu_model_runner_binding_stage(monkeypatch):
assert callable(dummy_module.router.capture_fn)
dummy_module.router.capture_fn(torch.tensor([[9, 10]]))
assert len(capturer.calls) == 1
def test_gpu_model_runner_v2_binds_router_capture(monkeypatch):
from vllm.v1.worker.gpu.routed_experts_utils import RoutedExpertsCaptureHelper
class DummyFusedMoE:
def __init__(self):
self.layer_id = 13
self.router = _make_router()
class DummyCapturer:
def __init__(self):
self.calls = []
def capture(self, layer_id, topk_ids):
self.calls.append((layer_id, topk_ids))
dummy_module = DummyFusedMoE()
import vllm.model_executor.layers.fused_moe.layer as fused_moe_layer
monkeypatch.setattr(fused_moe_layer, "FusedMoE", DummyFusedMoE)
dummy_self = types.SimpleNamespace(
compilation_config=types.SimpleNamespace(
static_forward_context={"dummy": dummy_module}
),
routed_experts=RoutedExpertsCaptureHelper(),
)
capturer = DummyCapturer()
dummy_self.routed_experts.bind(dummy_self, capturer)
assert dummy_module.router.capture_fn is not None
dummy_module.router.capture_fn(torch.tensor([[7, 8]]))
assert len(capturer.calls) == 1
layer_id, topk_ids = capturer.calls[0]
assert layer_id == 13
assert torch.equal(topk_ids, torch.tensor([[7, 8]]))
@@ -9,6 +9,7 @@ import torch
from vllm.v1.worker.gpu import eplb_utils as eplb
from vllm.v1.worker.gpu import model_runner as mrv2
from vllm.v1.worker.gpu.routed_experts_utils import RoutedExpertsCaptureHelper
class FakeMemoryProfiler:
@@ -185,3 +186,74 @@ def test_v2_sample_tokens_runs_eplb_on_non_last_pp_rank(monkeypatch):
assert mrv2.GPUModelRunner.sample_tokens(runner, None) is None
assert events == ["postprocess", "eplb"]
def test_v2_sample_tokens_saves_routed_experts(monkeypatch):
saved_indices = []
class DummyAsyncOutput:
def __init__(self, **kwargs: Any):
self.kwargs = kwargs
def get_output(self):
return self.kwargs["model_runner_output"]
runner = _make_runner(
is_last_pp_rank=True,
use_pp=False,
use_async_scheduling=False,
main_stream="main",
output_copy_stream="copy",
output_copy_event="event",
)
runner.routed_experts = RoutedExpertsCaptureHelper()
runner.routed_experts._initialized = True
runner.routed_experts._slot_mapping = torch.tensor(
[3, 5], dtype=torch.int32
).numpy()
runner.execute_model_state = mrv2.ExecuteModelState(
input_batch=SimpleNamespace(
req_ids=["req-0"],
req_id_to_index={"req-0": 0},
idx_mapping_np=torch.tensor([0], dtype=torch.int32).numpy(),
idx_mapping=torch.tensor([0], dtype=torch.int32),
num_reqs=1,
),
attn_metadata=None,
slot_mappings_by_layer=None,
hidden_states=torch.zeros((1, 4)),
aux_hidden_states=None,
kv_connector_output=None,
num_tokens_across_dp=None,
)
runner.sample = lambda *args, **kwargs: (
SimpleNamespace(sampled_token_ids=torch.tensor([[42]], dtype=torch.long)),
torch.tensor([1], dtype=torch.int32),
torch.tensor([0], dtype=torch.int32),
)
runner.postprocess = lambda *args, **kwargs: None
runner.prompt_logprobs_worker = SimpleNamespace(
compute_prompt_logprobs=lambda *args, **kwargs: {}
)
runner.model = SimpleNamespace(compute_logits=lambda x: x)
runner.req_states = SimpleNamespace(
all_token_ids=SimpleNamespace(gpu=None),
num_computed_tokens=SimpleNamespace(gpu=None),
prompt_len=SimpleNamespace(np=None),
prefill_len=SimpleNamespace(np=None),
num_computed_prefill_tokens=None,
)
runner.eplb.step = lambda *args, **kwargs: None
monkeypatch.setattr(
mrv2.RoutedExpertsCaptureHelper,
"save",
lambda self: saved_indices.append(self._slot_mapping.copy()),
)
monkeypatch.setattr(mrv2, "AsyncOutput", DummyAsyncOutput)
output = mrv2.GPUModelRunner.sample_tokens(runner, None)
assert output.req_ids == ["req-0"]
assert len(saved_indices) == 1
assert (saved_indices[0] == runner.routed_experts._slot_mapping).all()
+9
View File
@@ -84,6 +84,7 @@ from vllm.v1.worker.gpu.mm.encoder_cache import EncoderCache
from vllm.v1.worker.gpu.model_states import init_model_state
from vllm.v1.worker.gpu.pool.pooling_runner import PoolingRunner
from vllm.v1.worker.gpu.pp_utils import pp_broadcast, pp_receive
from vllm.v1.worker.gpu.routed_experts_utils import RoutedExpertsCaptureHelper
from vllm.v1.worker.gpu.sample.output import SamplerOutput
from vllm.v1.worker.gpu.sample.prompt_logprob import PromptLogprobsWorker
from vllm.v1.worker.gpu.sample.sampler import Sampler
@@ -247,6 +248,7 @@ class GPUModelRunner(LoRAModelRunnerMixin):
# Expert parallelism load balancer.
self.eplb = EPLBController(self.parallel_config, self.device)
self.routed_experts = RoutedExpertsCaptureHelper()
def update_max_model_len(self, max_model_len: int) -> None:
self.max_model_len = max_model_len
@@ -389,6 +391,9 @@ class GPUModelRunner(LoRAModelRunnerMixin):
)
self.kv_connector = get_kv_connector(self.vllm_config, kv_caches_dict)
def init_routed_experts_capturer(self) -> None:
self.routed_experts.init(self)
@torch.inference_mode()
@step_eplb_after(is_dummy=True)
def _dummy_run(
@@ -916,6 +921,8 @@ class GPUModelRunner(LoRAModelRunnerMixin):
dummy_run: bool = False,
skip_attn_for_dummy_run: bool = False,
) -> ModelRunnerOutput | IntermediateTensors | None:
self.routed_experts.before_execute()
if not dummy_run:
# Update the request states.
self.finish_requests(scheduler_output)
@@ -1000,6 +1007,7 @@ class GPUModelRunner(LoRAModelRunnerMixin):
slot_mappings_by_layer = None
if not (dummy_run and skip_attn_for_dummy_run):
assert slot_mappings is not None
self.routed_experts.record_slot_mapping(slot_mappings, num_toks)
slot_mappings_by_layer = build_slot_mappings_by_layer(
slot_mappings, self.kv_cache_config
)
@@ -1169,6 +1177,7 @@ class GPUModelRunner(LoRAModelRunnerMixin):
prompt_logprobs_dict=prompt_logprobs_dict, # type: ignore[arg-type]
kv_connector_output=kv_connector_output,
)
self.routed_experts.save()
async_output = AsyncOutput(
model_runner_output=model_runner_output,
sampler_output=sampler_output,
+106
View File
@@ -0,0 +1,106 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from __future__ import annotations
from typing import Any
import numpy as np
import torch
from vllm.logger import init_logger
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
RoutedExpertsCapturer,
)
from vllm.v1.kv_cache_interface import AttentionSpec
logger = init_logger(__name__)
class RoutedExpertsCaptureHelper:
def __init__(self) -> None:
self._initialized = False
self._attn_gid = 0
self._slot_mapping: np.ndarray | None = None
@property
def initialized(self) -> bool:
return self._initialized
def init(self, runner: Any) -> None:
logger.info(
"Initializing routed experts capturer, enable_return_routed_experts: %s",
runner.model_config.enable_return_routed_experts,
)
capturer = RoutedExpertsCapturer.create()
self._attn_gid = self._get_attention_kv_cache_gid(runner)
min_block_size = min(
group.kv_cache_spec.block_size
for group in runner.kv_cache_config.kv_cache_groups
)
num_groups = len(runner.kv_cache_config.kv_cache_groups)
max_num_kv_tokens = (
runner.kv_cache_config.num_blocks // num_groups
) * min_block_size
dcp_size = runner.vllm_config.parallel_config.decode_context_parallel_size
pcp_size = runner.vllm_config.parallel_config.prefill_context_parallel_size
if pcp_size * dcp_size > 1:
max_num_kv_tokens *= pcp_size * dcp_size
capturer.init_buffer(
max_num_batched_tokens=runner.scheduler_config.max_num_batched_tokens,
max_num_kv_tokens=max_num_kv_tokens,
vllm_config=runner.vllm_config,
)
self.bind(runner, capturer)
self._initialized = True
def bind(self, runner: Any, capturer: RoutedExpertsCapturer) -> None:
from vllm.model_executor.layers.fused_moe.layer import FusedMoE
from vllm.model_executor.layers.fused_moe.router.base_router import (
BaseRouter,
)
for module in runner.compilation_config.static_forward_context.values():
if isinstance(module, FusedMoE) and isinstance(module.router, BaseRouter):
layer_id = module.layer_id
def _capture_fn(topk_ids, _layer_id=layer_id, _capturer=capturer):
_capturer.capture(_layer_id, topk_ids)
module.router.set_capture_fn(_capture_fn)
def before_execute(self) -> None:
if not self._initialized:
return
capturer = RoutedExpertsCapturer.get_instance()
if capturer is None:
logger.error("RoutedExpertsCapturer not initialized.")
return
capturer.clear_buffer()
def record_slot_mapping(
self,
slot_mappings: tuple[torch.Tensor, ...],
num_tokens: int,
) -> None:
if not self._initialized:
return
slot_mapping_attn = slot_mappings[self._attn_gid]
self._slot_mapping = slot_mapping_attn[:num_tokens].cpu().numpy()
def save(self) -> None:
if not self._initialized or self._slot_mapping is None:
return
capturer = RoutedExpertsCapturer.get_instance()
if capturer is None:
logger.error("RoutedExpertsCapturer not initialized.")
return
capturer.save_captured_experts(indices=self._slot_mapping)
@staticmethod
def _get_attention_kv_cache_gid(runner: Any) -> int:
for gid, group in enumerate(runner.kv_cache_config.kv_cache_groups):
if isinstance(group.kv_cache_spec, AttentionSpec):
return gid
return 0