refactor(runner): move kernel warmup into the shared runner lifecycle (warmup()) (#28739)
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
d705a91de1
commit
856b0dc74b
@@ -416,12 +416,17 @@ def _deep_gemm_execution_hook(
|
||||
yield
|
||||
|
||||
|
||||
def pp_parallel_deep_gemm_warmup(model_runner) -> None:
|
||||
def pp_parallel_deep_gemm_warmup(runner) -> None:
|
||||
"""Run per-PP-rank dummy DECODE+EXTEND forwards so each rank's
|
||||
DeepGEMM JIT compiles in parallel instead of serially via the warmup
|
||||
/generate flowing through the pipeline. Opt-in via
|
||||
SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.
|
||||
|
||||
Driven from BaseRunner.warmup(), which passes the runner; the dummy
|
||||
forwards go through runner._dummy_run (the autotune/dummy-run machinery now
|
||||
lives on BaseRunner). ModelRunner state is read via runner.model_runner.
|
||||
"""
|
||||
model_runner = runner.model_runner
|
||||
# n_splits ~= n_sms / ceil(bs/block_m) with block_m=64; sweep 5 bs to
|
||||
# cover the brackets real /generate hits (smallest decode shape,
|
||||
# mid-low, two mid, and n_splits=1 for ~5K+ token prefill). Ceil-align
|
||||
@@ -466,11 +471,11 @@ def pp_parallel_deep_gemm_warmup(model_runner) -> None:
|
||||
with torch.inference_mode():
|
||||
for bs in batch_sizes:
|
||||
if run_decode:
|
||||
model_runner._dummy_run(
|
||||
runner._dummy_run(
|
||||
batch_size=bs, forward_mode_override=ForwardMode.DECODE
|
||||
)
|
||||
if run_extend:
|
||||
model_runner._dummy_run(
|
||||
runner._dummy_run(
|
||||
batch_size=bs, forward_mode_override=ForwardMode.EXTEND
|
||||
)
|
||||
|
||||
|
||||
@@ -18,7 +18,6 @@ from __future__ import annotations
|
||||
import contextlib
|
||||
import datetime
|
||||
import gc
|
||||
import hashlib
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
@@ -27,7 +26,6 @@ import threading
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
@@ -35,7 +33,6 @@ import torch.distributed as dist
|
||||
from torch import nn
|
||||
|
||||
from sglang.jit_kernel.ngram_embedding import update_token_table_decode
|
||||
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
|
||||
from sglang.srt.configs import (
|
||||
BailingHybridConfig,
|
||||
FalconH1Config,
|
||||
@@ -129,12 +126,8 @@ from sglang.srt.layers.cp.utils import (
|
||||
get_cp_strategy,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
get_attention_tp_group,
|
||||
get_attention_tp_size,
|
||||
initialize_dp_attention,
|
||||
set_dp_buffer_len,
|
||||
set_is_extend_in_batch,
|
||||
)
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.moe.hash_topk import HashTopK
|
||||
@@ -156,7 +149,6 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
cuda_graph_fully_disabled,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
PPProxyTensors,
|
||||
@@ -175,9 +167,6 @@ 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_loader.loader import DefaultModelLoader, get_model_loader
|
||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
RemoteInstanceWeightLoaderBackend,
|
||||
@@ -193,10 +182,7 @@ from sglang.srt.server_args import (
|
||||
get_global_server_args,
|
||||
set_global_server_args_for_scheduler,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import (
|
||||
SpeculativeAlgorithm,
|
||||
create_dummy_verify_input,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.state_capturer.base import TopkCaptureOutput
|
||||
from sglang.srt.state_capturer.indexer_topk import (
|
||||
create_indexer_capturer,
|
||||
@@ -213,7 +199,6 @@ from sglang.srt.utils import (
|
||||
broadcast_pyobj,
|
||||
cpu_has_amx_support,
|
||||
dynamic_import,
|
||||
empty_context,
|
||||
enable_show_time_cost,
|
||||
get_available_gpu_memory,
|
||||
get_bool_env_var,
|
||||
@@ -224,14 +209,11 @@ from sglang.srt.utils import (
|
||||
is_npu,
|
||||
log_info_on_rank0,
|
||||
monkey_patch_p2p_access_check,
|
||||
require_attn_tp_gather,
|
||||
require_gathered_buffer,
|
||||
require_mlp_tp_gather,
|
||||
reserve_rope_cache_for_long_sequences,
|
||||
set_cuda_arch,
|
||||
slow_rank_detector,
|
||||
)
|
||||
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
|
||||
@@ -889,22 +871,14 @@ 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)
|
||||
|
||||
# Device-specific attention-backend init. cg_supported gates decode
|
||||
# cuda-graph capture (out-of-tree / unknown platforms may not support it).
|
||||
cg_supported = True
|
||||
if self.device == "cuda" or self.device == "musa":
|
||||
self.init_cublas()
|
||||
self.init_attention_backend()
|
||||
self.kernel_warmup()
|
||||
self._pre_initialize_flashinfer_allreduce_workspace()
|
||||
if not disable_cuda_graph:
|
||||
self.init_decode_cuda_graph()
|
||||
elif self.device == "cpu":
|
||||
self.init_attention_backend()
|
||||
if not disable_cuda_graph:
|
||||
self.init_decode_cuda_graph()
|
||||
elif self.device == "npu":
|
||||
self.init_attention_backend()
|
||||
# lazy init for zbal with mix mode (before graph capture when enable_cuda_graph)
|
||||
@@ -918,32 +892,44 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
get_world_group().world_size,
|
||||
get_world_group().cpu_group,
|
||||
)
|
||||
if not disable_cuda_graph:
|
||||
self.init_decode_cuda_graph()
|
||||
elif current_platform.is_out_of_tree():
|
||||
self.init_attention_backend()
|
||||
if current_platform.support_cuda_graph() and not disable_cuda_graph:
|
||||
self.init_decode_cuda_graph()
|
||||
else:
|
||||
self.decode_cuda_graph_runner = None
|
||||
self.graph_mem_usage = 0
|
||||
cg_supported = current_platform.support_cuda_graph()
|
||||
else:
|
||||
self.init_attention_backend()
|
||||
cg_supported = False
|
||||
|
||||
# The eager (no-cuda-graph) phase runner, built AFTER the attention
|
||||
# backend so its __init__ can warm up kernels (run-once) and allocate the
|
||||
# fixed-max static buffer — both before the cuda-graph runners, so that
|
||||
# buffer is canonical in the shared pool and the cg runners coalesce onto
|
||||
# it. Always built: it serves both the fully-disabled case (decode/prefill
|
||||
# runners point at it) and the eager fallback when a cg runner can't run a
|
||||
# batch.
|
||||
self.eager_runner = EagerRunner(self)
|
||||
|
||||
# cuda-graph capture: prefill before decode, so both coalesce onto the
|
||||
# eager buffer allocated above. (init_prefill_cuda_graph routes prefill
|
||||
# to the eager runner when the prefill graph is disabled.)
|
||||
self.init_prefill_cuda_graph()
|
||||
if not disable_cuda_graph and cg_supported:
|
||||
self.init_decode_cuda_graph()
|
||||
else:
|
||||
self.decode_cuda_graph_runner = None
|
||||
self.graph_mem_usage = 0
|
||||
self.init_attention_backend()
|
||||
|
||||
if disable_cuda_graph:
|
||||
# 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
|
||||
|
||||
# Register forward hooks AFTER cuda-graph capture so their tensor ops are
|
||||
# not traced into any captured graph — capture stays hook-free and hooks
|
||||
# fire only on the eager forward path (capture replay never runs Python
|
||||
# hooks anyway).
|
||||
if server_args.forward_hooks:
|
||||
register_forward_hooks(self.model, server_args.forward_hooks)
|
||||
|
||||
self.init_prefill_cuda_graph()
|
||||
|
||||
self.prealloc_symmetric_memory_pool()
|
||||
|
||||
if self.canary_manager is not None and not self.is_draft_worker:
|
||||
@@ -2490,431 +2476,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
full_attention_backend = ATTENTION_BACKENDS[backend_str](self)
|
||||
return attn_backend_wrapper(self, full_attention_backend)
|
||||
|
||||
def kernel_warmup(self):
|
||||
"""Warmup and tune kernels before cuda graph capture."""
|
||||
if self.device != "cuda":
|
||||
return
|
||||
|
||||
if self._should_run_flashinfer_autotune():
|
||||
self._flashinfer_autotune()
|
||||
|
||||
if (
|
||||
envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get()
|
||||
and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
|
||||
and self.pp_size > 1
|
||||
and not self.spec_algorithm.is_speculative()
|
||||
):
|
||||
from sglang.srt.layers.deep_gemm_wrapper.compile_utils import (
|
||||
pp_parallel_deep_gemm_warmup,
|
||||
)
|
||||
|
||||
pp_parallel_deep_gemm_warmup(self)
|
||||
|
||||
def _pre_initialize_flashinfer_allreduce_workspace(self):
|
||||
"""Pre-initialize flashinfer allreduce fusion workspaces.
|
||||
|
||||
Must run before CUDA graph capture to avoid collective operations
|
||||
(broadcasts, barriers) inside the graph capture context, which can
|
||||
deadlock with custom_all_reduce.register_graph_buffers.
|
||||
"""
|
||||
if self.server_args.flashinfer_allreduce_fusion_backend is None:
|
||||
return
|
||||
|
||||
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
|
||||
from sglang.srt.layers.flashinfer_comm_fusion import pre_initialize_workspaces
|
||||
|
||||
pre_initialize_workspaces(
|
||||
max_token_num=FUSE_ALLREDUCE_MAX_BATCH_SIZE,
|
||||
hidden_dim=self.model_config.hidden_size,
|
||||
dtype=self.dtype,
|
||||
)
|
||||
|
||||
def _should_run_flashinfer_autotune(self) -> bool:
|
||||
"""Check if flashinfer autotune should be run."""
|
||||
if self.server_args.disable_flashinfer_autotune:
|
||||
return False
|
||||
|
||||
# CuteDSL v1 (cutedsl runner + deepep a2a) bypasses MoeRunner and must not
|
||||
# be autotuned -- its _dummy_run would dispatch more tokens per rank than
|
||||
# SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK, tripping a DeepEP assert.
|
||||
# Read server_args directly to avoid depending on initialize_moe_config()
|
||||
# having already populated the MoE backend globals.
|
||||
if (
|
||||
self.server_args.moe_runner_backend == "flashinfer_cutedsl"
|
||||
and self.server_args.moe_a2a_backend == "deepep"
|
||||
):
|
||||
return False
|
||||
|
||||
backend_str = self.server_args.moe_runner_backend
|
||||
|
||||
# TODO smor- support other cases for flashinfer autotune, such as, mamba backend
|
||||
|
||||
moe_needs_autotune = backend_str in [
|
||||
"flashinfer_trtllm",
|
||||
"flashinfer_trtllm_routed",
|
||||
"flashinfer_mxfp4",
|
||||
"flashinfer_cutedsl",
|
||||
"flashinfer_cutlass",
|
||||
]
|
||||
|
||||
from sglang.srt.layers.quantization.fp4_utils import (
|
||||
get_fp4_gemm_runner_backend,
|
||||
)
|
||||
|
||||
model_uses_fp4 = self.model_config.quantization in (
|
||||
"modelopt_fp4",
|
||||
"modelopt_mixed",
|
||||
)
|
||||
fp4_gemm_needs_autotune = model_uses_fp4 and (
|
||||
get_fp4_gemm_runner_backend().is_flashinfer_cutlass()
|
||||
or get_fp4_gemm_runner_backend().is_flashinfer_cutedsl()
|
||||
)
|
||||
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
get_fp8_gemm_runner_backend,
|
||||
)
|
||||
from sglang.srt.utils import is_sm100_supported
|
||||
|
||||
model_uses_modelopt_fp8 = self.model_config.quantization in (
|
||||
"modelopt",
|
||||
"modelopt_fp8",
|
||||
"modelopt_mixed",
|
||||
)
|
||||
fp8_gemm_needs_autotune = (
|
||||
get_fp8_gemm_runner_backend().is_flashinfer_cutlass()
|
||||
or (model_uses_modelopt_fp8 and is_sm100_supported())
|
||||
)
|
||||
|
||||
if not (
|
||||
moe_needs_autotune or fp4_gemm_needs_autotune or fp8_gemm_needs_autotune
|
||||
):
|
||||
return False
|
||||
|
||||
major, _ = torch.cuda.get_device_capability()
|
||||
if major < 9:
|
||||
return False
|
||||
|
||||
if self.spec_algorithm.is_speculative():
|
||||
return not self.is_draft_worker
|
||||
|
||||
return True
|
||||
|
||||
def _flashinfer_autotune(self):
|
||||
"""Run flashinfer autotune."""
|
||||
from flashinfer.autotuner import autotune
|
||||
|
||||
from sglang.srt.layers.logits_processor import autotune_dummy_run_mode
|
||||
|
||||
cache_path = self._flashinfer_autotune_cache_path()
|
||||
if envs.SGLANG_FLASHINFER_AUTOTUNE_CACHE.get():
|
||||
autotune_cache = cache_path
|
||||
logger.info("Running FlashInfer autotune with cache: %s", autotune_cache)
|
||||
else:
|
||||
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
runs_dir = cache_path.parent / "runs"
|
||||
runs_dir.mkdir(parents=True, exist_ok=True)
|
||||
autotune_cache = (
|
||||
runs_dir / f"{cache_path.stem}.{timestamp}{cache_path.suffix}"
|
||||
)
|
||||
logger.info(
|
||||
"Running FlashInfer autotune (cache reuse DISABLED via "
|
||||
"SGLANG_FLASHINFER_AUTOTUNE_CACHE=0); writing fresh result to: %s",
|
||||
autotune_cache,
|
||||
)
|
||||
|
||||
# Run warmup on the non-default stream to avoid NCCL 2.29+ cudaMemcpyBatchAsync
|
||||
# calls on default stream (unsupported by CUDA) when --enable-symm-mem is used.
|
||||
self.forward_stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.get_device_module(self.device).stream(self.forward_stream):
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
autotune(True, cache=str(autotune_cache)),
|
||||
autotune_dummy_run_mode(),
|
||||
):
|
||||
self._dummy_run(batch_size=self.req_to_token_pool.size)
|
||||
torch.cuda.current_stream().wait_stream(self.forward_stream)
|
||||
logger.info("FlashInfer autotune completed.")
|
||||
|
||||
def _flashinfer_autotune_cache_path(self) -> Path:
|
||||
import flashinfer
|
||||
|
||||
major, minor = torch.cuda.get_device_capability(self.device)
|
||||
arch = f"sm{major}{minor}"
|
||||
flashinfer_version = getattr(flashinfer, "__version__", "unknown")
|
||||
|
||||
server_args = self.server_args
|
||||
model_key = "|".join(
|
||||
[
|
||||
str(server_args.model_path),
|
||||
str(self.dtype),
|
||||
str(server_args.quantization),
|
||||
str(server_args.moe_runner_backend),
|
||||
str(self.tp_size),
|
||||
str(self.pp_size),
|
||||
str(self.dp_size),
|
||||
str(self.moe_ep_size),
|
||||
str(self.model_config.hf_config.__class__.__name__),
|
||||
]
|
||||
)
|
||||
cache_key = hashlib.sha256(model_key.encode()).hexdigest()[:16]
|
||||
cache_dir = (
|
||||
Path(envs.SGLANG_CACHE_DIR.get())
|
||||
/ "flashinfer"
|
||||
/ "autotune"
|
||||
/ flashinfer_version
|
||||
/ arch
|
||||
/ cache_key
|
||||
)
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
return (
|
||||
cache_dir
|
||||
/ f"rank_tp{self.tp_rank}_pp{self.pp_rank}_dp{self.dp_rank or 0}.json"
|
||||
)
|
||||
|
||||
def _dummy_run(
|
||||
self,
|
||||
batch_size: int,
|
||||
run_ctx=None,
|
||||
forward_mode_override: Optional[ForwardMode] = None,
|
||||
):
|
||||
"""Run a dummy forward pass for warmup/profiling.
|
||||
|
||||
forward_mode_override forces EXTEND/DECODE regardless of
|
||||
is_generation (used by the PP-parallel DeepGEMM warmup).
|
||||
"""
|
||||
if forward_mode_override is not None:
|
||||
capture_forward_mode = forward_mode_override
|
||||
elif self.is_generation:
|
||||
capture_forward_mode = ForwardMode.DECODE
|
||||
else:
|
||||
capture_forward_mode = ForwardMode.EXTEND
|
||||
capture_hidden_mode = CaptureHiddenMode.NULL
|
||||
num_tokens_per_bs = 1
|
||||
if self.spec_algorithm.is_speculative():
|
||||
if self.is_draft_worker:
|
||||
if not self.spec_algorithm.supports_target_verify_for_draft():
|
||||
raise RuntimeError("This should not happen")
|
||||
capture_forward_mode = ForwardMode.TARGET_VERIFY
|
||||
num_tokens_per_bs = (
|
||||
self.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
|
||||
self.server_args.speculative_num_draft_tokens, self.is_draft_worker
|
||||
)
|
||||
)
|
||||
|
||||
if self.server_args.enable_return_hidden_states:
|
||||
capture_hidden_mode = CaptureHiddenMode.FULL
|
||||
|
||||
num_tokens = batch_size * num_tokens_per_bs
|
||||
|
||||
# Keep warmup aligned with scheduler MLP-sync padding.
|
||||
if require_mlp_sync(self.server_args):
|
||||
attn_tp_size = get_attention_tp_size()
|
||||
if attn_tp_size > 1 and num_tokens % attn_tp_size != 0:
|
||||
num_tokens = ceil_align(num_tokens, attn_tp_size)
|
||||
batch_size = num_tokens // num_tokens_per_bs
|
||||
|
||||
seq_len_fill_value = self.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
|
||||
if self.server_args.enable_torch_compile:
|
||||
set_torch_compile_config()
|
||||
should_disable_torch_compile = not getattr(
|
||||
self.model, "_can_torch_compile", True
|
||||
)
|
||||
if should_disable_torch_compile:
|
||||
log_info_on_rank0(
|
||||
logger,
|
||||
"Transformers backend model reports it is not torch.compile "
|
||||
"compatible (e.g. dynamic rope scaling). Disabling torch.compile.",
|
||||
)
|
||||
self.server_args.enable_torch_compile = False
|
||||
|
||||
# NOTE: aux hidden state capture (eagle3/dflash) is already
|
||||
# configured by init_aux_hidden_state_capture() in initialize().
|
||||
|
||||
require_mlp_tp_gather_ = require_mlp_tp_gather(self.server_args)
|
||||
if require_gathered_buffer(self.server_args):
|
||||
assert require_mlp_tp_gather_ or require_attn_tp_gather(self.server_args)
|
||||
|
||||
buffers = _allocate_decode_buffers(
|
||||
device=self.device,
|
||||
max_bs=batch_size,
|
||||
max_num_token=num_tokens,
|
||||
hidden_size=self.model_config.hidden_size,
|
||||
vocab_size=self.model_config.vocab_size,
|
||||
dtype=self.model_config.dtype,
|
||||
dp_size=self.server_args.dp_size,
|
||||
pp_size=self.server_args.pp_size,
|
||||
is_encoder_decoder=self.model_config.is_encoder_decoder,
|
||||
require_mlp_tp_gather=require_mlp_tp_gather_,
|
||||
seq_len_fill_value=seq_len_fill_value,
|
||||
encoder_len_fill_value=(
|
||||
getattr(self.model_config.hf_config, "max_source_positions", 0)
|
||||
if self.model_config.is_encoder_decoder
|
||||
else 0
|
||||
),
|
||||
num_tokens_per_bs=num_tokens_per_bs,
|
||||
cache_loc_dtype=torch.int64,
|
||||
enable_mamba_track=False,
|
||||
hc_hidden_size=getattr(self.model_config, "hc_hidden_size", None),
|
||||
pp_proxy_topk_size=self.get_pp_proxy_topk_size(),
|
||||
)
|
||||
buffers.num_token_non_padded[...] = num_tokens
|
||||
|
||||
# For extend mode
|
||||
if capture_forward_mode == ForwardMode.EXTEND:
|
||||
extend_prefix_lens_cpu = [0] * batch_size
|
||||
extend_seq_lens_cpu = [seq_len_fill_value] * batch_size
|
||||
extend_num_tokens = num_tokens
|
||||
extend_seq_lens = torch.full(
|
||||
(batch_size,), seq_len_fill_value, dtype=torch.int32, device=self.device
|
||||
)
|
||||
extend_prefix_lens = torch.zeros(
|
||||
(batch_size,), dtype=torch.int32, device=self.device
|
||||
)
|
||||
extend_start_loc = torch.arange(
|
||||
0, num_tokens, num_tokens_per_bs, dtype=torch.int32, device=self.device
|
||||
)
|
||||
else:
|
||||
extend_prefix_lens_cpu = None
|
||||
extend_seq_lens_cpu = None
|
||||
extend_num_tokens = None
|
||||
extend_seq_lens = None
|
||||
extend_prefix_lens = None
|
||||
extend_start_loc = None
|
||||
|
||||
if self.server_args.pp_size > 1:
|
||||
# PP0 already cp-split hidden_states before send.
|
||||
pp_hidden_tokens = num_tokens
|
||||
if (
|
||||
capture_forward_mode == ForwardMode.EXTEND
|
||||
and self.pp_rank != 0
|
||||
and self.attn_cp_size > 1
|
||||
):
|
||||
pp_hidden_tokens = num_tokens // self.attn_cp_size
|
||||
pp_proxy_tensors = PPProxyTensors(
|
||||
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
|
||||
)
|
||||
|
||||
if require_mlp_tp_gather_:
|
||||
global_num_tokens_cpu = [num_tokens] * self.server_args.dp_size
|
||||
elif require_attn_tp_gather(self.server_args):
|
||||
global_num_tokens_cpu = [num_tokens]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
num_tokens_tensor = torch.tensor(
|
||||
global_num_tokens_cpu, dtype=torch.int32, device=self.device
|
||||
)
|
||||
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
spec_info = create_dummy_verify_input(
|
||||
self.spec_algorithm,
|
||||
self.server_args,
|
||||
buffers.custom_mask,
|
||||
num_tokens_per_bs,
|
||||
self.is_draft_worker,
|
||||
)
|
||||
if spec_info is not None and (
|
||||
self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone()
|
||||
):
|
||||
# MTP models (e.g. deepseek_nextn) read spec_info.hidden_states
|
||||
# during forward; provide a dummy so warmup doesn't crash.
|
||||
spec_info.hidden_states = torch.zeros(
|
||||
(num_tokens, self.model_config.hidden_size),
|
||||
dtype=self.dtype,
|
||||
device=self.device,
|
||||
)
|
||||
if capture_hidden_mode != CaptureHiddenMode.FULL:
|
||||
capture_hidden_mode = (
|
||||
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
|
||||
)
|
||||
|
||||
if self.server_args.enable_lora:
|
||||
lora_ids = [None] * batch_size
|
||||
else:
|
||||
lora_ids = None
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
forward_mode=capture_forward_mode,
|
||||
batch_size=batch_size,
|
||||
input_ids=buffers.input_ids,
|
||||
req_pool_indices=buffers.req_pool_indices,
|
||||
seq_lens=buffers.seq_lens,
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
next_token_logits_buffer=buffers.next_token_logits_buffer,
|
||||
orig_seq_lens=buffers.seq_lens,
|
||||
out_cache_loc=buffers.out_cache_loc,
|
||||
seq_lens_sum=buffers.seq_lens.sum().item(),
|
||||
encoder_lens=buffers.encoder_lens,
|
||||
return_logprob=False,
|
||||
positions=buffers.positions,
|
||||
extend_num_tokens=extend_num_tokens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
extend_prefix_lens=extend_prefix_lens,
|
||||
extend_start_loc=extend_start_loc,
|
||||
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
|
||||
extend_seq_lens_cpu=extend_seq_lens_cpu,
|
||||
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
|
||||
global_num_tokens_cpu=global_num_tokens_cpu,
|
||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||
global_dp_buffer_len=global_dp_buffer_len,
|
||||
mrope_positions=buffers.mrope_positions,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
spec_info=spec_info,
|
||||
capture_hidden_mode=capture_hidden_mode,
|
||||
num_token_non_padded=buffers.num_token_non_padded,
|
||||
global_forward_mode=capture_forward_mode,
|
||||
lora_ids=lora_ids,
|
||||
)
|
||||
|
||||
if lora_ids is not None:
|
||||
self.lora_manager.prepare_lora_batch(forward_batch)
|
||||
|
||||
self.attn_backend.init_forward_metadata(forward_batch)
|
||||
|
||||
def run_once():
|
||||
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
|
||||
set_dp_buffer_len(
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
kwargs = {}
|
||||
if (
|
||||
self.server_args.pp_size > 1
|
||||
and "pp_proxy_tensors"
|
||||
in inspect.signature(self.model.forward).parameters
|
||||
):
|
||||
kwargs["pp_proxy_tensors"] = PPProxyTensors(
|
||||
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
|
||||
)
|
||||
if not self.is_generation:
|
||||
kwargs["get_embedding"] = True
|
||||
|
||||
logits_output_or_pp_proxy_tensors = self.model.forward(
|
||||
buffers.input_ids,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
)
|
||||
return logits_output_or_pp_proxy_tensors
|
||||
|
||||
torch.get_device_module(self.device).synchronize()
|
||||
self.tp_group.barrier()
|
||||
with forward_context(ForwardContext(attn_backend=self.attn_backend)):
|
||||
with torch.inference_mode(), run_ctx or empty_context():
|
||||
run_once()
|
||||
|
||||
def maybe_init_ngram_embedding(self):
|
||||
self.use_ngram_embedding = self.model_config.use_ngram_embedding
|
||||
if self.use_ngram_embedding:
|
||||
@@ -3038,6 +2599,29 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
||||
if self.is_draft_worker and not force_for_draft_worker:
|
||||
return
|
||||
|
||||
# EAGLE-family target worker: the prefill graph captures
|
||||
# CaptureHiddenMode.NULL, but target prefill needs FULL hidden states to
|
||||
# feed the draft, so can_run_graph always rejects it and prefill runs
|
||||
# eagerly. The graph is therefore never used; capturing it (now before
|
||||
# the decode graph) can perturb backend state on FP4 / TRTLLM-MoE paths
|
||||
# and corrupt decode replay, so skip its capture and route prefill
|
||||
# through the eager runner — the runtime path either way. With
|
||||
# enable_return_hidden_states the prefill graph is FULL and usable, so
|
||||
# only skip when it would capture NULL.
|
||||
if (
|
||||
self.spec_algorithm.is_eagle()
|
||||
and not self.is_draft_worker
|
||||
and not self.server_args.enable_return_hidden_states
|
||||
):
|
||||
logger.info(
|
||||
"Disable prefill CUDA graph for the EAGLE target worker: target "
|
||||
"prefill needs FULL hidden states but the prefill graph captures "
|
||||
"NULL, so the graph is unused; skipping its capture keeps decode "
|
||||
"graph capture clean."
|
||||
)
|
||||
self.prefill_cuda_graph_runner = self.eager_runner
|
||||
return
|
||||
|
||||
# Disable piecewise CUDA graph for non-language models
|
||||
if not hasattr(self.model, "model"):
|
||||
logger.warning(
|
||||
|
||||
@@ -7,7 +7,7 @@ BaseCudaGraphBackend chosen via cuda_graph_config.
|
||||
|
||||
Public API:
|
||||
- BaseRunner — minimal abstract base shared by the cuda-graph runners
|
||||
and the eager runner (shared __init__ + abstract
|
||||
and the eager runner (shared __init__ + warmup + abstract
|
||||
can_run_graph/load_batch/execute).
|
||||
- BaseCudaGraphRunner — abstract cuda-graph base; bucket padding +
|
||||
capture-loop scaffolding on top of BaseRunner.
|
||||
|
||||
@@ -15,22 +15,176 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
import hashlib
|
||||
import inspect
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.batch_overlap.two_batch_overlap import TboCudaGraphRunnerPlugin
|
||||
from sglang.srt.compilation.torch_compile_decoration import set_torch_compile_config
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
set_dp_buffer_len,
|
||||
set_is_extend_in_batch,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
NgramEmbeddingInfo,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.speculative.spec_info import create_dummy_verify_input
|
||||
from sglang.srt.utils import (
|
||||
empty_context,
|
||||
log_info_on_rank0,
|
||||
require_attn_tp_gather,
|
||||
require_gathered_buffer,
|
||||
require_mlp_tp_gather,
|
||||
)
|
||||
from sglang.srt.utils.common import ceil_align, require_mlp_sync
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _allocate_decode_buffers(
|
||||
*,
|
||||
device: torch.device,
|
||||
max_bs: int,
|
||||
max_num_token: int,
|
||||
hidden_size: int,
|
||||
vocab_size: int,
|
||||
dtype: torch.dtype,
|
||||
dp_size: int,
|
||||
pp_size: int,
|
||||
is_encoder_decoder: bool,
|
||||
require_mlp_tp_gather: bool,
|
||||
seq_len_fill_value: int,
|
||||
encoder_len_fill_value: int,
|
||||
num_tokens_per_bs: int,
|
||||
cache_loc_dtype: torch.dtype,
|
||||
enable_mamba_track: bool,
|
||||
ne_token_table: Optional[torch.Tensor] = None,
|
||||
hc_hidden_size: Optional[int] = None,
|
||||
) -> SimpleNamespace:
|
||||
"""Allocate the FB-shared decode buffers."""
|
||||
with torch.device(device):
|
||||
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
|
||||
input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype)
|
||||
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
|
||||
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int64)
|
||||
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
|
||||
positions = torch.zeros((max_num_token,), dtype=torch.int64)
|
||||
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
|
||||
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
|
||||
custom_mask = torch.ones(
|
||||
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs,
|
||||
dtype=torch.bool,
|
||||
)
|
||||
next_token_logits_buffer = torch.zeros(
|
||||
(max_num_token, vocab_size),
|
||||
dtype=torch.float,
|
||||
)
|
||||
mamba_track_indices = (
|
||||
torch.zeros((max_bs,), dtype=torch.int64) if enable_mamba_track else None
|
||||
)
|
||||
mamba_track_mask = (
|
||||
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
|
||||
)
|
||||
|
||||
if pp_size > 1:
|
||||
# mHC (e.g. DSV4) flattens residual into hidden_states (size = hc_hidden_size).
|
||||
is_mhc = hc_hidden_size is not None
|
||||
hs = hc_hidden_size if is_mhc else hidden_size
|
||||
pp_proxy_tensors = {
|
||||
"hidden_states": torch.zeros((max_bs, hs), dtype=dtype),
|
||||
}
|
||||
if not is_mhc:
|
||||
pp_proxy_tensors["residual"] = torch.zeros(
|
||||
(max_bs, hidden_size), dtype=dtype
|
||||
)
|
||||
else:
|
||||
pp_proxy_tensors = None
|
||||
|
||||
if is_encoder_decoder:
|
||||
encoder_lens = torch.full(
|
||||
(max_bs,), encoder_len_fill_value, dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
encoder_lens = None
|
||||
|
||||
if require_mlp_tp_gather:
|
||||
global_num_tokens_gpu = torch.zeros((dp_size,), dtype=torch.int32)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(dp_size,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
|
||||
ngram_embedding_info = (
|
||||
NgramEmbeddingInfo(
|
||||
token_table=ne_token_table,
|
||||
column_starts=torch.zeros([max_bs], dtype=torch.int32),
|
||||
req_lens=torch.ones([max_bs], dtype=torch.int32),
|
||||
out_column_starts=torch.zeros([max_bs], dtype=torch.int32),
|
||||
out_req_lens=torch.ones([max_bs], dtype=torch.int32),
|
||||
)
|
||||
if ne_token_table is not None
|
||||
else None
|
||||
)
|
||||
|
||||
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
|
||||
rids_int = torch.zeros((max_bs,), dtype=torch.int64)
|
||||
bootstrap_room_ids_int = torch.full((max_bs,), -1, dtype=torch.int64)
|
||||
else:
|
||||
rids_int = None
|
||||
bootstrap_room_ids_int = None
|
||||
|
||||
seq_lens_cpu = torch.full(
|
||||
(max_bs,),
|
||||
seq_len_fill_value,
|
||||
dtype=torch.int64,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
return SimpleNamespace(
|
||||
input_ids=input_ids,
|
||||
input_embeds=input_embeds,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
out_cache_loc=out_cache_loc,
|
||||
positions=positions,
|
||||
mrope_positions=mrope_positions,
|
||||
num_token_non_padded=num_token_non_padded,
|
||||
custom_mask=custom_mask,
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
mamba_track_indices=mamba_track_indices,
|
||||
mamba_track_mask=mamba_track_mask,
|
||||
encoder_lens=encoder_lens,
|
||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||
pp_proxy_tensors=pp_proxy_tensors,
|
||||
ngram_embedding_info=ngram_embedding_info,
|
||||
rids_int=rids_int,
|
||||
bootstrap_room_ids_int=bootstrap_room_ids_int,
|
||||
)
|
||||
|
||||
|
||||
class BaseRunner(ABC):
|
||||
def __init__(self, model_runner: ModelRunner) -> None:
|
||||
self.model_runner = model_runner
|
||||
@@ -43,6 +197,438 @@ class BaseRunner(ABC):
|
||||
self.attn_tp_rank = get_parallel().attn_tp_rank
|
||||
self.tbo_plugin = TboCudaGraphRunnerPlugin()
|
||||
|
||||
def warmup(self) -> None:
|
||||
"""Run kernel warmup + autotune once, gated by mr._kernel_warmed_up."""
|
||||
mr = self.model_runner
|
||||
if getattr(mr, "_kernel_warmed_up", False):
|
||||
return
|
||||
mr._kernel_warmed_up = True
|
||||
|
||||
if mr.device != "cuda":
|
||||
return
|
||||
|
||||
self._pre_initialize_flashinfer_allreduce_workspace()
|
||||
|
||||
if self._should_run_flashinfer_autotune():
|
||||
self._flashinfer_autotune()
|
||||
|
||||
if (
|
||||
envs.SGLANG_PP_PARALLEL_DEEPGEMM_WARMUP.get()
|
||||
and deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM
|
||||
and mr.pp_size > 1
|
||||
and not mr.spec_algorithm.is_speculative()
|
||||
):
|
||||
from sglang.srt.layers.deep_gemm_wrapper.compile_utils import (
|
||||
pp_parallel_deep_gemm_warmup,
|
||||
)
|
||||
|
||||
pp_parallel_deep_gemm_warmup(self)
|
||||
|
||||
def _pre_initialize_flashinfer_allreduce_workspace(self):
|
||||
"""Allocate flashinfer allreduce workspaces; must run before CG capture
|
||||
to keep broadcasts/barriers outside the capture context (else deadlock
|
||||
with custom_all_reduce.register_graph_buffers).
|
||||
"""
|
||||
mr = self.model_runner
|
||||
if mr.server_args.flashinfer_allreduce_fusion_backend is None:
|
||||
return
|
||||
|
||||
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
|
||||
from sglang.srt.layers.flashinfer_comm_fusion import pre_initialize_workspaces
|
||||
|
||||
pre_initialize_workspaces(
|
||||
max_token_num=FUSE_ALLREDUCE_MAX_BATCH_SIZE,
|
||||
hidden_dim=mr.model_config.hidden_size,
|
||||
dtype=mr.dtype,
|
||||
)
|
||||
|
||||
def _should_run_flashinfer_autotune(self) -> bool:
|
||||
"""Check if flashinfer autotune should be run."""
|
||||
mr = self.model_runner
|
||||
if mr.server_args.disable_flashinfer_autotune:
|
||||
return False
|
||||
|
||||
# CuteDSL v1 (cutedsl runner + deepep a2a) bypasses MoeRunner and must not
|
||||
# be autotuned -- its _dummy_run would dispatch more tokens per rank than
|
||||
# SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK, tripping a DeepEP assert.
|
||||
# Read server_args directly to avoid depending on initialize_moe_config()
|
||||
# having already populated the MoE backend globals.
|
||||
if (
|
||||
mr.server_args.moe_runner_backend == "flashinfer_cutedsl"
|
||||
and mr.server_args.moe_a2a_backend == "deepep"
|
||||
):
|
||||
return False
|
||||
|
||||
backend_str = mr.server_args.moe_runner_backend
|
||||
|
||||
# TODO smor- support other cases for flashinfer autotune, such as, mamba backend
|
||||
|
||||
moe_needs_autotune = backend_str in [
|
||||
"flashinfer_trtllm",
|
||||
"flashinfer_trtllm_routed",
|
||||
"flashinfer_mxfp4",
|
||||
"flashinfer_cutedsl",
|
||||
"flashinfer_cutlass",
|
||||
]
|
||||
|
||||
from sglang.srt.layers.quantization.fp4_utils import (
|
||||
get_fp4_gemm_runner_backend,
|
||||
)
|
||||
|
||||
model_uses_fp4 = mr.model_config.quantization in (
|
||||
"modelopt_fp4",
|
||||
"modelopt_mixed",
|
||||
)
|
||||
fp4_gemm_needs_autotune = model_uses_fp4 and (
|
||||
get_fp4_gemm_runner_backend().is_flashinfer_cutlass()
|
||||
or get_fp4_gemm_runner_backend().is_flashinfer_cutedsl()
|
||||
)
|
||||
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
get_fp8_gemm_runner_backend,
|
||||
)
|
||||
from sglang.srt.utils import is_sm100_supported
|
||||
|
||||
model_uses_modelopt_fp8 = mr.model_config.quantization in (
|
||||
"modelopt",
|
||||
"modelopt_fp8",
|
||||
"modelopt_mixed",
|
||||
)
|
||||
fp8_gemm_needs_autotune = (
|
||||
get_fp8_gemm_runner_backend().is_flashinfer_cutlass()
|
||||
or (model_uses_modelopt_fp8 and is_sm100_supported())
|
||||
)
|
||||
|
||||
if not (
|
||||
moe_needs_autotune or fp4_gemm_needs_autotune or fp8_gemm_needs_autotune
|
||||
):
|
||||
return False
|
||||
|
||||
major, _ = torch.cuda.get_device_capability()
|
||||
if major < 9:
|
||||
return False
|
||||
|
||||
if mr.spec_algorithm.is_speculative():
|
||||
return not mr.is_draft_worker
|
||||
|
||||
return True
|
||||
|
||||
def _flashinfer_autotune(self):
|
||||
"""Run flashinfer autotune."""
|
||||
from flashinfer.autotuner import autotune
|
||||
|
||||
from sglang.srt.layers.logits_processor import autotune_dummy_run_mode
|
||||
|
||||
mr = self.model_runner
|
||||
cache_path = self._flashinfer_autotune_cache_path()
|
||||
if envs.SGLANG_FLASHINFER_AUTOTUNE_CACHE.get():
|
||||
autotune_cache = cache_path
|
||||
logger.info("Running FlashInfer autotune with cache: %s", autotune_cache)
|
||||
else:
|
||||
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
runs_dir = cache_path.parent / "runs"
|
||||
runs_dir.mkdir(parents=True, exist_ok=True)
|
||||
autotune_cache = (
|
||||
runs_dir / f"{cache_path.stem}.{timestamp}{cache_path.suffix}"
|
||||
)
|
||||
logger.info(
|
||||
"Running FlashInfer autotune (cache reuse DISABLED via "
|
||||
"SGLANG_FLASHINFER_AUTOTUNE_CACHE=0); writing fresh result to: %s",
|
||||
autotune_cache,
|
||||
)
|
||||
|
||||
# Run warmup on the non-default stream to avoid NCCL 2.29+ cudaMemcpyBatchAsync
|
||||
# calls on default stream (unsupported by CUDA) when --enable-symm-mem is used.
|
||||
mr.forward_stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.get_device_module(mr.device).stream(mr.forward_stream):
|
||||
with (
|
||||
torch.inference_mode(),
|
||||
autotune(True, cache=str(autotune_cache)),
|
||||
autotune_dummy_run_mode(),
|
||||
):
|
||||
self._dummy_run(batch_size=mr.req_to_token_pool.size)
|
||||
torch.cuda.current_stream().wait_stream(mr.forward_stream)
|
||||
logger.info("FlashInfer autotune completed.")
|
||||
|
||||
def _flashinfer_autotune_cache_path(self) -> Path:
|
||||
import flashinfer
|
||||
|
||||
mr = self.model_runner
|
||||
major, minor = torch.cuda.get_device_capability(mr.device)
|
||||
arch = f"sm{major}{minor}"
|
||||
flashinfer_version = getattr(flashinfer, "__version__", "unknown")
|
||||
|
||||
server_args = mr.server_args
|
||||
model_key = "|".join(
|
||||
[
|
||||
str(server_args.model_path),
|
||||
str(mr.dtype),
|
||||
str(server_args.quantization),
|
||||
str(server_args.moe_runner_backend),
|
||||
str(mr.tp_size),
|
||||
str(mr.pp_size),
|
||||
str(mr.dp_size),
|
||||
str(mr.moe_ep_size),
|
||||
str(mr.model_config.hf_config.__class__.__name__),
|
||||
]
|
||||
)
|
||||
cache_key = hashlib.sha256(model_key.encode()).hexdigest()[:16]
|
||||
cache_dir = (
|
||||
Path(envs.SGLANG_CACHE_DIR.get())
|
||||
/ "flashinfer"
|
||||
/ "autotune"
|
||||
/ flashinfer_version
|
||||
/ arch
|
||||
/ cache_key
|
||||
)
|
||||
cache_dir.mkdir(parents=True, exist_ok=True)
|
||||
return (
|
||||
cache_dir / f"rank_tp{mr.tp_rank}_pp{mr.pp_rank}_dp{mr.dp_rank or 0}.json"
|
||||
)
|
||||
|
||||
def _dummy_run(
|
||||
self,
|
||||
batch_size: int,
|
||||
run_ctx=None,
|
||||
forward_mode_override: Optional[ForwardMode] = None,
|
||||
):
|
||||
"""Run a dummy forward pass for warmup/profiling.
|
||||
|
||||
forward_mode_override forces EXTEND/DECODE regardless of
|
||||
is_generation (used by the PP-parallel DeepGEMM warmup).
|
||||
"""
|
||||
mr = self.model_runner
|
||||
if forward_mode_override is not None:
|
||||
capture_forward_mode = forward_mode_override
|
||||
elif mr.is_generation:
|
||||
capture_forward_mode = ForwardMode.DECODE
|
||||
else:
|
||||
capture_forward_mode = ForwardMode.EXTEND
|
||||
capture_hidden_mode = CaptureHiddenMode.NULL
|
||||
num_tokens_per_bs = 1
|
||||
if mr.spec_algorithm.is_speculative():
|
||||
if mr.is_draft_worker:
|
||||
if not mr.spec_algorithm.supports_target_verify_for_draft():
|
||||
raise RuntimeError("This should not happen")
|
||||
capture_forward_mode = ForwardMode.TARGET_VERIFY
|
||||
num_tokens_per_bs = (
|
||||
mr.spec_algorithm.get_num_tokens_per_bs_for_target_verify(
|
||||
mr.server_args.speculative_num_draft_tokens, mr.is_draft_worker
|
||||
)
|
||||
)
|
||||
|
||||
if mr.server_args.enable_return_hidden_states:
|
||||
capture_hidden_mode = CaptureHiddenMode.FULL
|
||||
|
||||
num_tokens = batch_size * num_tokens_per_bs
|
||||
|
||||
# Keep warmup aligned with scheduler MLP-sync padding.
|
||||
if require_mlp_sync(mr.server_args):
|
||||
attn_tp_size = get_parallel().attn_tp_size
|
||||
if attn_tp_size > 1 and num_tokens % attn_tp_size != 0:
|
||||
num_tokens = ceil_align(num_tokens, attn_tp_size)
|
||||
batch_size = num_tokens // num_tokens_per_bs
|
||||
|
||||
seq_len_fill_value = mr.attn_backend.get_cuda_graph_seq_len_fill_value()
|
||||
|
||||
if mr.server_args.enable_torch_compile:
|
||||
set_torch_compile_config()
|
||||
should_disable_torch_compile = not getattr(
|
||||
mr.model, "_can_torch_compile", True
|
||||
)
|
||||
if should_disable_torch_compile:
|
||||
log_info_on_rank0(
|
||||
logger,
|
||||
"Transformers backend model reports it is not torch.compile "
|
||||
"compatible (e.g. dynamic rope scaling). Disabling torch.compile.",
|
||||
)
|
||||
mr.server_args.enable_torch_compile = False
|
||||
|
||||
# NOTE: aux hidden state capture (eagle3/dflash) is already
|
||||
# configured by init_aux_hidden_state_capture() in initialize().
|
||||
|
||||
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args)
|
||||
if require_gathered_buffer(mr.server_args):
|
||||
assert require_mlp_tp_gather_ or require_attn_tp_gather(mr.server_args)
|
||||
|
||||
buffers = _allocate_decode_buffers(
|
||||
device=mr.device,
|
||||
max_bs=batch_size,
|
||||
max_num_token=num_tokens,
|
||||
hidden_size=mr.model_config.hidden_size,
|
||||
vocab_size=mr.model_config.vocab_size,
|
||||
dtype=mr.model_config.dtype,
|
||||
dp_size=mr.server_args.dp_size,
|
||||
pp_size=mr.server_args.pp_size,
|
||||
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
||||
require_mlp_tp_gather=require_mlp_tp_gather_,
|
||||
seq_len_fill_value=seq_len_fill_value,
|
||||
encoder_len_fill_value=(
|
||||
getattr(mr.model_config.hf_config, "max_source_positions", 0)
|
||||
if mr.model_config.is_encoder_decoder
|
||||
else 0
|
||||
),
|
||||
num_tokens_per_bs=num_tokens_per_bs,
|
||||
cache_loc_dtype=torch.int64,
|
||||
enable_mamba_track=False,
|
||||
hc_hidden_size=getattr(mr.model_config, "hc_hidden_size", None),
|
||||
)
|
||||
buffers.num_token_non_padded[...] = num_tokens
|
||||
|
||||
# For extend mode
|
||||
if capture_forward_mode == ForwardMode.EXTEND:
|
||||
extend_prefix_lens_cpu = [0] * batch_size
|
||||
extend_seq_lens_cpu = [seq_len_fill_value] * batch_size
|
||||
extend_num_tokens = num_tokens
|
||||
extend_seq_lens = torch.full(
|
||||
(batch_size,), seq_len_fill_value, dtype=torch.int32, device=mr.device
|
||||
)
|
||||
extend_prefix_lens = torch.zeros(
|
||||
(batch_size,), dtype=torch.int32, device=mr.device
|
||||
)
|
||||
extend_start_loc = torch.arange(
|
||||
0, num_tokens, num_tokens_per_bs, dtype=torch.int32, device=mr.device
|
||||
)
|
||||
else:
|
||||
extend_prefix_lens_cpu = None
|
||||
extend_seq_lens_cpu = None
|
||||
extend_num_tokens = None
|
||||
extend_seq_lens = None
|
||||
extend_prefix_lens = None
|
||||
extend_start_loc = None
|
||||
|
||||
if mr.server_args.pp_size > 1:
|
||||
# PP0 already cp-split hidden_states before send.
|
||||
pp_hidden_tokens = num_tokens
|
||||
if (
|
||||
capture_forward_mode == ForwardMode.EXTEND
|
||||
and mr.pp_rank != 0
|
||||
and mr.attn_cp_size > 1
|
||||
):
|
||||
pp_hidden_tokens = num_tokens // mr.attn_cp_size
|
||||
pp_proxy_tensors = PPProxyTensors(
|
||||
{k: v[:pp_hidden_tokens] for k, v in buffers.pp_proxy_tensors.items()}
|
||||
)
|
||||
|
||||
if require_mlp_tp_gather_:
|
||||
global_num_tokens_cpu = [num_tokens] * mr.server_args.dp_size
|
||||
elif require_attn_tp_gather(mr.server_args):
|
||||
global_num_tokens_cpu = [num_tokens]
|
||||
else:
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
if global_num_tokens_cpu is not None:
|
||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||
num_tokens_tensor = torch.tensor(
|
||||
global_num_tokens_cpu, dtype=torch.int32, device=mr.device
|
||||
)
|
||||
buffers.global_num_tokens_gpu.copy_(num_tokens_tensor)
|
||||
buffers.global_num_tokens_for_logprob_gpu.copy_(num_tokens_tensor)
|
||||
else:
|
||||
global_dp_buffer_len = None
|
||||
global_num_tokens_cpu = None
|
||||
|
||||
spec_info = create_dummy_verify_input(
|
||||
mr.spec_algorithm,
|
||||
mr.server_args,
|
||||
buffers.custom_mask,
|
||||
num_tokens_per_bs,
|
||||
mr.is_draft_worker,
|
||||
)
|
||||
if spec_info is not None and (
|
||||
mr.spec_algorithm.is_eagle() or mr.spec_algorithm.is_standalone()
|
||||
):
|
||||
# MTP models (e.g. deepseek_nextn) read spec_info.hidden_states
|
||||
# during forward; provide a dummy so warmup doesn't crash.
|
||||
spec_info.hidden_states = torch.zeros(
|
||||
(num_tokens, mr.model_config.hidden_size),
|
||||
dtype=mr.dtype,
|
||||
device=mr.device,
|
||||
)
|
||||
if capture_hidden_mode != CaptureHiddenMode.FULL:
|
||||
capture_hidden_mode = (
|
||||
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
|
||||
)
|
||||
|
||||
if mr.server_args.enable_lora:
|
||||
lora_ids = [None] * batch_size
|
||||
else:
|
||||
lora_ids = None
|
||||
|
||||
forward_batch = ForwardBatch(
|
||||
forward_mode=capture_forward_mode,
|
||||
batch_size=batch_size,
|
||||
input_ids=buffers.input_ids,
|
||||
req_pool_indices=buffers.req_pool_indices,
|
||||
seq_lens=buffers.seq_lens,
|
||||
seq_lens_cpu=buffers.seq_lens_cpu,
|
||||
next_token_logits_buffer=buffers.next_token_logits_buffer,
|
||||
orig_seq_lens=buffers.seq_lens,
|
||||
out_cache_loc=buffers.out_cache_loc,
|
||||
seq_lens_sum=buffers.seq_lens.sum().item(),
|
||||
encoder_lens=buffers.encoder_lens,
|
||||
return_logprob=False,
|
||||
positions=buffers.positions,
|
||||
extend_num_tokens=extend_num_tokens,
|
||||
extend_seq_lens=extend_seq_lens,
|
||||
extend_prefix_lens=extend_prefix_lens,
|
||||
extend_start_loc=extend_start_loc,
|
||||
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
|
||||
extend_seq_lens_cpu=extend_seq_lens_cpu,
|
||||
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
|
||||
global_num_tokens_cpu=global_num_tokens_cpu,
|
||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||
global_dp_buffer_len=global_dp_buffer_len,
|
||||
mrope_positions=buffers.mrope_positions,
|
||||
spec_algorithm=mr.spec_algorithm,
|
||||
spec_info=spec_info,
|
||||
capture_hidden_mode=capture_hidden_mode,
|
||||
num_token_non_padded=buffers.num_token_non_padded,
|
||||
global_forward_mode=capture_forward_mode,
|
||||
lora_ids=lora_ids,
|
||||
)
|
||||
|
||||
if lora_ids is not None:
|
||||
mr.lora_manager.prepare_lora_batch(forward_batch)
|
||||
|
||||
mr.attn_backend.init_forward_metadata(forward_batch)
|
||||
|
||||
def run_once():
|
||||
forward_batch.dp_local_start_pos = forward_batch.dp_local_num_tokens = None
|
||||
set_dp_buffer_len(
|
||||
global_dp_buffer_len,
|
||||
num_tokens,
|
||||
forward_batch.dp_padding_mode.is_max_len(),
|
||||
global_num_tokens_cpu,
|
||||
)
|
||||
set_is_extend_in_batch(False)
|
||||
|
||||
kwargs = {}
|
||||
if (
|
||||
mr.server_args.pp_size > 1
|
||||
and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters
|
||||
):
|
||||
kwargs["pp_proxy_tensors"] = PPProxyTensors(
|
||||
{k: v.clone() for k, v in pp_proxy_tensors.tensors.items()}
|
||||
)
|
||||
if not mr.is_generation:
|
||||
kwargs["get_embedding"] = True
|
||||
|
||||
logits_output_or_pp_proxy_tensors = mr.model.forward(
|
||||
buffers.input_ids,
|
||||
forward_batch.positions,
|
||||
forward_batch,
|
||||
**kwargs,
|
||||
)
|
||||
return logits_output_or_pp_proxy_tensors
|
||||
|
||||
torch.get_device_module(mr.device).synchronize()
|
||||
mr.tp_group.barrier()
|
||||
with forward_context(ForwardContext(attn_backend=mr.attn_backend)):
|
||||
with torch.inference_mode(), run_ctx or empty_context():
|
||||
run_once()
|
||||
|
||||
@abstractmethod
|
||||
def can_run_graph(self, forward_batch: ForwardBatch) -> bool: ...
|
||||
|
||||
|
||||
@@ -43,7 +43,6 @@ from sglang.srt.distributed.parallel_state import (
|
||||
set_pdmux_status,
|
||||
)
|
||||
from sglang.srt.dllm.config import DllmConfig
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
DpPaddingMode,
|
||||
@@ -62,7 +61,6 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
CaptureHiddenMode,
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
NgramEmbeddingInfo,
|
||||
PPProxyTensors,
|
||||
compute_local_num_token_non_padded,
|
||||
enable_num_token_non_padded,
|
||||
@@ -163,137 +161,6 @@ def build_replay_fb_view(
|
||||
)
|
||||
|
||||
|
||||
def _allocate_decode_buffers(
|
||||
*,
|
||||
device: torch.device,
|
||||
max_bs: int,
|
||||
max_num_token: int,
|
||||
hidden_size: int,
|
||||
vocab_size: int,
|
||||
dtype: torch.dtype,
|
||||
dp_size: int,
|
||||
pp_size: int,
|
||||
is_encoder_decoder: bool,
|
||||
require_mlp_tp_gather: bool,
|
||||
seq_len_fill_value: int,
|
||||
encoder_len_fill_value: int,
|
||||
num_tokens_per_bs: int,
|
||||
cache_loc_dtype: torch.dtype,
|
||||
enable_mamba_track: bool,
|
||||
ne_token_table: Optional[torch.Tensor] = None,
|
||||
hc_hidden_size: Optional[int] = None,
|
||||
pp_proxy_topk_size: Optional[int] = None,
|
||||
) -> SimpleNamespace:
|
||||
"""Allocate the FB-shared decode buffers as a namespace adopted by
|
||||
``build_decode_registry(source=...)``."""
|
||||
with torch.device(device):
|
||||
input_ids = torch.zeros((max_num_token,), dtype=torch.int64)
|
||||
input_embeds = torch.zeros((max_num_token, hidden_size), dtype=dtype)
|
||||
req_pool_indices = torch.zeros((max_bs,), dtype=torch.int64)
|
||||
seq_lens = torch.full((max_bs,), seq_len_fill_value, dtype=torch.int64)
|
||||
out_cache_loc = torch.zeros((max_num_token,), dtype=cache_loc_dtype)
|
||||
positions = torch.zeros((max_num_token,), dtype=torch.int64)
|
||||
mrope_positions = torch.zeros((3, max_num_token), dtype=torch.int64)
|
||||
num_token_non_padded = torch.zeros((1,), dtype=torch.int32)
|
||||
custom_mask = torch.ones(
|
||||
(max_bs * seq_len_fill_value + max_num_token) * num_tokens_per_bs,
|
||||
dtype=torch.bool,
|
||||
)
|
||||
next_token_logits_buffer = torch.zeros(
|
||||
(max_num_token, vocab_size),
|
||||
dtype=torch.float,
|
||||
)
|
||||
mamba_track_indices = (
|
||||
torch.zeros((max_bs,), dtype=torch.int64) if enable_mamba_track else None
|
||||
)
|
||||
mamba_track_mask = (
|
||||
torch.zeros((max_bs,), dtype=torch.bool) if enable_mamba_track else None
|
||||
)
|
||||
|
||||
if pp_size > 1:
|
||||
# mHC (e.g. DSV4) flattens residual into hidden_states (size = hc_hidden_size).
|
||||
is_mhc = hc_hidden_size is not None
|
||||
hs = hc_hidden_size if is_mhc else hidden_size
|
||||
pp_proxy_tensors = {
|
||||
"hidden_states": torch.zeros((max_bs, hs), dtype=dtype),
|
||||
}
|
||||
if not is_mhc:
|
||||
pp_proxy_tensors["residual"] = torch.zeros(
|
||||
(max_bs, hidden_size), dtype=dtype
|
||||
)
|
||||
if pp_proxy_topk_size is not None:
|
||||
pp_proxy_tensors["topk_indices"] = torch.zeros(
|
||||
(max_num_token, pp_proxy_topk_size), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
pp_proxy_tensors = None
|
||||
|
||||
if is_encoder_decoder:
|
||||
encoder_lens = torch.full(
|
||||
(max_bs,), encoder_len_fill_value, dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
encoder_lens = None
|
||||
|
||||
if require_mlp_tp_gather:
|
||||
global_num_tokens_gpu = torch.zeros((dp_size,), dtype=torch.int32)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros(
|
||||
(dp_size,), dtype=torch.int32
|
||||
)
|
||||
else:
|
||||
global_num_tokens_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
global_num_tokens_for_logprob_gpu = torch.zeros((1,), dtype=torch.int32)
|
||||
|
||||
ngram_embedding_info = (
|
||||
NgramEmbeddingInfo(
|
||||
token_table=ne_token_table,
|
||||
column_starts=torch.zeros([max_bs], dtype=torch.int32),
|
||||
req_lens=torch.ones([max_bs], dtype=torch.int32),
|
||||
out_column_starts=torch.zeros([max_bs], dtype=torch.int32),
|
||||
out_req_lens=torch.ones([max_bs], dtype=torch.int32),
|
||||
)
|
||||
if ne_token_table is not None
|
||||
else None
|
||||
)
|
||||
|
||||
if envs.SGLANG_KV_CANARY_ENABLE_TOKEN_ORACLE.get():
|
||||
rids_int = torch.zeros((max_bs,), dtype=torch.int64)
|
||||
bootstrap_room_ids_int = torch.full((max_bs,), -1, dtype=torch.int64)
|
||||
else:
|
||||
rids_int = None
|
||||
bootstrap_room_ids_int = None
|
||||
|
||||
seq_lens_cpu = torch.full(
|
||||
(max_bs,),
|
||||
seq_len_fill_value,
|
||||
dtype=torch.int64,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
return SimpleNamespace(
|
||||
input_ids=input_ids,
|
||||
input_embeds=input_embeds,
|
||||
req_pool_indices=req_pool_indices,
|
||||
seq_lens=seq_lens,
|
||||
seq_lens_cpu=seq_lens_cpu,
|
||||
out_cache_loc=out_cache_loc,
|
||||
positions=positions,
|
||||
mrope_positions=mrope_positions,
|
||||
num_token_non_padded=num_token_non_padded,
|
||||
custom_mask=custom_mask,
|
||||
next_token_logits_buffer=next_token_logits_buffer,
|
||||
mamba_track_indices=mamba_track_indices,
|
||||
mamba_track_mask=mamba_track_mask,
|
||||
encoder_lens=encoder_lens,
|
||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||
pp_proxy_tensors=pp_proxy_tensors,
|
||||
ngram_embedding_info=ngram_embedding_info,
|
||||
rids_int=rids_int,
|
||||
bootstrap_room_ids_int=bootstrap_room_ids_int,
|
||||
)
|
||||
|
||||
|
||||
class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
"""Decode-phase CUDA graph runner.
|
||||
|
||||
@@ -775,6 +642,18 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
return forward_batch, attn_backend, pp_proxy_tensors
|
||||
|
||||
def capture(self) -> None:
|
||||
# Warm up + autotune kernels once before capture (run-once across the
|
||||
# decode + prefill runners; see BaseRunner.warmup).
|
||||
self.warmup()
|
||||
# warmup() may disable torch.compile for a model whose _can_torch_compile
|
||||
# is False; recompute the compile bucket so capture matches.
|
||||
if self.enable_torch_compile and not (
|
||||
self.model_runner.server_args.enable_torch_compile
|
||||
):
|
||||
self.enable_torch_compile = False
|
||||
_, self.compile_bs = get_batch_sizes_to_capture(
|
||||
self.model_runner, self.num_tokens_per_bs
|
||||
)
|
||||
profile_context = empty_context()
|
||||
if self.enable_profile_cuda_graph:
|
||||
profile_context = self._init_profile_context_and_memory_record()
|
||||
|
||||
@@ -116,6 +116,8 @@ class EagerRunner(BaseRunner):
|
||||
),
|
||||
dp_size=sa.dp_size,
|
||||
)
|
||||
# Eager has no capture step, so warm up here (run-once via mr._kernel_warmed_up).
|
||||
self.warmup()
|
||||
|
||||
def can_run_graph(self, forward_batch: ForwardBatch) -> bool:
|
||||
# Eager never runs a cuda graph; callers dispatch on isinstance(...,
|
||||
|
||||
@@ -587,6 +587,9 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
return forward_batch, self.model_runner.attn_backend
|
||||
|
||||
def capture(self) -> None:
|
||||
# Warm up + autotune kernels once before capture (run-once across the
|
||||
# decode + prefill runners; see BaseRunner.warmup).
|
||||
self.warmup()
|
||||
with freeze_gc(self.model_runner.server_args.enable_cudagraph_gc):
|
||||
with graph_capture() as graph_capture_context:
|
||||
self.stream = graph_capture_context.stream
|
||||
|
||||
@@ -325,6 +325,11 @@ class MockModelRunner(ModelRunner):
|
||||
self.pp_size = 1
|
||||
self.is_draft_worker = False
|
||||
self.spec_algorithm = SpeculativeAlgorithm.NONE
|
||||
# The runner lifecycle warms up kernels in capture() / first execute()
|
||||
# via BaseRunner.warmup(); this mock never calls init_backends and has no
|
||||
# real kernels to warm up, so mark it done (warmup becomes a no-op for
|
||||
# the runner-mode attention tests that drive capture directly).
|
||||
self._kernel_warmed_up = True
|
||||
speculative_num_draft_tokens = (
|
||||
max(case.input_lens)
|
||||
if case.forward_mode.is_target_verify()
|
||||
|
||||
@@ -306,6 +306,7 @@ class DSAMockModelRunner(ModelRunner):
|
||||
self.page_size = case.page_size
|
||||
self.model_config = model_config
|
||||
self.tp_size = 1
|
||||
self._kernel_warmed_up = True
|
||||
self.dp_size = 1
|
||||
self.pp_size = 1
|
||||
self.server_args = make_mock_server_args(
|
||||
|
||||
@@ -412,6 +412,7 @@ class MockDSV4ModelRunner:
|
||||
self.sliding_window_size = DSV4_SWA_WINDOW
|
||||
self.use_mla_backend = True
|
||||
self.is_draft_worker = False
|
||||
self._kernel_warmed_up = True
|
||||
|
||||
@property
|
||||
def hybrid_gdn_config(self):
|
||||
|
||||
@@ -321,6 +321,7 @@ class DualChunkMockModelRunner(ModelRunner):
|
||||
self.page_size = case.page_size
|
||||
self.model_config = model_config
|
||||
self.tp_size = 1
|
||||
self._kernel_warmed_up = True
|
||||
self.dp_size = 1
|
||||
self.pp_size = 1
|
||||
self.server_args = make_mock_server_args(
|
||||
|
||||
@@ -304,6 +304,7 @@ class MockGDNModelRunner(ModelRunner):
|
||||
self.sliding_window_size = None
|
||||
self.use_mla_backend = False
|
||||
self.is_draft_worker = False
|
||||
self._kernel_warmed_up = True
|
||||
|
||||
@property
|
||||
def hybrid_gdn_config(self):
|
||||
|
||||
@@ -310,6 +310,7 @@ class MockKDAModelRunner(ModelRunner):
|
||||
self.sliding_window_size = None
|
||||
self.use_mla_backend = False
|
||||
self.is_draft_worker = False
|
||||
self._kernel_warmed_up = True
|
||||
|
||||
@property
|
||||
def hybrid_gdn_config(self):
|
||||
|
||||
@@ -319,6 +319,7 @@ class MockLightningModelRunner(ModelRunner):
|
||||
self.sliding_window_size = None
|
||||
self.use_mla_backend = False
|
||||
self.is_draft_worker = False
|
||||
self._kernel_warmed_up = True
|
||||
|
||||
@property
|
||||
def hybrid_gdn_config(self):
|
||||
|
||||
@@ -454,6 +454,7 @@ class MockMamba2ModelRunner(ModelRunner):
|
||||
self.sliding_window_size = None
|
||||
self.use_mla_backend = False
|
||||
self.is_draft_worker = False
|
||||
self._kernel_warmed_up = True
|
||||
|
||||
@property
|
||||
def hybrid_gdn_config(self):
|
||||
|
||||
@@ -306,6 +306,7 @@ class MockMLAModelRunner(ModelRunner):
|
||||
self.sliding_window_size = None
|
||||
self.use_mla_backend = True
|
||||
self.is_draft_worker = False
|
||||
self._kernel_warmed_up = True
|
||||
|
||||
@property
|
||||
def hybrid_gdn_config(self):
|
||||
|
||||
Reference in New Issue
Block a user