forked from Karylab-cklius/vllm
[EPLB] Mask padding in EPLB load recording (#38128)
Signed-off-by: ilmarkov <markovilya197@gmail.com> Signed-off-by: Markov Ilya <markovilya19@gmail.com> Co-authored-by: Markov Ilya <markovilya19@gmail.com>
This commit is contained in:
co-authored by
Markov Ilya
parent
311ad689ad
commit
a2abce646f
@@ -225,6 +225,12 @@ def _test_eplb_fml(env, world_size: int, test_config: TestConfig):
|
||||
logical_to_physical_map,
|
||||
logical_replica_count,
|
||||
)
|
||||
fml.router.eplb_state.should_record_tensor = torch.ones(
|
||||
(), dtype=torch.bool, device=device
|
||||
)
|
||||
fml.router.eplb_state.num_unpadded_tokens_tensors = [
|
||||
torch.tensor(0, dtype=torch.int32, device=device)
|
||||
]
|
||||
|
||||
out_after_shuffle = []
|
||||
with set_forward_context(
|
||||
|
||||
@@ -1332,6 +1332,9 @@ def _test_body_eplb(
|
||||
eplb_moe_layer.router.eplb_state.should_record_tensor = torch.ones(
|
||||
(), dtype=torch.bool, device=device
|
||||
)
|
||||
eplb_moe_layer.router.eplb_state.num_unpadded_tokens_tensors = [
|
||||
torch.tensor(0, dtype=torch.int32, device=device)
|
||||
]
|
||||
|
||||
# Get "after" output with rearranged weights and EPLB routing
|
||||
with set_forward_context(
|
||||
|
||||
@@ -61,12 +61,14 @@ def setup_eplb_state(
|
||||
global_num_experts, dtype=torch.int64, device="cuda"
|
||||
)
|
||||
should_record_tensor = torch.ones((), dtype=torch.bool, device="cuda")
|
||||
num_unpadded_tokens_tensors = [torch.tensor(0, dtype=torch.int32, device="cuda")]
|
||||
|
||||
return EplbLayerState(
|
||||
expert_load_view=expert_load_view,
|
||||
logical_to_physical_map=logical_to_physical_map,
|
||||
logical_replica_count=logical_replica_count,
|
||||
should_record_tensor=should_record_tensor,
|
||||
num_unpadded_tokens_tensors=num_unpadded_tokens_tensors,
|
||||
)
|
||||
|
||||
|
||||
@@ -782,3 +784,67 @@ def test_eplb_map_with_redundancy(
|
||||
torch.testing.assert_close(load, exp_load)
|
||||
else:
|
||||
assert load.sum().item() == 0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"l2p_map, replica_count, num_physical, topk_ids, "
|
||||
"num_unpadded, expected_out, expected_load",
|
||||
[
|
||||
pytest.param(
|
||||
[[0], [1], [2], [3]],
|
||||
[1, 1, 1, 1],
|
||||
4,
|
||||
[[0, 1], [2, 3], [0, 2], [1, 3]],
|
||||
2,
|
||||
[[0, 1], [2, 3], [0, 2], [1, 3]],
|
||||
# only rows 0,1 counted: expert 0→1, 1→1, 2→1, 3→1
|
||||
[1, 1, 1, 1],
|
||||
id="half_padded",
|
||||
),
|
||||
pytest.param(
|
||||
# record everything (None = no padding info)
|
||||
[[0], [1], [2], [3]],
|
||||
[1, 1, 1, 1],
|
||||
4,
|
||||
[[0, 1], [2, 3], [0, 2], [1, 3]],
|
||||
None,
|
||||
[[0, 1], [2, 3], [0, 2], [1, 3]],
|
||||
[2, 2, 2, 2],
|
||||
id="no_padding_info",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_eplb_map_num_unpadded_tokens(
|
||||
l2p_map,
|
||||
replica_count,
|
||||
num_physical,
|
||||
topk_ids,
|
||||
num_unpadded,
|
||||
expected_out,
|
||||
expected_load,
|
||||
):
|
||||
l2p = torch.tensor(l2p_map, dtype=torch.int64, device="cuda")
|
||||
rc = torch.tensor(replica_count, dtype=torch.int64, device="cuda")
|
||||
load = torch.zeros(num_physical, dtype=torch.int32, device="cuda")
|
||||
rec = torch.tensor(True, dtype=torch.bool, device="cuda")
|
||||
ids = torch.tensor(topk_ids, dtype=torch.int32, device="cuda")
|
||||
num_unpadded_t = (
|
||||
torch.tensor(num_unpadded, dtype=torch.int32, device="cuda")
|
||||
if num_unpadded is not None
|
||||
else None
|
||||
)
|
||||
|
||||
out = eplb_map_to_physical_and_record(
|
||||
topk_ids=ids,
|
||||
expert_load_view=load,
|
||||
logical_to_physical_map=l2p,
|
||||
logical_replica_count=rc,
|
||||
record_enabled=rec,
|
||||
num_unpadded_tokens=num_unpadded_t,
|
||||
)
|
||||
|
||||
exp_out = torch.tensor(expected_out, dtype=out.dtype, device="cuda")
|
||||
torch.testing.assert_close(out, exp_out)
|
||||
|
||||
exp_load = torch.tensor(expected_load, dtype=torch.int32, device="cuda")
|
||||
torch.testing.assert_close(load, exp_load)
|
||||
|
||||
@@ -91,6 +91,7 @@ def test_base_router_capture_with_eplb_enabled():
|
||||
eplb_state.logical_to_physical_map = torch.arange(32).view(32, 1)
|
||||
eplb_state.logical_replica_count = torch.ones(32, dtype=torch.int64)
|
||||
eplb_state.should_record_tensor = torch.ones((), dtype=torch.bool)
|
||||
eplb_state.num_unpadded_tokens_tensors = [torch.tensor(0, dtype=torch.int32)]
|
||||
router = _make_router(eplb_state=eplb_state)
|
||||
|
||||
captured = []
|
||||
|
||||
@@ -203,6 +203,26 @@ class SupportsHash(Protocol):
|
||||
def compute_hash(self) -> str: ...
|
||||
|
||||
|
||||
_config_hash_cache: dict[int, str] = {}
|
||||
|
||||
|
||||
def compute_hash_cached(config: SupportsHash) -> str:
|
||||
"""Cache config.compute_hash() by object identity.
|
||||
|
||||
Config objects (ModelConfig, etc.) are long-lived singletons that never
|
||||
mutate after construction, but compute_hash() is expensive (JSON
|
||||
serialization + SHA-256). This utility avoids recomputing the hash on
|
||||
every forward pass while keeping a single consistent key type for all
|
||||
lookup paths.
|
||||
"""
|
||||
key = id(config)
|
||||
result = _config_hash_cache.get(key)
|
||||
if result is None:
|
||||
result = config.compute_hash()
|
||||
_config_hash_cache[key] = result
|
||||
return result
|
||||
|
||||
|
||||
class SupportsMetricsInfo(Protocol):
|
||||
def metrics_info(self) -> dict[str, str]: ...
|
||||
|
||||
|
||||
@@ -458,7 +458,9 @@ class ElasticEPScalingExecutor:
|
||||
eplb_model_state.logical_to_physical_map,
|
||||
eplb_model_state.logical_replica_count,
|
||||
)
|
||||
eplb_state._init_should_record_tensor(model)
|
||||
eplb_state._propagate_shared_tensors(
|
||||
model, eplb_model_state.num_unpadded_tokens_tensors
|
||||
)
|
||||
model.update_physical_experts_metadata(
|
||||
num_physical_experts=num_physical_experts,
|
||||
num_local_physical_experts=num_local_experts,
|
||||
|
||||
@@ -35,6 +35,7 @@ import torch
|
||||
from torch.distributed import ProcessGroup, all_reduce
|
||||
|
||||
from vllm.config import ModelConfig, ParallelConfig
|
||||
from vllm.config.utils import compute_hash_cached
|
||||
from vllm.distributed.parallel_state import (
|
||||
get_ep_group,
|
||||
get_eplb_group,
|
||||
@@ -206,6 +207,13 @@ class EplbModelState:
|
||||
pending_result relies on the GIL to synchronize access between the main thread and
|
||||
the async worker.
|
||||
"""
|
||||
num_unpadded_tokens_tensors: list[torch.Tensor] | None = None
|
||||
"""
|
||||
Per-ubatch scalar int32 tensors holding the number of real (non-padding)
|
||||
tokens. Allocated once in :meth:`EplbState.add_model` so that device
|
||||
pointers remain stable across CUDA-graph replays. The router kernel
|
||||
indexes this list with ``dbo_current_ubatch_id()``.
|
||||
"""
|
||||
|
||||
|
||||
class EplbState:
|
||||
@@ -253,7 +261,7 @@ class EplbState:
|
||||
Shared scalar bool tensor for all layers. Every
|
||||
:class:`EplbLayerState` holds a reference to the **same** object so
|
||||
a single ``.fill_()`` updates all layers at once. Allocated on the
|
||||
first call to :meth:`_init_should_record_tensor`.
|
||||
first call to :meth:`_propagate_shared_tensors`.
|
||||
"""
|
||||
self.is_async: bool = False
|
||||
"""
|
||||
@@ -440,12 +448,19 @@ class EplbState:
|
||||
self.policy = EPLB_POLICIES[policy_type]
|
||||
logger.debug("Selected EPLB policy: %s", policy_type)
|
||||
|
||||
# num_ubatches is 0 when DBO is disabled.
|
||||
num_ubatches = max(1, self.parallel_config.num_ubatches)
|
||||
num_unpadded_tokens_tensors = [
|
||||
torch.tensor(0, dtype=torch.int32, device=self.device)
|
||||
for _ in range(num_ubatches)
|
||||
]
|
||||
|
||||
model.set_eplb_state(
|
||||
expert_load_pass,
|
||||
logical_to_physical_map,
|
||||
logical_replica_count,
|
||||
)
|
||||
self._init_should_record_tensor(model)
|
||||
self._propagate_shared_tensors(model, num_unpadded_tokens_tensors)
|
||||
expert_buffer = [torch.empty_like(w) for w in model.expert_weights[0]]
|
||||
|
||||
assert self.parallel_config.eplb_config.communicator is not None, (
|
||||
@@ -471,10 +486,43 @@ class EplbState:
|
||||
eplb_stats=None,
|
||||
cuda_device_index=self.cuda_device_index,
|
||||
communicator=communicator,
|
||||
num_unpadded_tokens_tensors=num_unpadded_tokens_tensors,
|
||||
)
|
||||
self.model_states[model_config.compute_hash()] = model_state
|
||||
self.num_valid_physical_experts = model.num_physical_experts
|
||||
|
||||
def prepare_forward(
|
||||
self,
|
||||
model_config: ModelConfig,
|
||||
num_unpadded_tokens: int,
|
||||
ubatch_slices: list | None = None,
|
||||
) -> None:
|
||||
"""Fill the per-[u]batch ``num_unpadded_tokens`` tensors before a
|
||||
forward pass.
|
||||
|
||||
Args:
|
||||
model_config: Identifies which ``EplbModelState`` to update.
|
||||
num_unpadded_tokens: Total number of real (non-padding) tokens
|
||||
in the batch.
|
||||
ubatch_slices: When DBO is active, a list of
|
||||
``UBatchSlice`` objects describing each micro-batch's
|
||||
token range. When ``None``, only ``tensors[0]`` is filled.
|
||||
"""
|
||||
model_state = self.model_states.get(compute_hash_cached(model_config))
|
||||
if model_state is None or model_state.num_unpadded_tokens_tensors is None:
|
||||
return
|
||||
tensors = model_state.num_unpadded_tokens_tensors
|
||||
if ubatch_slices is None:
|
||||
tensors[0].fill_(num_unpadded_tokens)
|
||||
else:
|
||||
for i, ubatch_slice in enumerate(ubatch_slices):
|
||||
ts = ubatch_slice.token_slice
|
||||
# Real tokens in this ubatch: clamp the global count into
|
||||
# the slice range so partially-filled ubatches get the
|
||||
# correct count.
|
||||
val = max(0, min(num_unpadded_tokens, ts.stop) - ts.start)
|
||||
tensors[i].fill_(val)
|
||||
|
||||
def step(
|
||||
self,
|
||||
is_dummy: bool = False,
|
||||
@@ -638,11 +686,20 @@ class EplbState:
|
||||
self._should_record_current_step(log_stats=log_stats)
|
||||
)
|
||||
|
||||
def _init_should_record_tensor(self, model: "MixtureOfExperts") -> None: # type: ignore[name-defined]
|
||||
"""Allocate (once) and propagate the shared ``should_record_tensor``.
|
||||
def _propagate_shared_tensors(
|
||||
self,
|
||||
model: "MixtureOfExperts", # type: ignore[name-defined]
|
||||
num_unpadded_tokens_tensors: list[torch.Tensor],
|
||||
) -> None:
|
||||
"""Propagate shared tensors to every :class:`EplbLayerState`.
|
||||
|
||||
Allocates ``should_record_tensor`` on the first call and then
|
||||
assigns both it and ``num_unpadded_tokens_tensors`` to every
|
||||
MoE layer's :class:`EplbLayerState`. All layers reference the
|
||||
**same** objects so a single update is visible everywhere.
|
||||
|
||||
Must be called after :meth:`model.set_eplb_state` so that each
|
||||
layer's ``eplb_state`` is already populated with the tensor views.
|
||||
layer's ``eplb_state`` is already populated.
|
||||
"""
|
||||
layer_states = [
|
||||
layer.eplb_state
|
||||
@@ -659,6 +716,7 @@ class EplbState:
|
||||
for ls in layer_states:
|
||||
if ls is not None:
|
||||
ls.should_record_tensor = self.should_record_tensor
|
||||
ls.num_unpadded_tokens_tensors = num_unpadded_tokens_tensors
|
||||
|
||||
def rearrange(
|
||||
self,
|
||||
@@ -985,6 +1043,11 @@ class EplbLayerState:
|
||||
sliding window before the next rearrangement, so recording them wastes
|
||||
GPU work.
|
||||
"""
|
||||
num_unpadded_tokens_tensors: list[torch.Tensor] | None = None
|
||||
"""
|
||||
Reference to the parent :class:`EplbModelState`'s tensor list so the
|
||||
router can read the correct per-[u]batch unpadded token count.
|
||||
"""
|
||||
|
||||
def set_layer_state(
|
||||
self,
|
||||
|
||||
@@ -11,6 +11,7 @@ from vllm.model_executor.layers.fused_moe.router.fused_moe_router import (
|
||||
)
|
||||
from vllm.platforms import current_platform
|
||||
from vllm.triton_utils import tl, triton
|
||||
from vllm.v1.worker.ubatching import dbo_current_ubatch_id
|
||||
|
||||
if current_platform.is_cuda_alike():
|
||||
|
||||
@@ -22,11 +23,13 @@ if current_platform.is_cuda_alike():
|
||||
out_ids_ptr,
|
||||
out_ptr,
|
||||
record_enabled_ptr,
|
||||
num_unpadded_tokens_ptr,
|
||||
num_logical_experts,
|
||||
map_slots,
|
||||
out_size,
|
||||
numel,
|
||||
num_active_experts,
|
||||
HAS_NUM_UNPADDED: tl.constexpr,
|
||||
BLOCK_SIZE: tl.constexpr,
|
||||
):
|
||||
pid = tl.program_id(0)
|
||||
@@ -50,6 +53,13 @@ if current_platform.is_cuda_alike():
|
||||
token_idx = (offs // num_active_experts).to(tl.int64)
|
||||
hashed = (token_idx * KNUTH_MULTIPLIER) & 0xFFFFFFFF
|
||||
replica_idx = hashed % replica_count
|
||||
map_index = safe_expert_id * map_slots + replica_idx
|
||||
physical_id = tl.load(
|
||||
logical_to_physical_ptr + map_index,
|
||||
mask=mask & valid_expert,
|
||||
other=-1,
|
||||
)
|
||||
tl.store(out_ids_ptr + offs, physical_id, mask=mask)
|
||||
|
||||
# 2. Record expert load metrics.
|
||||
|
||||
@@ -64,16 +74,21 @@ if current_platform.is_cuda_alike():
|
||||
# If later refactor moved all the MoE kernel calls
|
||||
# to the modular kernel, we can move this logic there
|
||||
# to achieve better efficiency.
|
||||
map_index = safe_expert_id * map_slots + replica_idx
|
||||
physical_id = tl.load(
|
||||
logical_to_physical_ptr + map_index,
|
||||
mask=mask & valid_expert,
|
||||
other=-1,
|
||||
)
|
||||
tl.store(out_ids_ptr + offs, physical_id, mask=mask)
|
||||
|
||||
record_enabled = tl.load(record_enabled_ptr) != 0
|
||||
valid = mask & record_enabled & (physical_id >= 0) & (physical_id < out_size)
|
||||
# Skip padded tokens when recording.
|
||||
if HAS_NUM_UNPADDED:
|
||||
num_unpadded_tokens = tl.load(num_unpadded_tokens_ptr)
|
||||
is_unpadded = offs < num_unpadded_tokens * num_active_experts
|
||||
else:
|
||||
is_unpadded = True
|
||||
valid = (
|
||||
mask
|
||||
& record_enabled
|
||||
& is_unpadded
|
||||
& (physical_id >= 0)
|
||||
& (physical_id < out_size)
|
||||
)
|
||||
safe_physical_id = tl.where(physical_id >= 0, physical_id, 0)
|
||||
tl.atomic_add(out_ptr + safe_physical_id, 1, mask=valid)
|
||||
|
||||
@@ -83,6 +98,7 @@ if current_platform.is_cuda_alike():
|
||||
logical_replica_count: torch.Tensor,
|
||||
expert_load_view: torch.Tensor,
|
||||
record_enabled: torch.Tensor,
|
||||
num_unpadded_tokens: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
topk_ids_in = topk_ids.contiguous().to(dtype=torch.int32)
|
||||
numel = topk_ids_in.numel()
|
||||
@@ -99,11 +115,13 @@ if current_platform.is_cuda_alike():
|
||||
out_flat,
|
||||
expert_load_view,
|
||||
record_enabled,
|
||||
num_unpadded_tokens,
|
||||
logical_replica_count.shape[0],
|
||||
logical_to_physical_map.shape[1],
|
||||
expert_load_view.shape[0],
|
||||
numel,
|
||||
num_active_experts,
|
||||
HAS_NUM_UNPADDED=num_unpadded_tokens is not None,
|
||||
BLOCK_SIZE=256,
|
||||
)
|
||||
return out_flat.reshape(topk_ids.shape)
|
||||
@@ -114,6 +132,7 @@ if current_platform.is_cuda_alike():
|
||||
logical_to_physical_map: torch.Tensor,
|
||||
logical_replica_count: torch.Tensor,
|
||||
record_enabled: torch.Tensor,
|
||||
num_unpadded_tokens: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
# Fused triton implementation: mapping + optional recording in one kernel.
|
||||
return _eplb_map_and_record_triton(
|
||||
@@ -122,6 +141,7 @@ if current_platform.is_cuda_alike():
|
||||
logical_replica_count=logical_replica_count,
|
||||
expert_load_view=expert_load_view,
|
||||
record_enabled=record_enabled,
|
||||
num_unpadded_tokens=num_unpadded_tokens,
|
||||
)
|
||||
else:
|
||||
|
||||
@@ -131,6 +151,7 @@ else:
|
||||
logical_to_physical_map: torch.Tensor,
|
||||
logical_replica_count: torch.Tensor,
|
||||
record_enabled: torch.Tensor,
|
||||
num_unpadded_tokens: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
return topk_ids
|
||||
|
||||
@@ -177,6 +198,8 @@ class BaseRouter(FusedMoERouter):
|
||||
raise ValueError("EPLB requires logical_replica_count != None")
|
||||
if eplb_state.should_record_tensor is None:
|
||||
raise ValueError("EPLB requires should_record_tensor != None")
|
||||
if eplb_state.num_unpadded_tokens_tensors is None:
|
||||
raise ValueError("EPLB requires num_unpadded_tokens_tensors != None")
|
||||
|
||||
def _apply_eplb_mapping(self, topk_ids: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply EPLB mapping to convert logical expert IDs to physical expert IDs."""
|
||||
@@ -186,12 +209,16 @@ class BaseRouter(FusedMoERouter):
|
||||
assert eplb_state.logical_to_physical_map is not None
|
||||
assert eplb_state.logical_replica_count is not None
|
||||
assert eplb_state.should_record_tensor is not None
|
||||
assert eplb_state.num_unpadded_tokens_tensors is not None
|
||||
return eplb_map_to_physical_and_record(
|
||||
topk_ids=topk_ids,
|
||||
logical_to_physical_map=eplb_state.logical_to_physical_map,
|
||||
logical_replica_count=eplb_state.logical_replica_count,
|
||||
expert_load_view=eplb_state.expert_load_view,
|
||||
record_enabled=eplb_state.should_record_tensor,
|
||||
num_unpadded_tokens=eplb_state.num_unpadded_tokens_tensors[
|
||||
dbo_current_ubatch_id()
|
||||
],
|
||||
)
|
||||
return topk_ids
|
||||
|
||||
|
||||
@@ -70,6 +70,7 @@ from vllm.platforms import current_platform
|
||||
from vllm.sequence import IntermediateTensors
|
||||
from vllm.utils.math_utils import cdiv
|
||||
from vllm.v1.attention.backends.registry import AttentionBackendEnum
|
||||
from vllm.v1.worker.ubatching import dbo_current_ubatch_id
|
||||
|
||||
|
||||
class DeepseekV4MLP(nn.Module):
|
||||
@@ -464,6 +465,11 @@ class DeepseekV4MegaMoEExperts(nn.Module):
|
||||
logical_to_physical_map=eplb_state.logical_to_physical_map,
|
||||
logical_replica_count=eplb_state.logical_replica_count,
|
||||
record_enabled=eplb_state.should_record_tensor,
|
||||
num_unpadded_tokens=eplb_state.num_unpadded_tokens_tensors[
|
||||
dbo_current_ubatch_id()
|
||||
]
|
||||
if eplb_state.num_unpadded_tokens_tensors is not None
|
||||
else None,
|
||||
)
|
||||
|
||||
prepare_megamoe_inputs(
|
||||
|
||||
@@ -9,6 +9,7 @@ import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from vllm.config import CUDAGraphMode, VllmConfig, get_layers_from_vllm_config
|
||||
from vllm.distributed.eplb.eplb_state import EplbState
|
||||
from vllm.forward_context import set_forward_context
|
||||
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
|
||||
from vllm.model_executor.model_loader import get_model
|
||||
@@ -43,6 +44,8 @@ class ExtractHiddenStatesProposer:
|
||||
self.dtype = vllm_config.model_config.dtype
|
||||
self.dp_rank = vllm_config.parallel_config.data_parallel_rank
|
||||
|
||||
self.eplb_state: EplbState | None = None
|
||||
|
||||
# Model and attention layer tracking (initialized in load_model)
|
||||
self.model: nn.Module | None = None
|
||||
self.attn_layer_names: list[str] = []
|
||||
@@ -83,6 +86,10 @@ class ExtractHiddenStatesProposer:
|
||||
self.max_num_tokens, dtype=torch.int64, device=device
|
||||
)
|
||||
|
||||
def set_eplb_state(self, eplb_state: EplbState) -> None:
|
||||
"""Inject EPLB state after construction."""
|
||||
self.eplb_state = eplb_state
|
||||
|
||||
def propose(
|
||||
self,
|
||||
num_speculative_tokens: int,
|
||||
@@ -145,6 +152,12 @@ class ExtractHiddenStatesProposer:
|
||||
if num_tokens_across_dp is not None:
|
||||
num_tokens_across_dp[self.dp_rank] = num_input_tokens
|
||||
|
||||
if self.eplb_state is not None:
|
||||
assert self.vllm_config.speculative_config is not None
|
||||
self.eplb_state.prepare_forward(
|
||||
self.vllm_config.speculative_config.draft_model_config,
|
||||
num_tokens,
|
||||
)
|
||||
with set_forward_context(
|
||||
per_layer_attn_metadata,
|
||||
self.vllm_config,
|
||||
|
||||
@@ -15,6 +15,7 @@ from vllm.config import (
|
||||
get_layers_from_vllm_config,
|
||||
replace,
|
||||
)
|
||||
from vllm.distributed.eplb.eplb_state import EplbState
|
||||
from vllm.distributed.parallel_state import get_pp_group
|
||||
from vllm.forward_context import set_forward_context
|
||||
from vllm.logger import init_logger
|
||||
@@ -79,6 +80,7 @@ class SpecDecodeBaseProposer:
|
||||
self.dtype = vllm_config.model_config.dtype
|
||||
self.max_model_len = vllm_config.model_config.max_model_len
|
||||
self.dp_rank = vllm_config.parallel_config.data_parallel_rank
|
||||
self.eplb_state: EplbState | None = None
|
||||
self.num_speculative_tokens = self.speculative_config.num_speculative_tokens
|
||||
|
||||
# We need to get the hidden size from the draft model config because
|
||||
@@ -328,6 +330,10 @@ class SpecDecodeBaseProposer:
|
||||
"does not support M-RoPE yet"
|
||||
)
|
||||
|
||||
def set_eplb_state(self, eplb_state: EplbState) -> None:
|
||||
"""Inject EPLB state after construction."""
|
||||
self.eplb_state = eplb_state
|
||||
|
||||
def _init_parallel_drafting_params(self):
|
||||
# For parallel drafting, we need the token ID to use for masked slots
|
||||
# And for EAGLE + parallel drafting, we need the hidden state tensor to use
|
||||
@@ -527,6 +533,12 @@ class SpecDecodeBaseProposer:
|
||||
if self._share_mtp_indices and hasattr(self.model.model, "set_skip_topk"):
|
||||
self.model.model.set_skip_topk(False)
|
||||
|
||||
if self.eplb_state is not None:
|
||||
self.eplb_state.prepare_forward(
|
||||
self.draft_model_config,
|
||||
num_tokens,
|
||||
)
|
||||
|
||||
with set_forward_context(
|
||||
per_layer_attn_metadata,
|
||||
self.vllm_config,
|
||||
@@ -672,6 +684,12 @@ class SpecDecodeBaseProposer:
|
||||
if self.pass_hidden_states_to_model:
|
||||
model_kwargs["hidden_states"] = self.hidden_states[:input_batch_size]
|
||||
|
||||
if self.eplb_state is not None:
|
||||
self.eplb_state.prepare_forward(
|
||||
self.draft_model_config,
|
||||
batch_size,
|
||||
)
|
||||
|
||||
with set_forward_context(
|
||||
per_layer_attn_metadata,
|
||||
self.vllm_config,
|
||||
|
||||
@@ -8,6 +8,7 @@ from typing import Any
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from vllm.config import ModelConfig
|
||||
from vllm.distributed.eplb.eplb_state import EplbState
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.models.interfaces import (
|
||||
@@ -90,6 +91,7 @@ class EPLBController:
|
||||
draft_model,
|
||||
speculative_config.draft_model_config,
|
||||
)
|
||||
speculator.set_eplb_state(self.state)
|
||||
self._has_registered_models = True
|
||||
return True
|
||||
|
||||
@@ -135,6 +137,16 @@ class EPLBController:
|
||||
log_stats=self.parallel_config.eplb_config.log_balancedness,
|
||||
)
|
||||
|
||||
def prepare_forward(
|
||||
self,
|
||||
model_config: ModelConfig,
|
||||
num_unpadded_tokens: int,
|
||||
ubatch_slices: list | None = None,
|
||||
) -> None:
|
||||
if self.state is None or not self.parallel_config.enable_eplb:
|
||||
return
|
||||
self.state.prepare_forward(model_config, num_unpadded_tokens, ubatch_slices)
|
||||
|
||||
def setup_from_mapping(
|
||||
self,
|
||||
model: nn.Module,
|
||||
|
||||
@@ -1284,6 +1284,9 @@ class GPUModelRunner(LoRAModelRunnerMixin):
|
||||
model_inputs["intermediate_tensors"] = IntermediateTensors(new_tensors)
|
||||
del intermediate_tensors
|
||||
|
||||
# Update the EPLB meta.
|
||||
self.eplb.prepare_forward(self.model_config, input_batch.num_tokens)
|
||||
|
||||
# Run model.
|
||||
if batch_desc.cg_mode == CUDAGraphMode.FULL:
|
||||
# Use explicit cudagraph replay for FULL mode.
|
||||
|
||||
@@ -213,6 +213,8 @@ class AutoRegressiveSpeculator(DraftModelSpeculator):
|
||||
need_eager=is_profile,
|
||||
)
|
||||
|
||||
self._prepare_eplb_forward(input_batch.num_tokens)
|
||||
|
||||
if prefill_batch_desc.cg_mode == CUDAGraphMode.FULL:
|
||||
# Replay the full graph for draft prefill.
|
||||
assert self.prefill_cudagraph_manager is not None
|
||||
@@ -424,6 +426,8 @@ class AutoRegressiveSpeculator(DraftModelSpeculator):
|
||||
num_tokens_across_dp: torch.Tensor | None,
|
||||
cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE,
|
||||
) -> None:
|
||||
self._prepare_eplb_forward(num_reqs)
|
||||
|
||||
idx_mapping = self.idx_mapping[:num_reqs]
|
||||
positions = self.input_buffers.positions[:num_reqs]
|
||||
# Run the draft model forward pass.
|
||||
|
||||
@@ -278,6 +278,9 @@ class DFlashSpeculator(DraftModelSpeculator):
|
||||
self.hidden_states[:num_target_tokens],
|
||||
self.context_positions[:num_target_tokens],
|
||||
)
|
||||
# DFlash processes all speculative tokens in one forward pass,
|
||||
# so the real token count is num_query_tokens.
|
||||
self._prepare_eplb_forward(num_query_tokens)
|
||||
self._generate_draft(
|
||||
num_reqs,
|
||||
num_query_tokens,
|
||||
@@ -354,6 +357,10 @@ class DFlashSpeculator(DraftModelSpeculator):
|
||||
self.kv_cache_config,
|
||||
)
|
||||
|
||||
# DFlash processes all speculative tokens in one forward pass,
|
||||
# so the real token count is num_query_tokens.
|
||||
self._prepare_eplb_forward(num_query_tokens)
|
||||
|
||||
if batch_desc.cg_mode == CUDAGraphMode.FULL:
|
||||
assert self.query_cudagraph_manager is not None
|
||||
self.query_cudagraph_manager.run_fullgraph(batch_desc)
|
||||
|
||||
@@ -8,6 +8,7 @@ import torch.nn as nn
|
||||
|
||||
from vllm.config import VllmConfig, get_layers_from_vllm_config
|
||||
from vllm.config.compilation import CUDAGraphMode
|
||||
from vllm.distributed.eplb.eplb_state import EplbState
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.layers.attention_layer_base import AttentionLayerBase
|
||||
from vllm.v1.kv_cache_interface import KVCacheConfig
|
||||
@@ -106,6 +107,8 @@ class DraftModelSpeculator(BaseSpeculator):
|
||||
self.dp_size = vllm_config.parallel_config.data_parallel_size
|
||||
self.dp_rank = vllm_config.parallel_config.data_parallel_rank
|
||||
|
||||
self.eplb_state: EplbState | None = None
|
||||
|
||||
self.input_buffers = InputBuffers(
|
||||
max_num_reqs=self.max_num_reqs,
|
||||
max_num_tokens=self.max_num_tokens,
|
||||
@@ -165,6 +168,18 @@ class DraftModelSpeculator(BaseSpeculator):
|
||||
)
|
||||
self.draft_attn_layer_names = all_attn_layers - target_attn_layer_names
|
||||
|
||||
def set_eplb_state(self, eplb_state: EplbState) -> None:
|
||||
"""Inject EPLB state after construction."""
|
||||
self.eplb_state = eplb_state
|
||||
|
||||
def _prepare_eplb_forward(self, num_unpadded_tokens: int) -> None:
|
||||
"""Call EPLB prepare_forward if EPLB is active for the draft model."""
|
||||
if self.eplb_state is not None:
|
||||
self.eplb_state.prepare_forward(
|
||||
self.speculative_config.draft_model_config,
|
||||
num_unpadded_tokens,
|
||||
)
|
||||
|
||||
def set_attn(
|
||||
self,
|
||||
model_state: ModelState,
|
||||
|
||||
@@ -4319,6 +4319,13 @@ class GPUModelRunner(
|
||||
# When spec decode is enabled, defer connector finalization
|
||||
# (wait_for_save + clear metadata) until after draft model runs.
|
||||
defer_kv_connector_finalize = self.speculative_config is not None
|
||||
# Update the EPLB meta.
|
||||
if self.eplb_state is not None:
|
||||
self.eplb_state.prepare_forward(
|
||||
self.model_config,
|
||||
num_tokens_unpadded,
|
||||
ubatch_slices_padded,
|
||||
)
|
||||
with (
|
||||
set_forward_context(
|
||||
attn_metadata,
|
||||
@@ -5215,6 +5222,8 @@ class GPUModelRunner(
|
||||
self.drafter.model,
|
||||
spec_config.draft_model_config,
|
||||
)
|
||||
assert hasattr(self.drafter, "set_eplb_state")
|
||||
self.drafter.set_eplb_state(self.eplb_state)
|
||||
eplb_models += 1
|
||||
|
||||
self._setup_eagle3_aux_hidden_state_outputs()
|
||||
|
||||
Reference in New Issue
Block a user