diff --git a/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py b/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py index 513b66873..eaaf8b858 100644 --- a/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py +++ b/python/sglang/srt/model_executor/cuda_graph_buffer_registry.py @@ -887,3 +887,48 @@ def build_prefill_registry( ) reg.register_slot(slot, bind=bind) 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, + ) diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index aca970269..020d71b35 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -26,7 +26,7 @@ import socket import threading import time from collections import defaultdict -from dataclasses import dataclass, replace +from dataclasses import dataclass from pathlib import Path 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.tbo_backend import TboAttnBackend from sglang.srt.layers.cp.utils import ( - cp_gather_after_forward, - cp_split_before_forward, get_cp_strategy, - is_cp_v2_active, - prepare_cp_forward, ) from sglang.srt.layers.dp_attention import ( DpPaddingMode, @@ -143,7 +139,6 @@ from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.layers.moe.hash_topk import HashTopK 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.sampler import create_sampler 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.memory_pool import ReqToTokenPool 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 ( Backend, 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.runner import ( + EagerRunner, PrefillCudaGraphRunner, ) from sglang.srt.model_executor.runner.decode_cuda_graph_runner import ( _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.remote_instance_weight_loader_utils import ( RemoteInstanceWeightLoaderBackend, @@ -244,7 +231,7 @@ from sglang.srt.utils import ( set_cuda_arch, 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.nvtx_pytorch_hooks import PytHooks from sglang.srt.utils.nvtx_utils import profile_range @@ -369,14 +356,6 @@ class ModelRunnerOutput: 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): """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.forward_pass_id = 0 self.init_new_workspace = False - self._eager_decode_registry = _EagerBufferRegistry() - self._eager_prefill_registry = _EagerBufferRegistry() self.draft_model_idx = draft_model_idx self.enable_hisparse = server_args.enable_hisparse @@ -912,6 +889,11 @@ class ModelRunner(ModelRunnerKVCacheMixin): # runs with aux hidden state capture enabled. 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": self.init_cublas() self.init_attention_backend() @@ -951,7 +933,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): self.init_attention_backend() 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 if server_args.forward_hooks: @@ -3040,6 +3025,11 @@ class ModelRunner(ModelRunnerKVCacheMixin): "resolved prefill.backend='disabled' (e.g. via " "--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 # 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): 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: """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 runner's capture/replay, so this is skipped there. """ + # For MLP sync if forward_batch.global_num_tokens_cpu is not None: forward_batch.prepare_mlp_sync_batch(self) else: @@ -3343,6 +3229,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): server_args=self.server_args, ) + # Hisparse coordinator — backends now read it from self.model_runner. if self.hisparse_coordinator is not None: 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 {} - def forward_decode( - self, - forward_batch: ForwardBatch, - pp_proxy_tensors=None, - ) -> Union[LogitsProcessorOutput, PPProxyTensors]: - if not self.server_args.enable_pdmux: - 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 + def _extend_forward_kwargs( + self, forward_batch: ForwardBatch, pp_proxy_tensors + ) -> dict: + """Build the extend/prefill model.forward kwargs (pp_proxy_tensors + + input_embeds / replace_embeds overrides + get_embedding), shared by the + prefill cuda-graph path and the EagerRunner's eager extend path.""" kwargs = self._pp_kwargs(pp_proxy_tensors) if forward_batch.input_embeds is not None: kwargs["input_embeds"] = forward_batch.input_embeds.bfloat16() @@ -3426,159 +3263,7 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) if not self.is_generation: kwargs["get_embedding"] = True - - # 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, - ) + return kwargs def forward_split_prefill( self, @@ -3737,34 +3422,48 @@ class ModelRunner(ModelRunnerKVCacheMixin): ) return ModelRunnerOutput(logits_output=ret, can_run_graph=can_run_graph) - # DP / MLP-sync padding + attn-tp normalization that the eager - # (non-graph) forward needs. The graph path skips it: capture/replay - # pads inside the runner. + # DP / MLP-sync padding + attn-tp normalization. Only the decode + # cuda-graph path above pre-pads its static buffers and returns + # 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) - # Forward without cuda graph - if forward_batch.forward_mode.is_decode(): - ret = self.forward_decode( - forward_batch, - pp_proxy_tensors=pp_proxy_tensors, - ) - elif forward_batch.forward_mode.is_split_prefill(): + if forward_batch.forward_mode.is_split_prefill(): + # Layer-split mode; stays on ModelRunner, not the eager runner. ret = self.forward_split_prefill( forward_batch, reinit_attn_backend=reinit_attn_backend, forward_count=split_forward_count, ) - elif forward_batch.forward_mode.is_extend(include_draft_extend_v2=True): - ret, can_run_graph = self.forward_extend( - forward_batch, - pp_proxy_tensors=pp_proxy_tensors, + elif ( + forward_batch.forward_mode.is_extend(include_draft_extend_v2=True) + and not isinstance(self.prefill_cuda_graph_runner, EagerRunner) + 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(): - ret = self.forward_idle( + with ctx: + 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 ) - else: - raise ValueError(f"Invalid forward mode: {forward_batch.forward_mode}") if ( forward_batch.global_num_tokens_cpu is not None diff --git a/python/sglang/srt/model_executor/runner/__init__.py b/python/sglang/srt/model_executor/runner/__init__.py index 6619b2657..0422f4457 100644 --- a/python/sglang/srt/model_executor/runner/__init__.py +++ b/python/sglang/srt/model_executor/runner/__init__.py @@ -13,6 +13,9 @@ Public API: capture-loop scaffolding on top of BaseRunner. - DecodeCudaGraphRunner — concrete decode-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, and the DeepEP adapter live in 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 ( 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 PrefillCudaGraphRunner, ) diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 163c04805..da4b877cb 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -779,6 +779,15 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): if self.enable_profile_cuda_graph: 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. # Capture the large shapes first so that the smaller shapes # can reuse the memory pool allocated for the large shapes. diff --git a/python/sglang/srt/model_executor/runner/eager_runner.py b/python/sglang/srt/model_executor/runner/eager_runner.py new file mode 100644 index 000000000..4755f9ad3 --- /dev/null +++ b/python/sglang/srt/model_executor/runner/eager_runner.py @@ -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, + )