refactor(runner): add EagerRunner, own the eager path, polymorphic dispatch (#28386)
Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
ab0714d0ee
commit
d705a91de1
@@ -887,3 +887,48 @@ def build_prefill_registry(
|
|||||||
)
|
)
|
||||||
reg.register_slot(slot, bind=bind)
|
reg.register_slot(slot, bind=bind)
|
||||||
return reg
|
return reg
|
||||||
|
|
||||||
|
|
||||||
|
def build_eager_registry(
|
||||||
|
*,
|
||||||
|
device: torch.device,
|
||||||
|
max_bs: int,
|
||||||
|
max_num_token: int,
|
||||||
|
cache_loc_dtype: torch.dtype,
|
||||||
|
enable_mamba_track: bool = False,
|
||||||
|
is_encoder_decoder: bool = False,
|
||||||
|
encoder_len_fill_value: int = 0,
|
||||||
|
dp_size: int = 1,
|
||||||
|
) -> CudaGraphBufferRegistry:
|
||||||
|
"""One fixed-max input registry for the ``EagerRunner``, serving BOTH eager
|
||||||
|
decode and eager prefill.
|
||||||
|
|
||||||
|
The decode slot set is a superset of eager prefill's needs (eager prefill
|
||||||
|
carries ``input_embeds`` from the batch and reads the bs-axis fields live),
|
||||||
|
so we reuse it, sized at ``(max_bs, max_num_token)`` where ``max_num_token``
|
||||||
|
is the prefill token ceiling. ``seq_len_fill_value=0`` because eager never
|
||||||
|
pads, so the sentinel tail is never read.
|
||||||
|
|
||||||
|
``share_pool=True`` so same-named / same-size slots coalesce through the
|
||||||
|
process-wide pool. The ``EagerRunner`` is built before the cuda-graph runners
|
||||||
|
(see ``ModelRunner.init_backends``), so its (largest) allocations are
|
||||||
|
canonical and the cg runners' matching slots (prefill's token-axis at
|
||||||
|
``max_num_token``, decode's bs-axis at ``max_bs``) adopt them.
|
||||||
|
"""
|
||||||
|
return build_decode_registry(
|
||||||
|
device=device,
|
||||||
|
max_bs=max_bs,
|
||||||
|
max_num_token=max_num_token,
|
||||||
|
seq_len_fill_value=0,
|
||||||
|
cache_loc_dtype=cache_loc_dtype,
|
||||||
|
enable_mamba_track=enable_mamba_track,
|
||||||
|
is_encoder_decoder=is_encoder_decoder,
|
||||||
|
encoder_len_fill_value=encoder_len_fill_value,
|
||||||
|
enable_num_token_non_padded=False,
|
||||||
|
register_global_num_tokens=False,
|
||||||
|
require_gathered_buffer=False,
|
||||||
|
require_mlp_tp_gather=False,
|
||||||
|
dp_size=dp_size,
|
||||||
|
share_pool=True,
|
||||||
|
source=None,
|
||||||
|
)
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ import socket
|
|||||||
import threading
|
import threading
|
||||||
import time
|
import time
|
||||||
from collections import defaultdict
|
from collections import defaultdict
|
||||||
from dataclasses import dataclass, replace
|
from dataclasses import dataclass
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Callable, List, Optional, Tuple, Union
|
from typing import Any, Callable, List, Optional, Tuple, Union
|
||||||
|
|
||||||
@@ -126,11 +126,7 @@ from sglang.srt.layers.attention.attention_registry import (
|
|||||||
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
||||||
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
from sglang.srt.layers.attention.tbo_backend import TboAttnBackend
|
||||||
from sglang.srt.layers.cp.utils import (
|
from sglang.srt.layers.cp.utils import (
|
||||||
cp_gather_after_forward,
|
|
||||||
cp_split_before_forward,
|
|
||||||
get_cp_strategy,
|
get_cp_strategy,
|
||||||
is_cp_v2_active,
|
|
||||||
prepare_cp_forward,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
@@ -143,7 +139,6 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.layers.moe.hash_topk import HashTopK
|
from sglang.srt.layers.moe.hash_topk import HashTopK
|
||||||
from sglang.srt.layers.moe.topk import TopK
|
from sglang.srt.layers.moe.topk import TopK
|
||||||
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
|
|
||||||
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
from sglang.srt.layers.quantization.fp8_kernel import fp8_dtype
|
||||||
from sglang.srt.layers.sampler import create_sampler
|
from sglang.srt.layers.sampler import create_sampler
|
||||||
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
|
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
|
||||||
@@ -154,11 +149,6 @@ from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
|
|||||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner
|
||||||
from sglang.srt.model_executor.cuda_graph_buffer_registry import (
|
|
||||||
CudaGraphBufferRegistry,
|
|
||||||
build_decode_registry,
|
|
||||||
build_prefill_registry,
|
|
||||||
)
|
|
||||||
from sglang.srt.model_executor.cuda_graph_config import (
|
from sglang.srt.model_executor.cuda_graph_config import (
|
||||||
Backend,
|
Backend,
|
||||||
Phase,
|
Phase,
|
||||||
@@ -182,15 +172,12 @@ from sglang.srt.model_executor.model_runner_kv_cache_mixin import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||||
from sglang.srt.model_executor.runner import (
|
from sglang.srt.model_executor.runner import (
|
||||||
|
EagerRunner,
|
||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||||
_allocate_decode_buffers,
|
_allocate_decode_buffers,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
|
||||||
enable_tc_piecewise_cuda_graph,
|
|
||||||
set_tc_piecewise_forward_context,
|
|
||||||
)
|
|
||||||
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
|
from sglang.srt.model_loader.loader import DefaultModelLoader, get_model_loader
|
||||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||||
RemoteInstanceWeightLoaderBackend,
|
RemoteInstanceWeightLoaderBackend,
|
||||||
@@ -244,7 +231,7 @@ from sglang.srt.utils import (
|
|||||||
set_cuda_arch,
|
set_cuda_arch,
|
||||||
slow_rank_detector,
|
slow_rank_detector,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.common import ceil_align, next_power_of_2, require_mlp_sync
|
from sglang.srt.utils.common import ceil_align, require_mlp_sync
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||||
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
|
||||||
from sglang.srt.utils.nvtx_utils import profile_range
|
from sglang.srt.utils.nvtx_utils import profile_range
|
||||||
@@ -369,14 +356,6 @@ class ModelRunnerOutput:
|
|||||||
indexer_topk_output: Optional[TopkCaptureOutput] = None
|
indexer_topk_output: Optional[TopkCaptureOutput] = None
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class _EagerBufferRegistry:
|
|
||||||
# Lazily-built eager input-buffer registry plus the capacity it was sized to.
|
|
||||||
registry: Optional[CudaGraphBufferRegistry] = None
|
|
||||||
max_bs: int = 0
|
|
||||||
max_num_tokens: int = 0
|
|
||||||
|
|
||||||
|
|
||||||
class ModelRunner(ModelRunnerKVCacheMixin):
|
class ModelRunner(ModelRunnerKVCacheMixin):
|
||||||
"""ModelRunner runs the forward passes of the models."""
|
"""ModelRunner runs the forward passes of the models."""
|
||||||
|
|
||||||
@@ -453,8 +432,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.enable_elastic_ep = server_args.elastic_ep_backend is not None
|
self.enable_elastic_ep = server_args.elastic_ep_backend is not None
|
||||||
self.forward_pass_id = 0
|
self.forward_pass_id = 0
|
||||||
self.init_new_workspace = False
|
self.init_new_workspace = False
|
||||||
self._eager_decode_registry = _EagerBufferRegistry()
|
|
||||||
self._eager_prefill_registry = _EagerBufferRegistry()
|
|
||||||
self.draft_model_idx = draft_model_idx
|
self.draft_model_idx = draft_model_idx
|
||||||
self.enable_hisparse = server_args.enable_hisparse
|
self.enable_hisparse = server_args.enable_hisparse
|
||||||
|
|
||||||
@@ -912,6 +889,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
# runs with aux hidden state capture enabled.
|
# runs with aux hidden state capture enabled.
|
||||||
self.init_aux_hidden_state_capture()
|
self.init_aux_hidden_state_capture()
|
||||||
|
|
||||||
|
# The eager (no-cuda-graph) phase runner. Always built: it serves both
|
||||||
|
# the fully-disabled case (decode/prefill runners point at it) and the
|
||||||
|
# eager fallback when a cuda-graph runner can't run a batch.
|
||||||
|
self.eager_runner = EagerRunner(self)
|
||||||
|
|
||||||
if self.device == "cuda" or self.device == "musa":
|
if self.device == "cuda" or self.device == "musa":
|
||||||
self.init_cublas()
|
self.init_cublas()
|
||||||
self.init_attention_backend()
|
self.init_attention_backend()
|
||||||
@@ -951,7 +933,10 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.init_attention_backend()
|
self.init_attention_backend()
|
||||||
|
|
||||||
if disable_cuda_graph:
|
if disable_cuda_graph:
|
||||||
self.decode_cuda_graph_runner = None
|
# Decode cuda graph disabled: route eager decode through the
|
||||||
|
# EagerRunner (the dispatch gate isinstance(..., EagerRunner) keeps
|
||||||
|
# _forward_raw off any replay branch).
|
||||||
|
self.decode_cuda_graph_runner = self.eager_runner
|
||||||
self.graph_mem_usage = 0
|
self.graph_mem_usage = 0
|
||||||
|
|
||||||
if server_args.forward_hooks:
|
if server_args.forward_hooks:
|
||||||
@@ -3040,6 +3025,11 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
"resolved prefill.backend='disabled' (e.g. via "
|
"resolved prefill.backend='disabled' (e.g. via "
|
||||||
"--cuda-graph-backend-prefill=disabled or auto-disable rules)."
|
"--cuda-graph-backend-prefill=disabled or auto-disable rules)."
|
||||||
)
|
)
|
||||||
|
# Prefill cuda graph disabled: route eager prefill through the
|
||||||
|
# EagerRunner (its can_run_graph returns False, so _forward_raw's
|
||||||
|
# extend branch falls through to the eager path).
|
||||||
|
if not self.is_draft_worker:
|
||||||
|
self.prefill_cuda_graph_runner = self.eager_runner
|
||||||
return
|
return
|
||||||
|
|
||||||
# Draft models skip here during __init__; the eagle worker calls
|
# Draft models skip here during __init__; the eagle worker calls
|
||||||
@@ -3208,111 +3198,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
def update_decode_attn_backend(self, stream_idx: int):
|
def update_decode_attn_backend(self, stream_idx: int):
|
||||||
self.decode_attn_backend = self.decode_attn_backend_group[stream_idx]
|
self.decode_attn_backend = self.decode_attn_backend_group[stream_idx]
|
||||||
|
|
||||||
def _ensure_eager_registry(
|
|
||||||
self,
|
|
||||||
cache: _EagerBufferRegistry,
|
|
||||||
raw_bs: int,
|
|
||||||
raw_num_tokens: int,
|
|
||||||
build: Callable[[int, int], CudaGraphBufferRegistry],
|
|
||||||
) -> CudaGraphBufferRegistry:
|
|
||||||
# Built on first use and grown (next power of two) when a batch exceeds
|
|
||||||
# the current capacity.
|
|
||||||
if (
|
|
||||||
cache.registry is not None
|
|
||||||
and raw_bs <= cache.max_bs
|
|
||||||
and raw_num_tokens <= cache.max_num_tokens
|
|
||||||
):
|
|
||||||
return cache.registry
|
|
||||||
cache.max_bs = next_power_of_2(max(raw_bs, cache.max_bs))
|
|
||||||
cache.max_num_tokens = next_power_of_2(
|
|
||||||
max(raw_num_tokens, cache.max_num_tokens)
|
|
||||||
)
|
|
||||||
cache.registry = build(cache.max_bs, cache.max_num_tokens)
|
|
||||||
return cache.registry
|
|
||||||
|
|
||||||
def _ensure_eager_decode_registry(
|
|
||||||
self, raw_bs: int, raw_num_tokens: int
|
|
||||||
) -> CudaGraphBufferRegistry:
|
|
||||||
is_encoder_decoder = self.model_config.is_encoder_decoder
|
|
||||||
return self._ensure_eager_registry(
|
|
||||||
self._eager_decode_registry,
|
|
||||||
raw_bs,
|
|
||||||
raw_num_tokens,
|
|
||||||
lambda bs, num_tokens: build_decode_registry(
|
|
||||||
device=self.device,
|
|
||||||
max_bs=bs,
|
|
||||||
max_num_token=num_tokens,
|
|
||||||
# Eager has no padding so this sentinel is never read; 0 avoids the
|
|
||||||
# cuda-graph-only fill-value method that some backends lack.
|
|
||||||
seq_len_fill_value=0,
|
|
||||||
cache_loc_dtype=torch.int64,
|
|
||||||
enable_mamba_track=(
|
|
||||||
self.server_args.enable_mamba_extra_buffer()
|
|
||||||
and self.spec_algorithm.is_none()
|
|
||||||
),
|
|
||||||
is_encoder_decoder=is_encoder_decoder,
|
|
||||||
encoder_len_fill_value=(
|
|
||||||
getattr(self.model_config.hf_config, "max_source_positions", 0)
|
|
||||||
if is_encoder_decoder
|
|
||||||
else 0
|
|
||||||
),
|
|
||||||
enable_num_token_non_padded=False,
|
|
||||||
register_global_num_tokens=False,
|
|
||||||
require_gathered_buffer=False,
|
|
||||||
require_mlp_tp_gather=False,
|
|
||||||
dp_size=self.server_args.dp_size,
|
|
||||||
share_pool=False,
|
|
||||||
source=None,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _ensure_eager_prefill_registry(
|
|
||||||
self, raw_bs: int, raw_num_tokens: int
|
|
||||||
) -> CudaGraphBufferRegistry:
|
|
||||||
return self._ensure_eager_registry(
|
|
||||||
self._eager_prefill_registry,
|
|
||||||
raw_bs,
|
|
||||||
raw_num_tokens,
|
|
||||||
lambda bs, num_tokens: build_prefill_registry(
|
|
||||||
device=self.device,
|
|
||||||
max_bs=bs,
|
|
||||||
max_num_token=num_tokens,
|
|
||||||
cache_loc_dtype=torch.int64,
|
|
||||||
is_multimodal=self.is_multimodal,
|
|
||||||
enable_mamba_track=False,
|
|
||||||
register_input_embeds=False,
|
|
||||||
share_pool=False,
|
|
||||||
source=None,
|
|
||||||
),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _eager_fb_view(
|
|
||||||
self, forward_batch: ForwardBatch, pp_proxy_tensors=None
|
|
||||||
) -> ForwardBatch:
|
|
||||||
if envs.SGLANG_EAGER_INPUT_NO_COPY.get():
|
|
||||||
return replace(forward_batch)
|
|
||||||
raw_bs = forward_batch.batch_size
|
|
||||||
raw_num_tokens = forward_batch.input_ids.shape[0]
|
|
||||||
ensure = (
|
|
||||||
self._ensure_eager_prefill_registry
|
|
||||||
if forward_batch.forward_mode.is_extend(include_draft_extend_v2=True)
|
|
||||||
else self._ensure_eager_decode_registry
|
|
||||||
)
|
|
||||||
registry = ensure(raw_bs, raw_num_tokens)
|
|
||||||
registry.fill_from(
|
|
||||||
forward_batch,
|
|
||||||
raw_bs=raw_bs,
|
|
||||||
padded_bs=raw_bs,
|
|
||||||
raw_num_tokens=raw_num_tokens,
|
|
||||||
padded_num_tokens=raw_num_tokens,
|
|
||||||
pp_proxy_tensors=pp_proxy_tensors,
|
|
||||||
)
|
|
||||||
return registry.extract_buffer(
|
|
||||||
padded_bs=raw_bs,
|
|
||||||
padded_num_tokens=raw_num_tokens,
|
|
||||||
forward_batch_template=forward_batch,
|
|
||||||
)
|
|
||||||
|
|
||||||
def _prepare_eager_forward_batch(self, forward_batch: ForwardBatch) -> None:
|
def _prepare_eager_forward_batch(self, forward_batch: ForwardBatch) -> None:
|
||||||
"""Pad / normalize a batch for the eager (non-cuda-graph) forward.
|
"""Pad / normalize a batch for the eager (non-cuda-graph) forward.
|
||||||
|
|
||||||
@@ -3321,6 +3206,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
forward path needs — the cuda-graph path does the equivalent inside the
|
forward path needs — the cuda-graph path does the equivalent inside the
|
||||||
runner's capture/replay, so this is skipped there.
|
runner's capture/replay, so this is skipped there.
|
||||||
"""
|
"""
|
||||||
|
# For MLP sync
|
||||||
if forward_batch.global_num_tokens_cpu is not None:
|
if forward_batch.global_num_tokens_cpu is not None:
|
||||||
forward_batch.prepare_mlp_sync_batch(self)
|
forward_batch.prepare_mlp_sync_batch(self)
|
||||||
else:
|
else:
|
||||||
@@ -3343,6 +3229,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Hisparse coordinator — backends now read it from self.model_runner.
|
||||||
if self.hisparse_coordinator is not None:
|
if self.hisparse_coordinator is not None:
|
||||||
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
|
self.hisparse_coordinator.num_real_reqs.fill_(forward_batch.batch_size)
|
||||||
|
|
||||||
@@ -3354,62 +3241,12 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
"""
|
"""
|
||||||
return {"pp_proxy_tensors": pp_proxy_tensors} if self.support_pp else {}
|
return {"pp_proxy_tensors": pp_proxy_tensors} if self.support_pp else {}
|
||||||
|
|
||||||
def forward_decode(
|
def _extend_forward_kwargs(
|
||||||
self,
|
self, forward_batch: ForwardBatch, pp_proxy_tensors
|
||||||
forward_batch: ForwardBatch,
|
) -> dict:
|
||||||
pp_proxy_tensors=None,
|
"""Build the extend/prefill model.forward kwargs (pp_proxy_tensors +
|
||||||
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
input_embeds / replace_embeds overrides + get_embedding), shared by the
|
||||||
if not self.server_args.enable_pdmux:
|
prefill cuda-graph path and the EagerRunner's eager extend path."""
|
||||||
forward_batch = self._eager_fb_view(forward_batch, pp_proxy_tensors)
|
|
||||||
# Set extra arguments
|
|
||||||
pdmux_override = False
|
|
||||||
if forward_batch.needs_forward_metadata_init():
|
|
||||||
if hasattr(self.model, "prepare_forward_batch"):
|
|
||||||
# Prepare model-specific attention metadata before planning,
|
|
||||||
# e.g. Moss-VL's prefill cross-attention custom mask.
|
|
||||||
self.model.prepare_forward_batch(forward_batch)
|
|
||||||
if self.server_args.enable_pdmux:
|
|
||||||
self.decode_attn_backend.init_forward_metadata(forward_batch)
|
|
||||||
# PDmux selects a per-stream backend; publish it to model-layer
|
|
||||||
# readers via the active ForwardContext so RadixAttention etc.
|
|
||||||
# dispatch against the right backend for this forward.
|
|
||||||
pdmux_override = True
|
|
||||||
else:
|
|
||||||
self.attn_backend.init_forward_metadata(forward_batch)
|
|
||||||
# FIXME: add pp_proxy_tensors arg to all models
|
|
||||||
kwargs = self._pp_kwargs(pp_proxy_tensors)
|
|
||||||
|
|
||||||
# Launch forward
|
|
||||||
ctx = (
|
|
||||||
self.device_timer.wrap(metadata={"category": "decode"})
|
|
||||||
if self.device_timer
|
|
||||||
else contextlib.nullcontext()
|
|
||||||
)
|
|
||||||
|
|
||||||
def _do_forward():
|
|
||||||
return self.model.forward(
|
|
||||||
forward_batch.input_ids,
|
|
||||||
forward_batch.positions,
|
|
||||||
forward_batch,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
with ctx:
|
|
||||||
if pdmux_override:
|
|
||||||
with forward_context(
|
|
||||||
ForwardContext(attn_backend=self.decode_attn_backend)
|
|
||||||
):
|
|
||||||
return _do_forward()
|
|
||||||
return _do_forward()
|
|
||||||
|
|
||||||
def forward_extend(
|
|
||||||
self,
|
|
||||||
forward_batch: ForwardBatch,
|
|
||||||
pp_proxy_tensors=None,
|
|
||||||
) -> Tuple[
|
|
||||||
Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput], bool
|
|
||||||
]:
|
|
||||||
# Setup extra arguments
|
|
||||||
kwargs = self._pp_kwargs(pp_proxy_tensors)
|
kwargs = self._pp_kwargs(pp_proxy_tensors)
|
||||||
if forward_batch.input_embeds is not None:
|
if forward_batch.input_embeds is not None:
|
||||||
kwargs["input_embeds"] = forward_batch.input_embeds.bfloat16()
|
kwargs["input_embeds"] = forward_batch.input_embeds.bfloat16()
|
||||||
@@ -3426,159 +3263,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
if not self.is_generation:
|
if not self.is_generation:
|
||||||
kwargs["get_embedding"] = True
|
kwargs["get_embedding"] = True
|
||||||
|
return kwargs
|
||||||
# Check piecewies cuda graph
|
|
||||||
can_run_graph = (
|
|
||||||
self.prefill_cuda_graph_runner is not None
|
|
||||||
and self.prefill_cuda_graph_runner.can_run_graph(forward_batch)
|
|
||||||
)
|
|
||||||
if get_cp_strategy() is not None:
|
|
||||||
can_run_graph = False
|
|
||||||
if can_run_graph:
|
|
||||||
# TODO: device_timer.wrap is too broad here — it also includes
|
|
||||||
# load_batch time. Move timing into the prefill cuda graph
|
|
||||||
# runner to capture only the model.forward part.
|
|
||||||
ctx = (
|
|
||||||
self.device_timer.wrap(metadata={"category": "extend"})
|
|
||||||
if self.device_timer
|
|
||||||
else contextlib.nullcontext()
|
|
||||||
)
|
|
||||||
with ctx:
|
|
||||||
ret = self.prefill_cuda_graph_runner.execute(forward_batch, **kwargs)
|
|
||||||
return (ret, can_run_graph)
|
|
||||||
|
|
||||||
if not self.server_args.enable_pdmux:
|
|
||||||
forward_batch = self._eager_fb_view(forward_batch, pp_proxy_tensors)
|
|
||||||
|
|
||||||
# Launch model forward
|
|
||||||
if forward_batch.needs_forward_metadata_init():
|
|
||||||
if hasattr(self.model, "prepare_forward_batch"):
|
|
||||||
# Prepare model-specific attention metadata before planning,
|
|
||||||
# e.g. Moss-VL's prefill cross-attention custom mask.
|
|
||||||
self.model.prepare_forward_batch(forward_batch)
|
|
||||||
self.attn_backend.init_forward_metadata(forward_batch)
|
|
||||||
cp_v2_active = is_cp_v2_active(forward_batch)
|
|
||||||
forward_positions = forward_batch.positions
|
|
||||||
if cp_v2_active:
|
|
||||||
prepare_cp_forward(forward_batch)
|
|
||||||
complete_hidden_states = kwargs.get("input_embeds")
|
|
||||||
if complete_hidden_states is None:
|
|
||||||
embed_layer = self.model.get_input_embeddings()
|
|
||||||
complete_hidden_states = embed_layer(forward_batch.input_ids)
|
|
||||||
sharded_hidden_states, sharded_positions = cp_split_before_forward(
|
|
||||||
complete_hidden_states,
|
|
||||||
forward_batch.positions,
|
|
||||||
forward_batch,
|
|
||||||
)
|
|
||||||
kwargs["input_embeds"] = sharded_hidden_states
|
|
||||||
forward_positions = sharded_positions
|
|
||||||
|
|
||||||
ctx = (
|
|
||||||
self.device_timer.wrap(metadata={"category": "extend"})
|
|
||||||
if self.device_timer
|
|
||||||
else contextlib.nullcontext()
|
|
||||||
)
|
|
||||||
with ctx:
|
|
||||||
if (
|
|
||||||
_is_hip
|
|
||||||
and self.prefill_cuda_graph_runner is not None
|
|
||||||
and not cp_v2_active
|
|
||||||
):
|
|
||||||
# AMD/HIP: when PCG is enabled but the batch exceeds max captured
|
|
||||||
# size, run eagerly under enable_tc_piecewise_cuda_graph() and
|
|
||||||
# set_tc_piecewise_forward_context() so that (a) Dynamo guards on
|
|
||||||
# _in_tc_piecewise_cuda_graph stay consistent with the PCG-traced
|
|
||||||
# graph (preventing runtime recompilation) and (b) PCG-specific
|
|
||||||
# code paths (MoE, attention) can access their layer objects.
|
|
||||||
with (
|
|
||||||
enable_tc_piecewise_cuda_graph(),
|
|
||||||
set_tc_piecewise_forward_context(
|
|
||||||
forward_batch,
|
|
||||||
self.attention_layers,
|
|
||||||
getattr(self.model, "quant_config", None),
|
|
||||||
self.moe_layers,
|
|
||||||
self.moe_fusions,
|
|
||||||
dsa_indexers=self.dsa_indexers,
|
|
||||||
),
|
|
||||||
):
|
|
||||||
ret = self.model.forward(
|
|
||||||
forward_batch.input_ids,
|
|
||||||
forward_positions,
|
|
||||||
forward_batch,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
elif cp_v2_active:
|
|
||||||
hidden_states = self.model.model(
|
|
||||||
forward_batch.input_ids,
|
|
||||||
forward_positions,
|
|
||||||
forward_batch,
|
|
||||||
input_embeds=kwargs.get("input_embeds"),
|
|
||||||
pp_proxy_tensors=kwargs.get("pp_proxy_tensors"),
|
|
||||||
)
|
|
||||||
|
|
||||||
aux_hidden_states = None
|
|
||||||
capture_aux_hidden_states = getattr(
|
|
||||||
self.model, "capture_aux_hidden_states", False
|
|
||||||
)
|
|
||||||
if capture_aux_hidden_states:
|
|
||||||
hidden_states, aux_hidden_states = hidden_states
|
|
||||||
|
|
||||||
if self.model.pp_group.is_last_rank:
|
|
||||||
hidden_states = cp_gather_after_forward(
|
|
||||||
hidden_states,
|
|
||||||
forward_batch,
|
|
||||||
torch.cuda.current_stream(),
|
|
||||||
)
|
|
||||||
ret = self.model.logits_processor(
|
|
||||||
forward_batch.input_ids,
|
|
||||||
hidden_states,
|
|
||||||
self.model.lm_head,
|
|
||||||
forward_batch,
|
|
||||||
aux_hidden_states,
|
|
||||||
)
|
|
||||||
elif capture_aux_hidden_states:
|
|
||||||
ret = hidden_states, aux_hidden_states
|
|
||||||
else:
|
|
||||||
ret = hidden_states
|
|
||||||
else:
|
|
||||||
ret = self.model.forward(
|
|
||||||
forward_batch.input_ids,
|
|
||||||
forward_positions,
|
|
||||||
forward_batch,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
return (ret, can_run_graph)
|
|
||||||
|
|
||||||
def forward_idle(
|
|
||||||
self, forward_batch: ForwardBatch, pp_proxy_tensors=None
|
|
||||||
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
|
||||||
# In DP Attention, IDLE batches may be padded (batch_size > 0) for MLP
|
|
||||||
# sync. Reinit metadata for the padded case so attention kernels see
|
|
||||||
# the right batch_size (e.g. DSA Indexer). For the unpadded case
|
|
||||||
# (batch_size == 0) explicitly drop any stale forward_metadata left
|
|
||||||
# over from the previous forward — without this, attention layers
|
|
||||||
# called from the idle path can re-read a prior batch's req_pool
|
|
||||||
# indices and trigger SWA mapping use-after-free.
|
|
||||||
if forward_batch.batch_size > 0:
|
|
||||||
if not self.server_args.enable_pdmux:
|
|
||||||
forward_batch = self._eager_fb_view(forward_batch, pp_proxy_tensors)
|
|
||||||
self.attn_backend.init_forward_metadata(forward_batch)
|
|
||||||
else:
|
|
||||||
self.attn_backend.forward_metadata = None
|
|
||||||
|
|
||||||
kwargs = self._pp_kwargs(pp_proxy_tensors)
|
|
||||||
ctx = (
|
|
||||||
self.device_timer.wrap(metadata={"category": "idle"})
|
|
||||||
if self.device_timer
|
|
||||||
else contextlib.nullcontext()
|
|
||||||
)
|
|
||||||
with ctx:
|
|
||||||
return self.model.forward(
|
|
||||||
forward_batch.input_ids,
|
|
||||||
forward_batch.positions,
|
|
||||||
forward_batch,
|
|
||||||
**kwargs,
|
|
||||||
)
|
|
||||||
|
|
||||||
def forward_split_prefill(
|
def forward_split_prefill(
|
||||||
self,
|
self,
|
||||||
@@ -3737,34 +3422,48 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
)
|
)
|
||||||
return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph)
|
return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph)
|
||||||
|
|
||||||
# DP / MLP-sync padding + attn-tp normalization that the eager
|
# DP / MLP-sync padding + attn-tp normalization. Only the decode
|
||||||
# (non-graph) forward needs. The graph path skips it: capture/replay
|
# cuda-graph path above pre-pads its static buffers and returns
|
||||||
# pads inside the runner.
|
# early; split prefill, the prefill cuda graph, and the eager
|
||||||
|
# forward all run the live batch and need this first — it sets
|
||||||
|
# global_dp_buffer_len / padded token counts that graph eligibility
|
||||||
|
# and the collectives depend on.
|
||||||
self._prepare_eager_forward_batch(forward_batch)
|
self._prepare_eager_forward_batch(forward_batch)
|
||||||
|
|
||||||
# Forward without cuda graph
|
if forward_batch.forward_mode.is_split_prefill():
|
||||||
if forward_batch.forward_mode.is_decode():
|
# Layer-split mode; stays on ModelRunner, not the eager runner.
|
||||||
ret = self.forward_decode(
|
|
||||||
forward_batch,
|
|
||||||
pp_proxy_tensors=pp_proxy_tensors,
|
|
||||||
)
|
|
||||||
elif forward_batch.forward_mode.is_split_prefill():
|
|
||||||
ret = self.forward_split_prefill(
|
ret = self.forward_split_prefill(
|
||||||
forward_batch,
|
forward_batch,
|
||||||
reinit_attn_backend=reinit_attn_backend,
|
reinit_attn_backend=reinit_attn_backend,
|
||||||
forward_count=split_forward_count,
|
forward_count=split_forward_count,
|
||||||
)
|
)
|
||||||
elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True):
|
elif (
|
||||||
ret, can_run_graph = self.forward_extend(
|
forward_batch.forward_mode.is_extend(include_draft_extend_v2=True)
|
||||||
forward_batch,
|
and not isinstance(self.prefill_cuda_graph_runner, EagerRunner)
|
||||||
pp_proxy_tensors=pp_proxy_tensors,
|
and self.prefill_cuda_graph_runner is not None
|
||||||
|
and self.prefill_cuda_graph_runner.can_run_graph(forward_batch)
|
||||||
|
and get_cp_strategy() is None
|
||||||
|
):
|
||||||
|
# Prefill cuda graph (piecewise).
|
||||||
|
kwargs = self._extend_forward_kwargs(forward_batch, pp_proxy_tensors)
|
||||||
|
# TODO: device_timer.wrap is too broad here — it also includes
|
||||||
|
# load_batch time. Move timing into the prefill cuda graph runner
|
||||||
|
# to capture only the model.forward part.
|
||||||
|
ctx = (
|
||||||
|
self.device_timer.wrap(metadata={"category": "extend"})
|
||||||
|
if self.device_timer
|
||||||
|
else contextlib.nullcontext()
|
||||||
)
|
)
|
||||||
elif forward_batch.forward_mode.is_idle():
|
with ctx:
|
||||||
ret = self.forward_idle(
|
ret = self.prefill_cuda_graph_runner.execute(
|
||||||
|
forward_batch, **kwargs
|
||||||
|
)
|
||||||
|
can_run_graph = True
|
||||||
|
else:
|
||||||
|
# Eager: decode / extend / idle dispatched inside the runner.
|
||||||
|
ret = self.eager_runner.execute(
|
||||||
forward_batch, pp_proxy_tensors=pp_proxy_tensors
|
forward_batch, pp_proxy_tensors=pp_proxy_tensors
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode}")
|
|
||||||
|
|
||||||
if (
|
if (
|
||||||
forward_batch.global_num_tokens_cpu is not None
|
forward_batch.global_num_tokens_cpu is not None
|
||||||
|
|||||||
@@ -13,6 +13,9 @@ Public API:
|
|||||||
capture-loop scaffolding on top of BaseRunner.
|
capture-loop scaffolding on top of BaseRunner.
|
||||||
- DecodeCudaGraphRunner — concrete decode-phase runner.
|
- DecodeCudaGraphRunner — concrete decode-phase runner.
|
||||||
- PrefillCudaGraphRunner — concrete prefill-phase runner.
|
- PrefillCudaGraphRunner — concrete prefill-phase runner.
|
||||||
|
- EagerRunner — no-cuda-graph runner; runs model.forward live (the
|
||||||
|
eager dual of the cuda-graph runners), mode-dispatched over decode +
|
||||||
|
extend + idle.
|
||||||
- Buffer dataclasses, capture-mode flags, the global memory pool,
|
- Buffer dataclasses, capture-mode flags, the global memory pool,
|
||||||
and the DeepEP adapter live in
|
and the DeepEP adapter live in
|
||||||
sglang.srt.model_executor.runner_utils; they are
|
sglang.srt.model_executor.runner_utils; they are
|
||||||
@@ -29,6 +32,7 @@ from sglang.srt.model_executor.runner.base_runner import BaseRunner # noqa: F40
|
|||||||
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
from sglang.srt.model_executor.runner.decode_cuda_graph_runner import (
|
||||||
DecodeCudaGraphRunner,
|
DecodeCudaGraphRunner,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.runner.eager_runner import EagerRunner # noqa: F401
|
||||||
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( # noqa: F401
|
from sglang.srt.model_executor.runner.prefill_cuda_graph_runner import ( # noqa: F401
|
||||||
PrefillCudaGraphRunner,
|
PrefillCudaGraphRunner,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -779,6 +779,15 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
if self.enable_profile_cuda_graph:
|
if self.enable_profile_cuda_graph:
|
||||||
profile_context = self._init_profile_context_and_memory_record()
|
profile_context = self._init_profile_context_and_memory_record()
|
||||||
|
|
||||||
|
# share_buffers() coalesces seq_lens / seq_lens_cpu through the process-
|
||||||
|
# wide pool, so they may alias a buffer seeded by an earlier runner (the
|
||||||
|
# eager registry fills them with 0). The capture-time attention-metadata
|
||||||
|
# plan reads these as the per-request KV length, and the prefill wrapper
|
||||||
|
# (DLLM_EXTEND) asserts kv_len >= qo_len, so restore the fill value the
|
||||||
|
# captured graph needs before capturing.
|
||||||
|
self.buffers.seq_lens.fill_(self.seq_len_fill_value)
|
||||||
|
self.buffers.seq_lens_cpu.fill_(self.seq_len_fill_value)
|
||||||
|
|
||||||
# Trigger CUDA graph capture for specific shapes.
|
# Trigger CUDA graph capture for specific shapes.
|
||||||
# Capture the large shapes first so that the smaller shapes
|
# Capture the large shapes first so that the smaller shapes
|
||||||
# can reuse the memory pool allocated for the large shapes.
|
# can reuse the memory pool allocated for the large shapes.
|
||||||
|
|||||||
@@ -0,0 +1,340 @@
|
|||||||
|
# Copyright 2023-2026 SGLang Team
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
# ==============================================================================
|
||||||
|
"""No-cuda-graph phase runner; the eager dual of BaseCudaGraphRunner."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextlib
|
||||||
|
import logging
|
||||||
|
from dataclasses import replace
|
||||||
|
from typing import TYPE_CHECKING, Any, Tuple, Union
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.layers.cp.utils import (
|
||||||
|
cp_gather_after_forward,
|
||||||
|
cp_split_before_forward,
|
||||||
|
is_cp_v2_active,
|
||||||
|
prepare_cp_forward,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.pooler import EmbeddingPoolerOutput
|
||||||
|
from sglang.srt.model_executor.cuda_graph_buffer_registry import (
|
||||||
|
build_eager_registry,
|
||||||
|
)
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||||
|
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||||
|
from sglang.srt.model_executor.runner.base_runner import BaseRunner
|
||||||
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||||
|
enable_tc_piecewise_cuda_graph,
|
||||||
|
set_tc_piecewise_forward_context,
|
||||||
|
)
|
||||||
|
from sglang.srt.utils import is_hip
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
_is_hip = is_hip()
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
|
||||||
|
class EagerRunner(BaseRunner):
|
||||||
|
def __init__(self, model_runner: ModelRunner) -> None:
|
||||||
|
super().__init__(model_runner)
|
||||||
|
mr = model_runner
|
||||||
|
sa = mr.server_args
|
||||||
|
# Built first so the cg runners coalesce onto its buffers via the shared
|
||||||
|
# input pool; size to the largest tokens/req across modes the worker hits.
|
||||||
|
num_tokens_per_bs = 1
|
||||||
|
if mr.spec_algorithm.is_speculative():
|
||||||
|
# speculative_adaptive can grow draft tokens at runtime; size to the max.
|
||||||
|
num_draft_tokens = sa.max_speculative_num_draft_tokens or 1
|
||||||
|
if mr.is_draft_worker:
|
||||||
|
num_tokens_per_bs = max(
|
||||||
|
sa.speculative_eagle_topk or 1,
|
||||||
|
num_draft_tokens,
|
||||||
|
(
|
||||||
|
2 * (sa.speculative_num_steps or 0)
|
||||||
|
if sa.enable_multi_layer_eagle
|
||||||
|
else 0
|
||||||
|
),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
num_tokens_per_bs = (
|
||||||
|
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
|
||||||
|
num_draft_tokens, mr.is_draft_worker
|
||||||
|
)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
dllm_config = DllmConfig.from_server_args(sa)
|
||||||
|
if dllm_config is not None:
|
||||||
|
# dLLM runs block_size tokens/request (DLLM_EXTEND).
|
||||||
|
num_tokens_per_bs = dllm_config.block_size
|
||||||
|
max_bs = mr.max_running_requests
|
||||||
|
if (
|
||||||
|
mr.is_draft_worker
|
||||||
|
and mr.spec_algorithm.is_frozen_kv_mtp()
|
||||||
|
and sa.speculative_eagle_topk > 1
|
||||||
|
):
|
||||||
|
# Frozen-KV MTP expands the draft batch by topk on the bs axis
|
||||||
|
# (expand_for_topk_draft) before the eager fallback.
|
||||||
|
max_bs *= sa.speculative_eagle_topk
|
||||||
|
prefill_ceiling = (
|
||||||
|
sa.chunked_prefill_size
|
||||||
|
if sa.chunked_prefill_size and sa.chunked_prefill_size > 0
|
||||||
|
else mr.max_total_num_tokens
|
||||||
|
)
|
||||||
|
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_bs)
|
||||||
|
is_encoder_decoder = mr.model_config.is_encoder_decoder
|
||||||
|
self._eager_registry = build_eager_registry(
|
||||||
|
device=mr.device,
|
||||||
|
max_bs=max_bs,
|
||||||
|
max_num_token=max_num_token,
|
||||||
|
cache_loc_dtype=torch.int64,
|
||||||
|
enable_mamba_track=(
|
||||||
|
sa.enable_mamba_extra_buffer() and mr.spec_algorithm.is_none()
|
||||||
|
),
|
||||||
|
is_encoder_decoder=is_encoder_decoder,
|
||||||
|
encoder_len_fill_value=(
|
||||||
|
getattr(mr.model_config.hf_config, "max_source_positions", 0)
|
||||||
|
if is_encoder_decoder
|
||||||
|
else 0
|
||||||
|
),
|
||||||
|
dp_size=sa.dp_size,
|
||||||
|
)
|
||||||
|
|
||||||
|
def can_run_graph(self, forward_batch: ForwardBatch) -> bool:
|
||||||
|
# Eager never runs a cuda graph; callers dispatch on isinstance(...,
|
||||||
|
# EagerRunner) and must not route an eager batch into a replay branch.
|
||||||
|
return False
|
||||||
|
|
||||||
|
def load_batch(
|
||||||
|
self, forward_batch: ForwardBatch, pp_proxy_tensors=None, **kwargs
|
||||||
|
) -> ForwardBatch:
|
||||||
|
"""Copy the live batch into the fixed-max eager static buffers (sliced to
|
||||||
|
this batch's shape) — the eager counterpart of the cuda-graph runners'
|
||||||
|
load_batch."""
|
||||||
|
if envs.SGLANG_EAGER_INPUT_NO_COPY.get():
|
||||||
|
return replace(forward_batch)
|
||||||
|
raw_bs = forward_batch.batch_size
|
||||||
|
raw_num_tokens = forward_batch.input_ids.shape[0]
|
||||||
|
registry = self._eager_registry
|
||||||
|
registry.fill_from(
|
||||||
|
forward_batch,
|
||||||
|
raw_bs=raw_bs,
|
||||||
|
padded_bs=raw_bs,
|
||||||
|
raw_num_tokens=raw_num_tokens,
|
||||||
|
padded_num_tokens=raw_num_tokens,
|
||||||
|
pp_proxy_tensors=pp_proxy_tensors,
|
||||||
|
)
|
||||||
|
return registry.extract_buffer(
|
||||||
|
padded_bs=raw_bs,
|
||||||
|
padded_num_tokens=raw_num_tokens,
|
||||||
|
forward_batch_template=forward_batch,
|
||||||
|
)
|
||||||
|
|
||||||
|
def execute(
|
||||||
|
self, forward_batch: ForwardBatch, pp_proxy_tensors=None, **kwargs
|
||||||
|
) -> Any:
|
||||||
|
mode = forward_batch.forward_mode
|
||||||
|
if mode.is_decode():
|
||||||
|
return self._execute_decode(forward_batch, pp_proxy_tensors)
|
||||||
|
if mode.is_idle():
|
||||||
|
return self._execute_idle(forward_batch, pp_proxy_tensors)
|
||||||
|
if mode.is_extend(include_draft_extend_v2=True):
|
||||||
|
return self._execute_extend(forward_batch, pp_proxy_tensors)
|
||||||
|
raise ValueError(f"Invalid forward mode for eager runner: {mode}")
|
||||||
|
|
||||||
|
def _resolve_decode_pdmux(
|
||||||
|
self,
|
||||||
|
) -> Tuple[Any, contextlib.AbstractContextManager]:
|
||||||
|
"""Resolve the (attn_backend, forward_context) the eager decode forward
|
||||||
|
runs under. PDmux selects a per-stream backend and publishes it via an
|
||||||
|
active ForwardContext; non-pdmux uses attn_backend + the ambient ctx."""
|
||||||
|
model_runner = self.model_runner
|
||||||
|
if model_runner.server_args.enable_pdmux:
|
||||||
|
return model_runner.decode_attn_backend, forward_context(
|
||||||
|
ForwardContext(attn_backend=model_runner.decode_attn_backend)
|
||||||
|
)
|
||||||
|
return model_runner.attn_backend, contextlib.nullcontext()
|
||||||
|
|
||||||
|
def _execute_decode(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
pp_proxy_tensors=None,
|
||||||
|
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||||
|
model_runner = self.model_runner
|
||||||
|
enable_pdmux = model_runner.server_args.enable_pdmux
|
||||||
|
attn_backend, pdmux_ctx = self._resolve_decode_pdmux()
|
||||||
|
if not enable_pdmux:
|
||||||
|
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||||
|
if forward_batch.needs_forward_metadata_init():
|
||||||
|
if hasattr(model_runner.model, "prepare_forward_batch"):
|
||||||
|
# Prepare model-specific attention metadata before planning,
|
||||||
|
# e.g. Moss-VL's prefill cross-attention custom mask.
|
||||||
|
model_runner.model.prepare_forward_batch(forward_batch)
|
||||||
|
attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
# FIXME: add pp_proxy_tensors arg to all models
|
||||||
|
kwargs = model_runner._pp_kwargs(pp_proxy_tensors)
|
||||||
|
|
||||||
|
ctx = (
|
||||||
|
model_runner.device_timer.wrap(metadata={"category": "decode"})
|
||||||
|
if model_runner.device_timer
|
||||||
|
else contextlib.nullcontext()
|
||||||
|
)
|
||||||
|
|
||||||
|
with ctx, pdmux_ctx:
|
||||||
|
return model_runner.model.forward(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_batch.positions,
|
||||||
|
forward_batch,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _execute_extend(
|
||||||
|
self,
|
||||||
|
forward_batch: ForwardBatch,
|
||||||
|
pp_proxy_tensors=None,
|
||||||
|
) -> Union[LogitsProcessorOutput, PPProxyTensors, EmbeddingPoolerOutput]:
|
||||||
|
model_runner = self.model_runner
|
||||||
|
kwargs = model_runner._extend_forward_kwargs(forward_batch, pp_proxy_tensors)
|
||||||
|
|
||||||
|
if not model_runner.server_args.enable_pdmux:
|
||||||
|
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||||
|
|
||||||
|
if forward_batch.needs_forward_metadata_init():
|
||||||
|
if hasattr(model_runner.model, "prepare_forward_batch"):
|
||||||
|
# Prepare model-specific attention metadata before planning,
|
||||||
|
# e.g. Moss-VL's prefill cross-attention custom mask.
|
||||||
|
model_runner.model.prepare_forward_batch(forward_batch)
|
||||||
|
model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
|
||||||
|
cp_v2_active = is_cp_v2_active(forward_batch)
|
||||||
|
forward_positions = forward_batch.positions
|
||||||
|
if cp_v2_active:
|
||||||
|
prepare_cp_forward(forward_batch)
|
||||||
|
complete_hidden_states = kwargs.get("input_embeds")
|
||||||
|
if complete_hidden_states is None:
|
||||||
|
embed_layer = model_runner.model.get_input_embeddings()
|
||||||
|
complete_hidden_states = embed_layer(forward_batch.input_ids)
|
||||||
|
sharded_hidden_states, sharded_positions = cp_split_before_forward(
|
||||||
|
complete_hidden_states,
|
||||||
|
forward_batch.positions,
|
||||||
|
forward_batch,
|
||||||
|
)
|
||||||
|
kwargs["input_embeds"] = sharded_hidden_states
|
||||||
|
forward_positions = sharded_positions
|
||||||
|
|
||||||
|
ctx = (
|
||||||
|
model_runner.device_timer.wrap(metadata={"category": "extend"})
|
||||||
|
if model_runner.device_timer
|
||||||
|
else contextlib.nullcontext()
|
||||||
|
)
|
||||||
|
with ctx:
|
||||||
|
pcg_runner = model_runner.prefill_cuda_graph_runner
|
||||||
|
if (
|
||||||
|
_is_hip
|
||||||
|
and pcg_runner is not None
|
||||||
|
and not isinstance(pcg_runner, EagerRunner)
|
||||||
|
and not cp_v2_active
|
||||||
|
):
|
||||||
|
# HIP PCG eager fallback: enter the PCG context so Dynamo guards
|
||||||
|
# and PCG-specific MoE/attention paths stay consistent.
|
||||||
|
with (
|
||||||
|
enable_tc_piecewise_cuda_graph(),
|
||||||
|
set_tc_piecewise_forward_context(
|
||||||
|
forward_batch,
|
||||||
|
model_runner.attention_layers,
|
||||||
|
getattr(model_runner.model, "quant_config", None),
|
||||||
|
model_runner.moe_layers,
|
||||||
|
model_runner.moe_fusions,
|
||||||
|
dsa_indexers=model_runner.dsa_indexers,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
ret = model_runner.model.forward(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_positions,
|
||||||
|
forward_batch,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
elif cp_v2_active:
|
||||||
|
# CP-V2: drive .model directly to gather across CP ranks before logits.
|
||||||
|
hidden_states = model_runner.model.model(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_positions,
|
||||||
|
forward_batch,
|
||||||
|
input_embeds=kwargs.get("input_embeds"),
|
||||||
|
pp_proxy_tensors=kwargs.get("pp_proxy_tensors"),
|
||||||
|
)
|
||||||
|
aux_hidden_states = None
|
||||||
|
capture_aux_hidden_states = getattr(
|
||||||
|
model_runner.model, "capture_aux_hidden_states", False
|
||||||
|
)
|
||||||
|
if capture_aux_hidden_states:
|
||||||
|
hidden_states, aux_hidden_states = hidden_states
|
||||||
|
if model_runner.model.pp_group.is_last_rank:
|
||||||
|
hidden_states = cp_gather_after_forward(
|
||||||
|
hidden_states,
|
||||||
|
forward_batch,
|
||||||
|
torch.cuda.current_stream(),
|
||||||
|
)
|
||||||
|
ret = model_runner.model.logits_processor(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
hidden_states,
|
||||||
|
model_runner.model.lm_head,
|
||||||
|
forward_batch,
|
||||||
|
aux_hidden_states,
|
||||||
|
)
|
||||||
|
elif capture_aux_hidden_states:
|
||||||
|
ret = hidden_states, aux_hidden_states
|
||||||
|
else:
|
||||||
|
ret = hidden_states
|
||||||
|
else:
|
||||||
|
ret = model_runner.model.forward(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_positions,
|
||||||
|
forward_batch,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
return ret
|
||||||
|
|
||||||
|
def _execute_idle(
|
||||||
|
self, forward_batch: ForwardBatch, pp_proxy_tensors=None
|
||||||
|
) -> Union[LogitsProcessorOutput, PPProxyTensors]:
|
||||||
|
model_runner = self.model_runner
|
||||||
|
# Padded idle (DP-attn MLP sync) needs metadata reinit; unpadded must
|
||||||
|
# drop stale forward_metadata to avoid an SWA use-after-free on req_pool.
|
||||||
|
if forward_batch.batch_size > 0:
|
||||||
|
if not model_runner.server_args.enable_pdmux:
|
||||||
|
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
|
||||||
|
model_runner.attn_backend.init_forward_metadata(forward_batch)
|
||||||
|
else:
|
||||||
|
model_runner.attn_backend.forward_metadata = None
|
||||||
|
|
||||||
|
kwargs = model_runner._pp_kwargs(pp_proxy_tensors)
|
||||||
|
ctx = (
|
||||||
|
model_runner.device_timer.wrap(metadata={"category": "idle"})
|
||||||
|
if model_runner.device_timer
|
||||||
|
else contextlib.nullcontext()
|
||||||
|
)
|
||||||
|
with ctx:
|
||||||
|
return model_runner.model.forward(
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_batch.positions,
|
||||||
|
forward_batch,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
Reference in New Issue
Block a user