[Core] Replace routing replay with device cache and async D2H pipeline (#39917)

Signed-off-by: Tomer Barnatan <[email protected]>
This commit is contained in:
TomerBN-Nvidia
2026-05-07 11:24:56 -07:00
committed by GitHub
parent 8eb401134e
commit 8189a15914
15 changed files with 1398 additions and 642 deletions
+289
View File
@@ -0,0 +1,289 @@
# Routed Experts Replay
## Overview
Routed experts replay captures which MoE (Mixture of Experts) experts process each token during inference and returns this information alongside the generated text. This is essential for **reinforcement learning (RL) training pipelines** (such as GRPO and RLHF) where the training step needs to reconstruct expert routing decisions from the inference pass.
When enabled, each API response includes:
- **`prompt_routed_experts`**: A `[prompt_len, num_moe_layers, top_k]` array of expert IDs for the prompt tokens (at the response level, shared across completions).
- **`routed_experts`**: A `[gen_len, num_moe_layers, top_k]` array of expert IDs for the generated tokens (per completion).
For example, a model with 40 MoE layers and top-22 routing that processes a 100-token prompt and generates 50 tokens would return:
- `prompt_routed_experts`: shape `[100, 40, 22]`
- `routed_experts`: shape `[50, 40, 22]`
Each value is an int16 expert ID in the range `[0, num_experts)`.
## Quickstart
### OpenAI API Server
```bash
vllm serve <MODEL> \
--enable-return-routed-experts \
--tensor-parallel-size 4 \
--enable-expert-parallel
```
Then query the `/v1/completions` endpoint as usual. The response includes routing data:
```python
import requests
resp = requests.post("http://localhost:8000/v1/completions", json={
"model": "<MODEL>",
"prompt": "Explain quantum computing.",
"max_tokens": 64,
"temperature": 0.0,
}).json()
# Generation routing (per completion choice)
gen_routing = resp["choices"][0]["routed_experts"] # [gen_len, layers, top_k]
# Prompt routing (shared across all choices)
prompt_routing = resp["prompt_routed_experts"] # [prompt_len, layers, top_k]
print(f"Prompt routing shape: [{len(prompt_routing)}, "
f"{len(prompt_routing[0])}, {len(prompt_routing[0][0])}]")
print(f"Gen routing shape: [{len(gen_routing)}, "
f"{len(gen_routing[0])}, {len(gen_routing[0][0])}]")
```
### Python SDK (Offline Inference)
```python
from vllm import LLM, SamplingParams
llm = LLM(
model="<MODEL>",
enable_return_routed_experts=True,
tensor_parallel_size=4,
enable_expert_parallel=True,
)
outputs = llm.generate(
["Explain quantum computing."],
SamplingParams(temperature=0, max_tokens=64),
)
result = outputs[0]
# Prompt routing: numpy array, shape [prompt_len, num_moe_layers, top_k]
prompt_routing = result.prompt_routed_experts
print(f"Prompt routing: {prompt_routing.shape}, dtype={prompt_routing.dtype}")
# Generation routing: numpy array, shape [gen_len, num_moe_layers, top_k]
gen_routing = result.outputs[0].routed_experts
print(f"Gen routing: {gen_routing.shape}, dtype={gen_routing.dtype}")
```
## Output Format
### `CompletionOutput.routed_experts`
- **Type**: `numpy.ndarray` (Python SDK) or `list[list[list[int]]]` (JSON API)
- **Shape**: `[gen_len, num_moe_layers, top_k]`
- **Dtype**: `int16`
- **Content**: Expert IDs for **generated tokens only**. `gen_len` matches the number of generated tokens (i.e., `usage.completion_tokens` or fewer).
### `RequestOutput.prompt_routed_experts`
- **Type**: `numpy.ndarray` (Python SDK) or `list[list[list[int]]]` (JSON API)
- **Shape**: `[prompt_len, num_moe_layers, top_k]`
- **Dtype**: `int16`
- **Content**: Expert IDs for **prompt tokens only**. `prompt_len` matches `usage.prompt_tokens`. This field lives on the request-level response (not per-choice), because prompt routing is shared across all completions when `n > 1`.
### Why Separate Prompt and Generation Routing?
When a request has multiple completions (`n > 1`), each completion shares the same prompt but produces different generated text. Storing prompt routing once on the `RequestOutput` (rather than duplicating it on every `CompletionOutput`) avoids redundant data. For RL training, the consumer typically needs:
1. The prompt routing (once) to reconstruct the forward pass for the shared prefix.
2. The per-completion generation routing to reconstruct each completion's forward pass.
## Architecture
### Data Flow
```text
Forward Pass Async D2H Pipeline Output
───────────── ────────────────── ──────
FusedMoE layer After forward pass: On request finish:
writes topk_ids ──────► D2H copy to pinned ──────► Extract from host cache
to device buffer staging buffer Split at prompt_len
(L, N, K) int16 (via CUDA stream) Trim gen to output len
Scatter to per-request Serialize to API response
host cache (numpy)
```
### Device Cache
A pre-allocated GPU buffer with layout `(L, N, K)` where:
- `L` = number of MoE layers
- `N` = `max_num_batched_tokens`
- `K` = `num_experts_per_tok` (top-k)
The `(L, N, K)` layout ensures that `buffer[layer_id]` gives a contiguous `(N, K)` view per layer. Each `FusedMoE` layer gets a persistent reference to its slice via `module._routing_replay_out = buffer[layer_id]`.
**Dtype**: `int16` — sufficient for expert IDs (max ~512 experts in practice) and half the memory of `int32`.
### Host Cache
Per-request numpy arrays for accumulating routing data across decode steps. Each request gets a lazily allocated `(seq_len, L, K)` int16 buffer that grows as the sequence lengthens. Buffers are freed when a request completes.
### Async D2H Pipeline
After each forward pass, the model runner issues a non-blocking device-to-host copy on a dedicated CUDA stream:
1. **Copy**: `pinned_staging[:, :total_tokens, :].copy_(device_buffer[:, :total_tokens, :])` on a separate stream, recorded with a CUDA event.
2. **Scatter** (deferred to next step): On the *next* forward pass, synchronize the event (effectively free — an entire forward pass has elapsed) and scatter the staging data into per-request host cache buffers using the token positions.
This design ensures the D2H copy overlaps with the next forward pass, minimizing GPU stall time.
### CUDA Graph Compatibility
CUDA graph compatibility requires two mechanisms:
1. **Persistent tensor attribute**: Each `FusedMoE` layer stores a reference to its buffer slice as `module._routing_replay_out`. Because `torch.compile` captures module attributes by reference, graph replay always writes to the live buffer — not a stale snapshot.
2. **Static marking**: Both the full `(L, N, K)` buffer and each per-layer `(N, K)` view are marked with `cudagraph_mark_tensor_static()`. This prevents CUDA graphs from snapshot/restore behavior that would zero the buffer on replay.
### Multi-Node Support
On multi-node tensor-parallel setups, all TP ranks allocate a device buffer (required for symmetric CUDA graph structure), but only TP rank 0 runs the D2H pipeline and host cache. Routing data flows from the model runner through `ModelRunnerOutput` via Ray DAG to the scheduler — no shared memory or file locks needed.
### Routing Capture Path
For the **non-monolithic (Triton) kernel path** (e.g., BF16 MoE), routing is captured after `select_experts()` in the MoE runner:
```python
routing_replay_out = getattr(layer, "_routing_replay_out", None)
topk_weights, topk_ids = self.router.select_experts(...)
if routing_replay_out is not None:
routing_replay_out[:topk_ids.shape[0]].copy_(topk_ids.to(torch.int16))
```
For the **monolithic kernel path** (e.g., FP8/MXFP8 via FlashInfer), `routing_replay_out` is threaded through the `apply_monolithic()` call chain and FlashInfer writes expert IDs directly during routing inside the fused kernel.
### MTP (Multi-Token Prediction) Handling
With MTP speculative decoding, the model captures routing for all tokens including speculative ones that may later be rejected. When a request finishes, the generation routing is trimmed to match the actual number of accepted output tokens:
```python
num_gen = self.detokenizer.num_output_tokens()
if gen_routed_experts.shape[0] > num_gen and num_gen > 0:
gen_routed_experts = gen_routed_experts[:num_gen]
```
This ensures the routing array length always matches the token IDs in the response.
## Design Decisions
### Why Replace SharedMemory with Device Cache?
The previous implementation used `multiprocessing.SharedMemory` with `fcntl` file locking to transfer routing data from GPU workers to the scheduler. This approach had fundamental problems:
- **Multi-node**: `SharedMemory` is node-local. On multi-node TP setups (required for 400B+ parameter models), the scheduler on node 0 cannot read shared memory from workers on other nodes.
- **Performance**: Synchronous `.cpu().numpy()` D2H transfers block the GPU. File-based locking adds further overhead.
- **CUDA graphs**: The callback-based capture mechanism bakes tensor references at trace time, causing stale data on graph replay.
The device cache approach solves all three: data flows through Ray DAG (works multi-node), D2H is async (non-blocking), and persistent tensor attributes work with CUDA graphs.
### Why `(L, N, K)` Layout Instead of `(N, L, K)`?
FlashInfer's `routing_replay_out` parameter expects a contiguous `(N, K)` tensor per layer. With `(L, N, K)` layout, `buffer[layer_id]` gives a contiguous `(N, K)` view with zero-copy slicing. The previous `(N, L, K)` layout would require non-contiguous indexing or an explicit copy.
### Why int16 Instead of int32?
Expert IDs are small integers (typically 0-255 for models with up to 256 experts). `int16` supports up to 32,767 experts — far more than any current model — while halving GPU memory usage and D2H bandwidth compared to `int32`.
### Why Split Prompt and Generation Routing?
RL training pipelines process prompt and generation routing separately:
- Prompt routing reconstructs the shared forward pass for the input.
- Generation routing reconstructs each sampled trajectory.
With `n > 1` completions, all completions share the same prompt routing. Duplicating it per completion would waste memory proportional to `n * prompt_len * L * K`. Instead, `prompt_routed_experts` is stored once on `RequestOutput` and shared.
### Why Async D2H Instead of Synchronous Copy?
A synchronous `.cpu()` call forces the GPU to drain its command queue before the copy can begin, stalling the pipeline. The async approach:
1. Issues the copy on a separate CUDA stream (non-blocking to the main compute stream).
2. Defers the host-side scatter to the *next* step, by which time the copy has finished.
This means the D2H transfer overlaps entirely with the next forward pass, adding near-zero latency to the critical path.
### Why All TP Ranks Get a Device Buffer?
CUDA graph capture records the exact sequence of kernel calls and their arguments. If only rank 0 had a device buffer, the `FusedMoE` layer would take a different code path on rank 0 vs. other ranks (one writes to a buffer, others don't). This asymmetry causes different CUDA graph structures across ranks, which can lead to NCCL deadlocks during collective operations inside the graph. Giving all ranks a real buffer ensures symmetric graph structure. Only rank 0 does the D2H copy and host cache management.
## Performance
Routing replay adds a small overhead from the device buffer writes and async D2H copies. On tested configurations:
- **Throughput overhead** (random data, ISL=1024, OSL=1024): **~2%**
- **Memory overhead** (int16 buffer, 40 layers, 8192 tokens, top-22): **~14 MB per GPU**
- **Accuracy impact** (GSM8K): **Zero** (pass@1 identical with and without routing replay)
The overhead is dominated by the per-layer `.copy_()` during the forward pass. The async D2H pipeline runs entirely in the background.
## Supported Configurations
| Configuration | Supported |
|------------------------------------------|-----------------------------------------------------------|
| BF16 Triton MoE (non-monolithic) | Yes |
| FP8/MXFP8 FlashInfer MoE (monolithic) | Yes (requires FlashInfer with `routing_replay_out`) |
| CUDA graphs | Yes |
| Multi-node tensor parallelism | Yes |
| Data parallelism (DP) | Yes |
| Expert parallelism (EP) | Yes |
| Prefix caching | Yes (cached positions marked with `-1` sentinel) |
| MTP speculative decoding | Yes (gen routing trimmed to accepted tokens) |
| `n > 1` (multiple completions) | Yes (prompt routing shared, gen routing per-completion) |
## Limitations
- **Streaming**: Routing data is only available when the request finishes (not streamed incrementally).
- **V1 engine only**: Routing replay is implemented for the vLLM V1 engine.
- **Preempted requests**: When a request is preempted by the scheduler (and later resumed via re-prefill), any routing already accumulated in the worker's host cache for that request is dropped without being emitted. The consumer sees `routed_experts=None` for the resumed request with no other signal. Partial-rollout and async-RL pipelines that rely on routing for preempted requests should either disable preemption (`--no-enable-chunked-prefill` / sufficient KV headroom) or reconstruct routing on the resumed prefill.
- **Async scheduling**: Not supported; rejected at config time. The worker-side stop predicate reads `req_state.output_token_ids[-1]`, which under async scheduling is the placeholder `-1` until `AsyncGPUModelRunnerOutput` resolves the real sampled token, so EOS / stop-token finishes would silently drop routing. Use sync scheduling (the default when `--enable-return-routed-experts` is set, or set explicitly with the appropriate scheduler config).
- **Sequence parallelism / naive DP MoE dispatch**: Not supported on the FusedMoE layer; rejected at bind time. SP shards `topk_ids` along dim 0 across the TP group so each rank only captures `1/sp_size` of the rows; naive DP dispatch all-gathers tokens across DP ranks before routing, so `topk_ids.shape[0]` exceeds the per-rank buffer size. Both raise `NotImplementedError` from `bind_routing_capture_to_model`.
- **Pipeline / prefill-context / decode-context parallelism**: Not yet validated; rejected at config time.
## CLI Reference
| Flag | Description |
|------------------------------------|------------------------------------------------------------------------|
| `--enable-return-routed-experts` | Enable routing replay capture and return expert IDs in API responses. |
## API Reference
### Completions (`/v1/completions`)
**Response-level field:**
| Field | Type | Description |
|---------------------------|-------------------------------------|-----------------------------------------------------------------------------|
| `prompt_routed_experts` | `list[list[list[int]]]` or `null` | Expert IDs for prompt tokens. Shape: `[prompt_len, num_moe_layers, top_k]`. |
**Choice-level field:**
| Field | Type | Description |
|--------------------|-------------------------------------|-------------------------------------------------------------------------------|
| `routed_experts` | `list[list[list[int]]]` or `null` | Expert IDs for generated tokens. Shape: `[gen_len, num_moe_layers, top_k]`. |
### Chat Completions (`/v1/chat/completions`)
Same fields as above on `ChatCompletionResponse` and `ChatCompletionResponseChoice`.
### Python SDK
| Object | Field | Type | Description |
|----------------------|---------------------------|--------------------------|-----------------------------|
| `RequestOutput` | `prompt_routed_experts` | `np.ndarray` or `None` | `[prompt_len, L, K]` i16 |
| `CompletionOutput` | `routed_experts` | `np.ndarray` or `None` | `[gen_len, L, K]` int16 |
@@ -1,245 +1,162 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
import types
from types import SimpleNamespace
from unittest.mock import patch
import pytest
import torch
from vllm.distributed.eplb.eplb_state import EplbLayerState
from vllm.model_executor.layers.fused_moe.config import RoutingMethodType
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
RoutedExpertsCapturer,
)
from vllm.model_executor.layers.fused_moe.router.base_router import BaseRouter
pytestmark = pytest.mark.cpu_test
_REC_MODULE = "vllm.model_executor.layers.fused_moe.routed_experts_capturer"
def test_bind_routing_capture_to_model_sets_layer_view(monkeypatch):
import vllm.model_executor.layers.fused_moe.layer as fused_moe_layer
import vllm.model_executor.layers.fused_moe.routed_experts_capturer as rec_mod
def _capturer_with_buffer(
*,
max_tokens: int = 8,
num_layers: int = 4,
num_experts_per_tok: int = 2,
dp_rank: int = 0,
) -> RoutedExpertsCapturer:
c = RoutedExpertsCapturer()
c.dp_rank = dp_rank
c._device_buffer = torch.full(
(max_tokens, num_layers, num_experts_per_tok),
-1,
dtype=torch.int32,
)
return c
class _DummyMoEConfig:
is_sequence_parallel = False
dp_size = 1
class DummyRouter(BaseRouter):
@property
def routing_method_type(self) -> RoutingMethodType:
return RoutingMethodType.FUSED_TOPK
def _compute_routing(
self, hidden_states, router_logits, indices_type, *, input_ids=None
):
topk_ids = torch.tensor([[1, 2], [3, 4]], dtype=torch.int64)
topk_weights = torch.ones_like(topk_ids, dtype=torch.float32)
return topk_weights, topk_ids
def _apply_eplb_mapping(self, topk_ids: torch.Tensor) -> torch.Tensor:
# Make mapping observable without requiring CUDA EPLB path.
return topk_ids + 10
def _make_router() -> DummyRouter:
return DummyRouter(
top_k=2,
global_num_experts=16,
eplb_state=EplbLayerState(),
enable_eplb=False,
indices_type_getter=None,
)
def test_base_router_capture_pre_eplb_mapping():
router = _make_router()
captured = []
def capture_fn(ids):
captured.append(ids.clone())
router.set_capture_fn(capture_fn)
topk_weights, topk_ids = router.select_experts(
hidden_states=torch.empty(1),
router_logits=torch.empty(1),
)
assert topk_weights.shape == topk_ids.shape
assert len(captured) == 1
assert torch.equal(captured[0], torch.tensor([[1, 2], [3, 4]]))
assert torch.equal(topk_ids, torch.tensor([[11, 12], [13, 14]]))
def test_base_router_capture_with_eplb_enabled():
router = _make_router()
router.enable_eplb = True
router.eplb_state.expert_load_view = torch.zeros(32, dtype=torch.int64)
router.eplb_state.logical_to_physical_map = torch.arange(32).view(32, 1)
router.eplb_state.logical_replica_count = torch.ones(32, dtype=torch.int64)
router.eplb_state.should_record_tensor = torch.ones((), dtype=torch.bool)
captured = []
def capture_fn(ids):
captured.append(ids.clone())
router.set_capture_fn(capture_fn)
_, topk_ids = router.select_experts(
hidden_states=torch.empty(1),
router_logits=torch.empty(1),
)
assert len(captured) == 1
# Capture should see logical ids pre-EPLB mapping.
assert torch.equal(captured[0], torch.tensor([[1, 2], [3, 4]]))
# Our DummyRouter mapping adds +10.
assert torch.equal(topk_ids, torch.tensor([[11, 12], [13, 14]]))
def test_gpu_model_runner_binds_router_capture(monkeypatch):
from vllm.v1.worker import gpu_model_runner as gmr
class _DummyQuantMethod:
supports_internal_mk = True
class DummyFusedMoE:
def __init__(self):
self.layer_id = 7
self.router = _make_router()
_routing_replay_out: torch.Tensor
class DummyCapturer:
def __init__(self):
self.calls = []
def capture(self, layer_id, topk_ids):
self.calls.append((layer_id, topk_ids))
dummy_module = DummyFusedMoE()
# Patch the runtime import inside _bind_routed_experts_capturer.
import vllm.model_executor.layers.fused_moe.layer as fused_moe_layer
def __init__(self, moe_layer_id):
self.moe_layer_id = moe_layer_id
self.moe_config = _DummyMoEConfig()
self.quant_method = _DummyQuantMethod()
monkeypatch.setattr(fused_moe_layer, "FusedMoE", DummyFusedMoE)
dummy_self = types.SimpleNamespace(
compilation_config=types.SimpleNamespace(
static_forward_context={"dummy": dummy_module}
)
)
num_layers, num_tokens, top_k = 4, 8, 2
buffer = torch.zeros((num_layers, num_tokens, top_k), dtype=torch.int16)
capturer = DummyCapturer()
gmr.GPUModelRunner._bind_routed_experts_capturer(dummy_self, capturer)
assert dummy_module.router.capture_fn is not None
dummy_module.router.capture_fn(torch.tensor([[5, 6]]))
assert len(capturer.calls) == 1
layer_id, topk_ids = capturer.calls[0]
assert layer_id == 7
assert torch.equal(topk_ids, torch.tensor([[5, 6]]))
def test_gpu_model_runner_binding_stage(monkeypatch):
from vllm.v1.worker import gpu_model_runner as gmr
class DummyFusedMoE:
def __init__(self):
self.layer_id = 11
self.router = _make_router()
class DummyDeviceCache:
def __init__(self, buf):
self.buffer = buf
class DummyCapturer:
def __init__(self):
self.calls = []
def get_device_cache(self):
return DummyDeviceCache(buffer)
def capture(self, layer_id, topk_ids):
self.calls.append((layer_id, topk_ids))
monkeypatch.setattr(rec_mod, "get_global_experts_capturer", lambda: DummyCapturer())
dummy_module = DummyFusedMoE()
m0 = DummyFusedMoE(moe_layer_id=0)
m2 = DummyFusedMoE(moe_layer_id=2)
import vllm.model_executor.layers.fused_moe.layer as fused_moe_layer
class DummyModel:
def modules(self):
return iter([m0, m2])
monkeypatch.setattr(fused_moe_layer, "FusedMoE", DummyFusedMoE)
rec_mod.bind_routing_capture_to_model(DummyModel())
dummy_self = types.SimpleNamespace(
compilation_config=types.SimpleNamespace(
static_forward_context={"dummy": dummy_module}
assert torch.equal(m0._routing_replay_out, buffer[0])
assert torch.equal(m2._routing_replay_out, buffer[2])
def test_bind_routing_capture_to_model_noop_when_disabled(monkeypatch):
import vllm.model_executor.layers.fused_moe.routed_experts_capturer as rec_mod
class DummyCapturer:
def get_device_cache(self):
return None
monkeypatch.setattr(rec_mod, "get_global_experts_capturer", lambda: DummyCapturer())
class DummyModel:
def modules(self):
return iter([])
rec_mod.bind_routing_capture_to_model(DummyModel())
# =========================================================================
# Tests for device-cache routing replay architecture
# =========================================================================
class TestRoutedExpertsDeviceCache:
"""Tests for _RoutedExpertsDeviceCache (GPU buffer for routing data)."""
def test_allocation_shape_and_dtype(self):
"""Device cache allocates (L, N, K) int16 buffer."""
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
_RoutedExpertsDeviceCache,
)
)
# Before binding, no capture hook.
assert dummy_module.router.capture_fn is None
cache = _RoutedExpertsDeviceCache(
num_hidden_layers=40,
max_num_batched_tokens=8192,
num_experts_per_tok=8,
device="cpu",
)
assert cache.buffer.shape == (40, 8192, 8)
assert cache.buffer.dtype == torch.int16
capturer = DummyCapturer()
gmr.GPUModelRunner._bind_routed_experts_capturer(dummy_self, capturer)
def test_per_layer_view_is_contiguous(self):
"""buffer[layer_id] gives contiguous (N, K) view for FlashInfer."""
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
_RoutedExpertsDeviceCache,
)
# After binding, hook should exist and be callable.
assert callable(dummy_module.router.capture_fn)
dummy_module.router.capture_fn(torch.tensor([[9, 10]]))
assert len(capturer.calls) == 1
cache = _RoutedExpertsDeviceCache(
num_hidden_layers=40,
max_num_batched_tokens=8192,
num_experts_per_tok=8,
device="cpu",
)
layer_view = cache.buffer[0]
assert layer_view.is_contiguous()
assert layer_view.shape == (8192, 8)
def test_routed_experts_capturer_single_dp_no_metadata():
"""dp_metadata is None: capture writes the full topk_ids rows."""
capturer = _capturer_with_buffer(dp_rank=0)
topk = torch.tensor([[1, 2], [3, 4], [5, 6]], dtype=torch.int32)
ctx = SimpleNamespace(dp_metadata=None)
with patch(f"{_REC_MODULE}.get_forward_context", return_value=ctx):
capturer.capture(layer_id=0, topk_ids=topk)
assert torch.equal(capturer._device_buffer[:3, 0, :], topk)
assert capturer._device_buffer[3, 0, 0].item() == -1
class TestRoutedExpertsHostCache:
"""Tests for _RoutedExpertsHostCache (per-request numpy buffer)."""
def test_sentinel_initialization(self):
"""Host cache initializes with zeros by default."""
import numpy as np
def test_routed_experts_capturer_dp_naive_concatenated_all_ranks():
"""n == sum(num_tokens_dp): slice this rank's segment from concatenated topk."""
capturer = _capturer_with_buffer(dp_rank=1)
num_tokens_dp = torch.tensor([2, 3], dtype=torch.int32)
ctx = SimpleNamespace(
dp_metadata=SimpleNamespace(num_tokens_across_dp_cpu=num_tokens_dp)
)
# Concatenated order: rank0 rows then rank1 rows.
topk = torch.tensor(
[[0, 1], [2, 3], [10, 11], [12, 13], [14, 15]], dtype=torch.int32
)
with patch(f"{_REC_MODULE}.get_forward_context", return_value=ctx):
capturer.capture(layer_id=0, topk_ids=topk)
want = topk[2:5]
assert torch.equal(capturer._device_buffer[:3, 0, :], want)
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
_RoutedExpertsHostCache,
)
cache = _RoutedExpertsHostCache(
num_hidden_layers=40,
num_experts_per_tok=8,
max_model_len=1024,
)
buf = cache.get_or_grow_buffer("req1", max_pos=100)
assert buf.dtype == np.int16
assert (buf == 0).all(), "Host cache must initialize with zeros"
def test_routed_experts_capturer_dp_modular_local_tokens():
"""n == token_num_per_dp: topk is already local to this DP rank."""
capturer = _capturer_with_buffer(dp_rank=1)
num_tokens_dp = torch.tensor([2, 3], dtype=torch.int32)
ctx = SimpleNamespace(
dp_metadata=SimpleNamespace(num_tokens_across_dp_cpu=num_tokens_dp)
)
topk = torch.tensor([[10, 11], [12, 13], [14, 15]], dtype=torch.int32)
with patch(f"{_REC_MODULE}.get_forward_context", return_value=ctx):
capturer.capture(layer_id=0, topk_ids=topk)
assert torch.equal(capturer._device_buffer[:3, 0, :], topk)
def test_grow_preserves_existing_data(self):
"""Growing the buffer preserves previously written data."""
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
_RoutedExpertsHostCache,
)
cache = _RoutedExpertsHostCache(
num_hidden_layers=40,
num_experts_per_tok=8,
max_model_len=1024,
)
buf = cache.get_or_grow_buffer("req1", max_pos=50)
buf[0, 0, 0] = 42
buf2 = cache.get_or_grow_buffer("req1", max_pos=200)
assert buf2[0, 0, 0] == 42, "Data lost during buffer grow"
def test_routed_experts_capturer_dp_unexpected_batch_raises():
"""Mismatch between topk batch dim and DP layout: fail fast."""
capturer = _capturer_with_buffer(dp_rank=0)
num_tokens_dp = torch.tensor([2, 3], dtype=torch.int32)
ctx = SimpleNamespace(
dp_metadata=SimpleNamespace(num_tokens_across_dp_cpu=num_tokens_dp)
)
# total=5, local=2: n=1 matches neither naive (5) nor modular (2).
topk = torch.tensor([[1, 2]], dtype=torch.int32)
with (
patch(f"{_REC_MODULE}.get_forward_context", return_value=ctx),
pytest.raises(AssertionError, match="unexpected topk_ids batch dim"),
):
capturer.capture(layer_id=0, topk_ids=topk)
assert capturer._device_buffer[0, 0, 0].item() == -1
def test_free_request_removes_buffer(self):
"""Freeing a request removes its buffer."""
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
_RoutedExpertsHostCache,
)
cache = _RoutedExpertsHostCache(
num_hidden_layers=40,
num_experts_per_tok=8,
max_model_len=1024,
)
cache.get_or_grow_buffer("req1", max_pos=50)
cache.free_request("req1")
assert cache.get_buffer("req1") is None
+51
View File
@@ -1157,6 +1157,12 @@ class VllmConfig:
if envs.VLLM_USE_V2_MODEL_RUNNER:
self._validate_v2_model_runner()
if (
self.model_config is not None
and self.model_config.enable_return_routed_experts
):
self._validate_return_routed_experts()
# Re-compute compile ranges after platform-specific config updates
# (e.g., XPU may lower max_num_batched_tokens when MLA is enabled)
self._set_compile_ranges()
@@ -1846,6 +1852,51 @@ class VllmConfig:
+ ", ".join(unsupported)
)
def _validate_return_routed_experts(self) -> None:
"""Reject parallelism configurations not yet validated with
--enable-return-routed-experts.
Validated scope (PR #39917): TP, EP, DP, single-node and multi-node,
prefix caching, and speculative decoding (MTP validated end-to-end;
Eagle/Eagle3/Ngram/Medusa supported by construction since the
routing buffer is bound only to the target model and verified-token
routing lands at the correct positions during the main forward).
Out-of-scope (block until validated): PP > 1, prefill context
parallelism (PCP) > 1, decode context parallelism (DCP) > 1,
async scheduling.
"""
unsupported: list[str] = []
if self.parallel_config.pipeline_parallel_size > 1:
unsupported.append(
"pipeline parallelism "
f"(pipeline_parallel_size="
f"{self.parallel_config.pipeline_parallel_size})"
)
if self.parallel_config.prefill_context_parallel_size > 1:
unsupported.append(
"prefill context parallelism "
f"(prefill_context_parallel_size="
f"{self.parallel_config.prefill_context_parallel_size})"
)
if self.parallel_config.decode_context_parallel_size > 1:
unsupported.append(
"decode context parallelism "
f"(decode_context_parallel_size="
f"{self.parallel_config.decode_context_parallel_size})"
)
if self.scheduler_config.async_scheduling:
unsupported.append("async scheduling")
if unsupported:
raise ValueError(
"--enable-return-routed-experts is not yet validated with: "
+ ", ".join(unsupported)
+ ". Disable these features or omit "
"--enable-return-routed-experts."
)
def validate_block_size(self) -> None:
"""Validate block_size against DCP and mamba constraints.
@@ -92,12 +92,16 @@ class ChatCompletionResponseChoice(OpenAIBaseModel):
# not part of the OpenAI spec but is useful for tracing the tokens
# in agent scenarios
token_ids: list[int] | None = None
routed_experts: list[list[list[int]]] | None = None # [gen_len, num_layers, top_k]
class ChatCompletionResponse(OpenAIBaseModel):
id: str = Field(default_factory=lambda: f"chatcmpl-{random_uuid()}")
object: Literal["chat.completion"] = "chat.completion"
created: int = Field(default_factory=lambda: int(time.time()))
prompt_routed_experts: list[list[list[int]]] | None = (
None # [prompt_len, num_layers, top_k]
)
model: str
choices: list[ChatCompletionResponseChoice]
service_tier: Literal["auto", "default", "flex", "scale", "priority"] | None = None
@@ -1088,6 +1088,11 @@ class OpenAIServingChat(OpenAIServing):
token_ids=(
as_list(output.token_ids) if request.return_token_ids else None
),
routed_experts=(
output.routed_experts.tolist()
if output.routed_experts is not None
else None
),
)
choices.append(choice_data)
continue
@@ -1309,6 +1314,11 @@ class OpenAIServingChat(OpenAIServing):
token_ids=(
as_list(output.token_ids) if request.return_token_ids else None
),
routed_experts=(
output.routed_experts.tolist()
if output.routed_experts is not None
else None
),
)
choice_data = maybe_filter_parallel_tool_calls(choice_data, request)
@@ -1348,6 +1358,10 @@ class OpenAIServingChat(OpenAIServing):
request_metadata.final_usage_info = usage
prompt_routed_experts = None
if final_res.prompt_routed_experts is not None:
prompt_routed_experts = final_res.prompt_routed_experts.tolist()
response = ChatCompletionResponse(
id=request_id,
created=created_time,
@@ -1360,6 +1374,7 @@ class OpenAIServingChat(OpenAIServing):
final_res.prompt_token_ids if request.return_token_ids else None
),
kv_transfer_params=final_res.kv_transfer_params,
prompt_routed_experts=prompt_routed_experts,
)
# Log complete response if output logging is enabled
@@ -468,12 +468,16 @@ class CompletionResponseChoice(OpenAIBaseModel):
token_ids: list[int] | None = None # For response
prompt_logprobs: list[dict[int, Logprob] | None] | None = None
prompt_token_ids: list[int] | None = None # For prompt
routed_experts: list[list[list[int]]] | None = None # [gen_len, num_layers, top_k]
class CompletionResponse(OpenAIBaseModel):
id: str = Field(default_factory=lambda: f"cmpl-{random_uuid()}")
object: Literal["text_completion"] = "text_completion"
created: int = Field(default_factory=lambda: int(time.time()))
prompt_routed_experts: list[list[list[int]]] | None = (
None # [prompt_len, num_layers, top_k]
)
model: str
choices: list[CompletionResponseChoice]
service_tier: Literal["auto", "default", "flex", "scale", "priority"] | None = None
@@ -542,6 +542,11 @@ class OpenAIServingCompletion(OpenAIServing):
token_ids=(
as_list(output.token_ids) if request.return_token_ids else None
),
routed_experts=(
output.routed_experts.tolist()
if output.routed_experts is not None
else None
),
)
choices.append(choice_data)
@@ -565,8 +570,13 @@ class OpenAIServingCompletion(OpenAIServing):
)
request_metadata.final_usage_info = usage
prompt_routed_experts = None
if final_res_batch:
kv_transfer_params = final_res_batch[0].kv_transfer_params
pre = final_res_batch[0].prompt_routed_experts
if pre is not None:
prompt_routed_experts = pre.tolist()
return CompletionResponse(
id=request_id,
created=created_time,
@@ -575,6 +585,7 @@ class OpenAIServingCompletion(OpenAIServing):
usage=usage,
system_fingerprint=self.system_fingerprint,
kv_transfer_params=kv_transfer_params,
prompt_routed_experts=prompt_routed_experts,
)
def _create_completion_logprobs(
@@ -246,6 +246,9 @@ class FusedMoE(PluggableLayer):
not supported by the router (or the experts).
"""
# Auto-incrementing layer ID for routing replay buffer binding.
_next_moe_layer_id: int = 0
# --8<-- [end:fused_moe]
def __init__(
@@ -290,6 +293,10 @@ class FusedMoE(PluggableLayer):
):
super().__init__()
# Assign unique layer ID for routing replay buffer binding.
self.moe_layer_id = FusedMoE._next_moe_layer_id
FusedMoE._next_moe_layer_id += 1
if params_dtype is None:
params_dtype = torch.get_default_dtype()
self.params_dtype = params_dtype
File diff suppressed because it is too large Load Diff
@@ -451,6 +451,10 @@ class MoERunner(MoERunnerInterface):
shared_experts_input, SharedExpertsOrder.NO_OVERLAP
)
# Get routing replay buffer from persistent layer attribute
# (set by bind_routing_capture_to_model during capturer init)
routing_replay_out = getattr(layer, "_routing_replay_out", None)
if self._quant_method.is_monolithic:
fused_out = self._quant_method.apply_monolithic(
layer=layer,
@@ -465,6 +469,10 @@ class MoERunner(MoERunnerInterface):
input_ids=input_ids,
)
# Write routing data for non-monolithic path (Triton, etc.)
if routing_replay_out is not None:
routing_replay_out[: topk_ids.shape[0]].copy_(topk_ids.to(torch.int16))
# Passing shared_experts_input in case SharedExpertsOrder is
# MK_INTERNAL_OVERLAPPED.
fused_out = self._quant_method.apply(
+4
View File
@@ -121,6 +121,7 @@ class RequestOutput:
num_cached_tokens: int | None = None,
*,
kv_transfer_params: dict[str, Any] | None = None,
prompt_routed_experts: np.ndarray | None = None,
# Forward compatibility, code that uses args added in new release can
# still run with older versions of vLLM without breaking.
**kwargs: Any,
@@ -141,12 +142,15 @@ class RequestOutput:
self.encoder_prompt_token_ids = encoder_prompt_token_ids
self.num_cached_tokens = num_cached_tokens
self.kv_transfer_params = kv_transfer_params
self.prompt_routed_experts = prompt_routed_experts
def add(self, next_output: "RequestOutput", aggregate: bool) -> None:
"""Merge subsequent RequestOutput into this one"""
self.finished |= next_output.finished
self.kv_transfer_params = next_output.kv_transfer_params
if next_output.prompt_routed_experts is not None:
self.prompt_routed_experts = next_output.prompt_routed_experts
for next_completion in next_output.outputs:
for i, completion in enumerate(self.outputs):
+7 -70
View File
@@ -7,8 +7,6 @@ from collections.abc import Iterable
from dataclasses import replace
from typing import Any
import numpy as np
from vllm import envs
from vllm.compilation.cuda_graph import CUDAGraphStat
from vllm.config import VllmConfig
@@ -27,9 +25,6 @@ from vllm.distributed.kv_transfer.kv_connector.v1 import (
from vllm.distributed.kv_transfer.kv_connector.v1.base import KVConnectorMetadata
from vllm.distributed.kv_transfer.kv_connector.v1.metrics import KVConnectorStats
from vllm.logger import init_logger
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
RoutedExpertsReader,
)
from vllm.multimodal import MULTIMODAL_REGISTRY, MultiModalRegistry
from vllm.multimodal.encoder_budget import MultiModalBudget
from vllm.v1.core.encoder_cache_manager import (
@@ -52,7 +47,7 @@ from vllm.v1.core.sched.request_queue import (
)
from vllm.v1.core.sched.utils import check_stop, remove_all
from vllm.v1.engine import EngineCoreEventType, EngineCoreOutput, EngineCoreOutputs
from vllm.v1.kv_cache_interface import AttentionSpec, KVCacheConfig
from vllm.v1.kv_cache_interface import KVCacheConfig
from vllm.v1.metrics.perf import ModelMetrics, PerfStats
from vllm.v1.metrics.stats import PrefixCacheStats, SchedulerStats
from vllm.v1.outputs import DraftTokenIds, KVConnectorOutput, ModelRunnerOutput
@@ -260,43 +255,6 @@ class Scheduler(SchedulerInterface):
if self.log_stats and vllm_config.observability_config.enable_mfu_metrics:
self.perf_metrics = ModelMetrics(vllm_config)
if self.vllm_config.model_config.enable_return_routed_experts:
assert self.dcp_world_size == 1 and self.pcp_world_size == 1, (
"enable_return_routed_experts does not support context parallelism "
"(dcp_world_size > 1 or pcp_world_size > 1)"
)
self.routed_experts_reader = RoutedExpertsReader.create()
assert len(kv_cache_config.kv_cache_groups) > 0, (
"enable_return_routed_experts requires at least one kv cache group"
)
# Find the attention group for routed experts indexing.
self.routed_experts_attn_gid = 0
for gid, group in enumerate(kv_cache_config.kv_cache_groups):
if isinstance(group.kv_cache_spec, AttentionSpec):
self.routed_experts_attn_gid = gid
break
min_block_size = min(
[
group.kv_cache_spec.block_size
for group in kv_cache_config.kv_cache_groups
]
)
num_groups = len(kv_cache_config.kv_cache_groups)
self.max_num_kv_tokens = (
kv_cache_config.num_blocks // num_groups
) * min_block_size
dcp_size = self.vllm_config.parallel_config.decode_context_parallel_size
pcp_size = self.vllm_config.parallel_config.prefill_context_parallel_size
if pcp_size * dcp_size > 1:
self.max_num_kv_tokens *= pcp_size * dcp_size
self.routed_experts_reader.attach_buffer(
max_num_kv_tokens=self.max_num_kv_tokens,
vllm_config=self.vllm_config,
)
self._pause_state: PauseState = PauseState.UNPAUSED
def _mamba_block_aligned_split(
@@ -1415,11 +1373,15 @@ class Scheduler(SchedulerInterface):
request.resumable = False
stopped = True
# Get routing data from ModelRunnerOutput (via worker D2H pipeline)
routed_experts = None
if (
model_runner_output.routed_experts_dict is not None
and req_id in model_runner_output.routed_experts_dict
):
routed_experts = model_runner_output.routed_experts_dict[req_id]
finish_reason = None
if stopped:
routed_experts = self._get_routed_experts(request)
# Capture finish_reason BEFORE _handle_stopped_request, which may
# reset the status to WAITING for streaming requests that continue.
finish_reason = request.get_finished_reason()
@@ -1594,31 +1556,6 @@ class Scheduler(SchedulerInterface):
self._enqueue_waiting_request(request)
return False
def _get_routed_experts(self, request: Request) -> np.ndarray | None:
if not self.vllm_config.model_config.enable_return_routed_experts:
return None
kv_blocks = self.kv_cache_manager.get_blocks(request.request_id)
block_ids = kv_blocks.get_block_ids()[self.routed_experts_attn_gid]
num_tokens = request.num_tokens - 1
# compute slot mapping using attention group's block_size
block_ids_array = np.array(block_ids, dtype=np.int32)
num_blocks = len(block_ids)
attn_group = self.kv_cache_config.kv_cache_groups[self.routed_experts_attn_gid]
block_size = attn_group.kv_cache_spec.block_size
# generate block offsets
block_offsets = np.arange(0, block_size)
# compute slot mapping: slot = block_id * block_size + offset
slot_mapping = (
block_offsets.reshape((1, block_size))
+ block_ids_array.reshape((num_blocks, 1)) * block_size
).flatten()[:num_tokens]
return self.routed_experts_reader.get_routed_experts(indices=slot_mapping)
def _update_request_with_output(
self, request: Request, new_token_ids: list[int]
) -> tuple[list[int], bool]:
+27 -2
View File
@@ -11,6 +11,9 @@ import numpy as np
import torch
from vllm.lora.request import LoRARequest
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
split_routed_experts,
)
from vllm.outputs import (
STREAM_FINISHED,
CompletionOutput,
@@ -314,8 +317,24 @@ class RequestState:
finished,
)
# Split routing data into prompt and generation portions.
# Prompt routing lives on RequestOutput (shared across n>1
# completions); generation routing lives on each CompletionOutput.
prompt_routed_experts = None
gen_routed_experts = None
if routed_experts is not None:
prompt_len = len(self.prompt_token_ids) if self.prompt_token_ids else 0
num_gen = (
self.detokenizer.num_output_tokens()
if self.detokenizer is not None
else None
)
prompt_routed_experts, gen_routed_experts = split_routed_experts(
routed_experts, prompt_len, num_gen
)
output = self._new_completion_output(
new_token_ids, finish_reason, stop_reason, routed_experts
new_token_ids, finish_reason, stop_reason, gen_routed_experts
)
if self.parent_req is None:
@@ -327,7 +346,11 @@ class RequestState:
external_req_id = self.parent_req.external_req_id
return self._new_request_output(
external_req_id, outputs, finished, kv_transfer_params
external_req_id,
outputs,
finished,
kv_transfer_params,
prompt_routed_experts,
)
def _new_request_output(
@@ -336,6 +359,7 @@ class RequestState:
outputs: list[CompletionOutput] | list[PoolingOutput],
finished: bool,
kv_transfer_params: dict[str, Any] | None = None,
prompt_routed_experts: np.ndarray | None = None,
) -> RequestOutput | PoolingRequestOutput:
# If prompt embeds were used, put placeholder prompt token ids
prompt_token_ids = self.prompt_token_ids
@@ -371,6 +395,7 @@ class RequestState:
kv_transfer_params=kv_transfer_params,
num_cached_tokens=self.num_cached_tokens,
metrics=self.stats,
prompt_routed_experts=prompt_routed_experts,
)
def _new_completion_output(
+3
View File
@@ -198,6 +198,9 @@ class ModelRunnerOutput:
# req_id -> num_nans_in_logits
num_nans_in_logits: dict[str, int] | None = None
# req_id -> routed experts ndarray of shape (seq_len, num_moe_layers, top_k)
routed_experts_dict: dict[str, np.ndarray] | None = None
# information related to cudagraph execution
cudagraph_stats: CUDAGraphStat | None = None
+64 -45
View File
@@ -54,7 +54,11 @@ from vllm.lora.layers import LoRAMapping, LoRAMappingType
from vllm.model_executor.layers.attention import Attention, MLAAttention
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
RoutedExpertsCapturer,
extract_routed_experts_for_current_batch,
free_routing_buffers,
get_global_experts_capturer,
init_routed_experts_capturer_with_shared_cache,
issue_routing_d2h_copy,
)
from vllm.model_executor.layers.mamba.ops.ssu_dispatch import (
initialize_mamba_ssu_backend,
@@ -1092,6 +1096,12 @@ class GPUModelRunner(
for req_id in scheduler_output.finished_req_ids:
self.input_batch.remove_request(req_id)
if self.routed_experts_initialized:
free_routing_buffers(
scheduler_output.finished_req_ids,
scheduler_output.preempted_req_ids,
)
# Zero GPU memory for freshly allocated cache blocks to prevent
# stale NaN/data from corrupting attention or SSM computation.
if scheduler_output.new_block_ids_to_zero:
@@ -2175,10 +2185,6 @@ class GPUModelRunner(
block_table_gid_0 = _get_block_table(0)
slot_mapping_gid_0 = slot_mappings[0]
if self.routed_experts_initialized:
attn_gid = self.routed_experts_attn_gid
slot_mapping_attn = slot_mappings[attn_gid]
self.slot_mapping = slot_mapping_attn[:num_tokens].cpu().numpy()
num_computed_tokens_cpu = self.input_batch.num_computed_tokens_cpu_tensor[
:num_reqs_padded
]
@@ -3859,11 +3865,9 @@ class GPUModelRunner(
)
if self.routed_experts_initialized:
capturer = RoutedExpertsCapturer.get_instance()
capturer = get_global_experts_capturer()
if capturer is not None:
capturer.clear_buffer() # noqa
else:
logger.error("RoutedExpertsCapturer not initialized.")
capturer.finalize_pending_copy()
# If ngram_gpu is used, we need to copy the scheduler_output to avoid
# the modification has influence on the scheduler_output in engine core process.
@@ -4373,6 +4377,14 @@ class GPUModelRunner(
scheduler_output.total_num_scheduled_tokens,
)
if self.routed_experts_initialized:
issue_routing_d2h_copy(
input_batch_req_ids=self.input_batch.req_ids,
num_scheduled_tokens=scheduler_output.num_scheduled_tokens,
positions=self.positions,
positions_cpu=self._positions_cpu,
)
if propose_drafts_after_bookkeeping:
# ngram and other speculative decoding methods use the sampled
# tokens on the CPU, so they are run after bookkeeping.
@@ -4392,12 +4404,15 @@ class GPUModelRunner(
self.kv_connector_output = None
with record_function_or_nullcontext("gpu_model_runner: ModelRunnerOutput"):
routed_experts_dict = None
if self.routed_experts_initialized:
capturer = RoutedExpertsCapturer.get_instance()
if capturer is not None:
capturer.save_captured_experts(indices=self.slot_mapping) # noqa
else:
logger.error("RoutedExpertsCapturer not initialized.")
routed_experts_dict = extract_routed_experts_for_current_batch(
req_ids=req_ids_output_copy,
requests=self.requests,
req_id_to_index=self.input_batch.req_id_to_index,
num_tokens_no_spec=self.input_batch.num_tokens_no_spec,
max_model_len=self.max_model_len,
)
output = ModelRunnerOutput(
req_ids=req_ids_output_copy,
@@ -4411,6 +4426,7 @@ class GPUModelRunner(
else None,
num_nans_in_logits=num_nans_in_logits,
cudagraph_stats=cudagraph_stats,
routed_experts_dict=routed_experts_dict,
)
if not self.use_async_scheduling:
@@ -6138,6 +6154,7 @@ class GPUModelRunner(
"Skipping CUDA graph capture. To turn on CUDA graph capture, "
"ensure `cudagraph_mode` was not manually set to `NONE`"
)
self.init_routed_experts_capturer()
return 0
# Initialize encoder CUDA graph manager if enabled.
@@ -6171,6 +6188,13 @@ class GPUModelRunner(
start_time = time.perf_counter()
# Initialize the routed experts capturer once before any CUDA graph
# capture. Must happen before graphs are captured so the buffer
# address is baked into the graph. Do NOT call this inside
# _capture_cudagraphs() -- creating the capturer twice replaces
# the device buffer, causing graphs to write to a dead buffer.
self.init_routed_experts_capturer()
# Trigger CUDA graph capture for specific shapes.
# Capture the large shapes first so that the smaller shapes
# can reuse the memory pool allocated for the large shapes.
@@ -6965,45 +6989,40 @@ class GPUModelRunner(
"Initializing routed experts capturer, enable_return_routed_experts: %s",
self.model_config.enable_return_routed_experts,
)
routed_experts_capturer = RoutedExpertsCapturer.create()
self.routed_experts_attn_gid = self._get_attention_kv_cache_gid()
min_block_size = min(
[
group.kv_cache_spec.block_size
for group in self.kv_cache_config.kv_cache_groups
]
)
num_groups = len(self.kv_cache_config.kv_cache_groups)
self.max_num_kv_tokens = (
self.kv_cache_config.num_blocks // num_groups
) * min_block_size
dcp_size = self.vllm_config.parallel_config.decode_context_parallel_size
pcp_size = self.vllm_config.parallel_config.prefill_context_parallel_size
if pcp_size * dcp_size > 1:
self.max_num_kv_tokens *= pcp_size * dcp_size
from vllm.distributed import get_tp_group
routed_experts_capturer.init_buffer(
if hasattr(self.model_config.hf_text_config, "n_shared_experts"):
num_fused_shared_experts = 1
else:
num_fused_shared_experts = 0
tp_group = get_tp_group()
init_routed_experts_capturer_with_shared_cache(
enable=self.model_config.enable_return_routed_experts,
model_config=self.model_config,
num_fused_shared_experts=num_fused_shared_experts,
max_num_batched_tokens=self.scheduler_config.max_num_batched_tokens,
max_num_kv_tokens=self.max_num_kv_tokens,
vllm_config=self.vllm_config,
max_model_len=self.max_model_len,
device=self.device,
rank=tp_group.rank_in_group,
world_size=tp_group.world_size,
)
self._bind_routed_experts_capturer(routed_experts_capturer)
self._bind_routed_experts_capturer()
self.routed_experts_initialized = True
def _bind_routed_experts_capturer(self, 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,
# Pinned CPU buffer for async positions D2H (avoids sync .cpu() call)
self._positions_cpu = torch.empty(
self.scheduler_config.max_num_batched_tokens,
dtype=torch.long,
pin_memory=True,
)
for module in self.compilation_config.static_forward_context.values():
if isinstance(module, FusedMoE) and isinstance(module.router, BaseRouter):
layer_id = module.layer_id
def _bind_routed_experts_capturer(self) -> None:
from vllm.model_executor.layers.fused_moe.routed_experts_capturer import (
bind_routing_capture_to_model,
)
def _capture_fn(topk_ids, _layer_id=layer_id, _capturer=capturer):
_capturer.capture(_layer_id, topk_ids)
module.router.set_capture_fn(_capture_fn)
bind_routing_capture_to_model(self.model)
def may_add_encoder_only_layers_to_kv_cache_config(self) -> None:
"""