config: borrowed-record reads follow the config bags (#35908)
Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 5
parent
64aa859da2
commit
362c2ee849
@@ -79,10 +79,10 @@ from sglang.srt.layers.quantization.fp8_utils import initialize_fp8_gemm_config
|
||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
||||
from sglang.srt.managers.scheduler_components.dp_attn import prepare_mlp_sync_batch_raw
|
||||
from sglang.srt.mem_cache.base_prefix_cache import EvictParams
|
||||
from sglang.srt.model_executor.cuda_graph_config import Phase
|
||||
from sglang.srt.model_executor.cuda_graph_config import Phase, cuda_graph_fully_disabled
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
@@ -487,7 +487,7 @@ class TreeCacheNamespace(SimpleNamespace):
|
||||
def extend(reqs, model_runner):
|
||||
# Create dummy tree_cache for benchmarks (no prefix caching, just allocation)
|
||||
dummy_tree_cache = TreeCacheNamespace(
|
||||
page_size=model_runner.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
device=model_runner.device,
|
||||
token_to_kv_pool_allocator=model_runner.token_to_kv_pool_allocator,
|
||||
)
|
||||
@@ -542,14 +542,14 @@ def _maybe_prepare_mlp_sync_batch(batch: ScheduleBatch, model_runner):
|
||||
prepare_mlp_sync_batch_raw(
|
||||
batch,
|
||||
model_runner=model_runner,
|
||||
dp_size=model_runner.server_args.dp_size,
|
||||
dp_size=get_parallel().dp_size,
|
||||
attn_tp_size=get_parallel().attn_tp_size,
|
||||
attn_cp_size=model_runner.ps.attn_cp_size,
|
||||
tp_group=model_runner.tp_group,
|
||||
get_idle_batch=None,
|
||||
disable_cuda_graph=model_runner.server_args.disable_cuda_graph,
|
||||
disable_cuda_graph=cuda_graph_fully_disabled(),
|
||||
require_mlp_tp_gather=require_mlp_tp_gather(model_runner.server_args),
|
||||
disable_overlap_schedule=model_runner.server_args.disable_overlap_schedule,
|
||||
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||
offload_tags=set(),
|
||||
)
|
||||
|
||||
|
||||
@@ -16,6 +16,10 @@ from typing import TYPE_CHECKING, List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.runtime_context import (
|
||||
get_schedule,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Bounded wait for a watermark advance before re-enqueueing a deferred staging
|
||||
@@ -548,7 +552,7 @@ class PrefillStagingStrategy:
|
||||
self.staging_buffer = staging_buffer
|
||||
page_size = kv_manager.kv_buffer_tensors["page_size"]
|
||||
self.full_chunk_pages = (
|
||||
staging_grid_tokens(kv_manager.server_args.chunked_prefill_size, page_size)
|
||||
staging_grid_tokens(get_schedule().chunked_prefill_size, page_size)
|
||||
// page_size
|
||||
)
|
||||
|
||||
|
||||
@@ -94,7 +94,11 @@ from sglang.srt.observability.req_time_stats import (
|
||||
set_schedule_time_batch,
|
||||
set_time_batch,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_disagg, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_memory,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.utils import ceil_align, get_num_new_pages, is_npu
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
@@ -477,7 +481,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
full_len = ceil_align(full_len, page_size)
|
||||
swa_len = ceil_align(swa_len, page_size)
|
||||
swa_reserved = self.num_reserved_decode_tokens
|
||||
if self.scheduler.server_args.disable_radix_cache:
|
||||
if get_memory().disable_radix_cache:
|
||||
swa_reserved = 0
|
||||
return (
|
||||
full_len + self.num_reserved_decode_tokens,
|
||||
@@ -556,7 +560,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
req_to_token_pool=getattr(self, "req_to_token_pool", None),
|
||||
)
|
||||
|
||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
||||
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
||||
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
||||
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
||||
kv_manager = kv_manager_class(
|
||||
@@ -608,7 +612,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
)
|
||||
|
||||
# NOTE: fake transfer does not need to resolve prefill dp rank in the pending queue
|
||||
if _is_fake_transfer(req, self.scheduler.server_args):
|
||||
if _is_fake_transfer(req):
|
||||
decode_req.kv_receiver.init(0)
|
||||
return
|
||||
|
||||
@@ -663,9 +667,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
self, req: Req, is_rebootstrap: bool = False
|
||||
) -> DecodeRequest:
|
||||
backend = (
|
||||
TransferBackend.FAKE
|
||||
if _is_fake_transfer(req, self.scheduler.server_args)
|
||||
else self.transfer_backend
|
||||
TransferBackend.FAKE if _is_fake_transfer(req) else self.transfer_backend
|
||||
)
|
||||
kv_receiver_class = get_kv_class(backend, KVClassType.RECEIVER)
|
||||
|
||||
@@ -1455,7 +1457,7 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
||||
if (
|
||||
self.scheduler.enable_hisparse
|
||||
and isinstance(self.token_to_kv_pool, DeepSeekV4TokenToKVPool)
|
||||
and not _is_fake_transfer(decode_req.req, self.scheduler.server_args)
|
||||
and not _is_fake_transfer(decode_req.req)
|
||||
):
|
||||
# alloc_logical_only() already allocated the shared logical pages
|
||||
# used by C4 indexer and C128 KV. These device buffers do not use
|
||||
@@ -2059,7 +2061,7 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
||||
else 0
|
||||
)
|
||||
|
||||
if _is_fake_transfer(decode_req.req, self.scheduler.server_args):
|
||||
if _is_fake_transfer(decode_req.req):
|
||||
pass
|
||||
elif actual_room == 0:
|
||||
# Should never happen: _poll_with_metadata_gate already confirmed
|
||||
@@ -2197,7 +2199,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
||||
self.gloo_group,
|
||||
decode_reqs=self.queue,
|
||||
metadata_buffers=self.metadata_buffers,
|
||||
server_args=self.scheduler.server_args,
|
||||
)
|
||||
|
||||
def _poll_with_staging(self) -> list:
|
||||
@@ -2206,7 +2207,6 @@ class DecodeTransferQueue(DecodeHiCacheTransferMixin):
|
||||
self.staging_handler,
|
||||
self.gloo_group,
|
||||
metadata_buffers=self.metadata_buffers,
|
||||
server_args=self.scheduler.server_args,
|
||||
)
|
||||
|
||||
def _init_staging_handler(self, kv_manager):
|
||||
|
||||
@@ -253,7 +253,7 @@ class PrefillBootstrapQueue:
|
||||
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
||||
self.metadata_buffers.get_buf_infos()
|
||||
)
|
||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
||||
kv_args.ib_device = get_disagg().disaggregation_ib_device
|
||||
kv_args.gpu_id = self.scheduler.ps.gpu_id
|
||||
|
||||
req_to_token_pool = getattr(self.scheduler, "req_to_token_pool", None)
|
||||
@@ -438,8 +438,7 @@ class PrefillBootstrapQueue:
|
||||
failed_reqs.append(req)
|
||||
elif poll == KVPoll.Bootstrapping:
|
||||
if (
|
||||
req.prefill_attempt_count
|
||||
< self.scheduler.server_args.optimistic_prefill_attempts
|
||||
req.prefill_attempt_count < get_disagg().optimistic_prefill_attempts
|
||||
and not req.is_retracted # engine paused
|
||||
):
|
||||
if not self.ensure_metadata_buffer(req):
|
||||
|
||||
@@ -22,6 +22,9 @@ import torch.distributed as dist
|
||||
from sglang.srt.configs.model_config import get_dsa_index_topk
|
||||
from sglang.srt.disaggregation.base import KVPoll
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
)
|
||||
from sglang.srt.utils import is_hip, is_npu
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -33,7 +36,6 @@ if TYPE_CHECKING:
|
||||
CommonKVSender,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
if is_npu():
|
||||
from sglang.srt.hardware_backend.npu.dsv4.dsv4_memory_pool import (
|
||||
@@ -168,14 +170,14 @@ def _poll_with_failure_injection(pollers) -> List[int]:
|
||||
return [int(poller.poll()) for poller in pollers]
|
||||
|
||||
|
||||
def _is_fake_transfer(req: Req, server_args: ServerArgs) -> bool:
|
||||
def _is_fake_transfer(req: Req) -> bool:
|
||||
return req.bootstrap_host == FAKE_BOOTSTRAP_HOST or (
|
||||
req.bootstrap_host is None
|
||||
and server_args.disaggregation_transfer_backend == "fake"
|
||||
and get_disagg().disaggregation_transfer_backend == "fake"
|
||||
)
|
||||
|
||||
|
||||
def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> None:
|
||||
def _apply_metadata_gate(polls, decode_reqs, metadata_buffers) -> None:
|
||||
"""Downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||
|
||||
Mutates `polls` in-place. Called before all-reduce so that MIN across TP
|
||||
@@ -184,7 +186,7 @@ def _apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args) -> N
|
||||
for i, poll_val in enumerate(polls):
|
||||
if poll_val == int(KVPoll.Success):
|
||||
decode_req = decode_reqs[i]
|
||||
if _is_fake_transfer(decode_req.req, server_args):
|
||||
if _is_fake_transfer(decode_req.req):
|
||||
continue
|
||||
actual_room = metadata_buffers.bootstrap_room[
|
||||
decode_req.metadata_buffer_index, 0
|
||||
@@ -205,18 +207,13 @@ def poll_and_all_reduce(
|
||||
gloo_group: dist.ProcessGroup,
|
||||
decode_reqs=None,
|
||||
metadata_buffers: Optional[MetadataBuffers] = None,
|
||||
server_args: Optional[ServerArgs] = None,
|
||||
):
|
||||
# at a certain prob, the poll is failed to simulate failure
|
||||
polls = _poll_with_failure_injection(pollers)
|
||||
|
||||
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||
if (
|
||||
decode_reqs is not None
|
||||
and metadata_buffers is not None
|
||||
and server_args is not None
|
||||
):
|
||||
_apply_metadata_gate(polls, decode_reqs, metadata_buffers, server_args)
|
||||
if decode_reqs is not None and metadata_buffers is not None:
|
||||
_apply_metadata_gate(polls, decode_reqs, metadata_buffers)
|
||||
return _all_reduce_polls(polls, gloo_group)
|
||||
|
||||
|
||||
@@ -239,7 +236,6 @@ def poll_and_all_reduce_with_staging(
|
||||
staging_handler,
|
||||
gloo_group: dist.ProcessGroup,
|
||||
metadata_buffers: Optional[MetadataBuffers] = None,
|
||||
server_args: Optional[ServerArgs] = None,
|
||||
):
|
||||
"""Staging-aware polling: advance scatter, demote incomplete transfers, all_reduce."""
|
||||
for decode_req in decode_reqs:
|
||||
@@ -265,8 +261,8 @@ def poll_and_all_reduce_with_staging(
|
||||
):
|
||||
raw_polls[i] = int(KVPoll.Transferring)
|
||||
# Apply metadata gate on the decode requests to downgrade Success → Transferring for requests whose metadata hasn't landed.
|
||||
if metadata_buffers is not None and server_args is not None:
|
||||
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers, server_args)
|
||||
if metadata_buffers is not None:
|
||||
_apply_metadata_gate(raw_polls, decode_reqs, metadata_buffers)
|
||||
return _all_reduce_polls(raw_polls, gloo_group)
|
||||
|
||||
|
||||
|
||||
@@ -27,7 +27,10 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import initialize_dp_attention
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
@@ -111,7 +114,7 @@ def init_torch_distributed(
|
||||
|
||||
# Pre-warm NCCL/RCCL/HCCL to eliminate cold-start latency in first request
|
||||
# Controlled by --pre-warm-nccl flag (default: enabled on AMD GPUs)
|
||||
if server_args.pre_warm_nccl and (
|
||||
if get_exec().comm.pre_warm_nccl and (
|
||||
ps.tp_size > 1 or ps.pp_size > 1 or ps.moe_ep_size > 1
|
||||
):
|
||||
_prewarm_nccl(
|
||||
@@ -185,11 +188,11 @@ def _resolve_dist_init_method(*, server_args: ServerArgs, dist_port: int) -> str
|
||||
|
||||
|
||||
def _set_all_reduce_flags(*, server_args: ServerArgs) -> None:
|
||||
set_custom_all_reduce(not server_args.disable_custom_all_reduce)
|
||||
set_custom_all_reduce(not get_exec().comm.disable_custom_all_reduce)
|
||||
set_mscclpp_all_reduce(server_args.enable_mscclpp)
|
||||
set_torch_symm_mem_all_reduce(server_args.enable_torch_symm_mem)
|
||||
set_torch_symm_mem_all_reduce(get_exec().comm.enable_torch_symm_mem)
|
||||
set_flashinfer_allreduce_only(
|
||||
server_args.flashinfer_allreduce_fusion_backend is not None
|
||||
get_exec().comm.flashinfer_allreduce_fusion_backend is not None
|
||||
)
|
||||
|
||||
|
||||
@@ -242,7 +245,7 @@ def _init_parallel_groups(
|
||||
local_rank=gpu_id,
|
||||
distributed_init_method=dist_init_method,
|
||||
timeout=server_args.dist_timeout,
|
||||
moe_a2a_backend=server_args.moe_a2a_backend,
|
||||
moe_a2a_backend=get_exec().moe.moe_a2a_backend,
|
||||
recovered_rank=is_ep_joiner,
|
||||
max_world_size=server_args.max_ep_size,
|
||||
)
|
||||
@@ -260,7 +263,7 @@ def _init_parallel_groups(
|
||||
and server_args.enable_two_batch_overlap
|
||||
and get_parallel().enable_dsa_prefill_context_parallel
|
||||
),
|
||||
enable_symm_mem=server_args.enable_symm_mem,
|
||||
enable_symm_mem=get_exec().comm.enable_symm_mem,
|
||||
recovered_rank=is_ep_joiner,
|
||||
rank_offset=rank_offset,
|
||||
max_world_size=server_args.max_ep_size,
|
||||
|
||||
@@ -16,6 +16,10 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
|
||||
from pydantic import ValidationError
|
||||
|
||||
from sglang.srt.configs.embedding_model_spec import resolved_embedding_plan
|
||||
from sglang.srt.runtime_context import (
|
||||
get_lora,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -400,7 +404,7 @@ class RuntimeHandle:
|
||||
result = {
|
||||
"model_path": self.tokenizer_manager.model_path,
|
||||
"served_model_name": self.tokenizer_manager.served_model_name,
|
||||
"tokenizer_path": self.tokenizer_manager.server_args.tokenizer_path,
|
||||
"tokenizer_path": get_serving().tokenizer_path,
|
||||
"is_generation": self.tokenizer_manager.is_generation,
|
||||
"weight_version": self.tokenizer_manager.config_value("weight_version"),
|
||||
"load_format": self.tokenizer_manager.config_value("load_format"),
|
||||
@@ -461,9 +465,7 @@ class RuntimeHandle:
|
||||
"max_model_len": self.tokenizer_manager.model_config.context_len,
|
||||
}
|
||||
]
|
||||
if self.tokenizer_manager.server_args.enable_lora and hasattr(
|
||||
self.tokenizer_manager, "lora_registry"
|
||||
):
|
||||
if get_lora().enable_lora and hasattr(self.tokenizer_manager, "lora_registry"):
|
||||
lora_registry = self.tokenizer_manager.lora_registry
|
||||
for _, lora_ref in lora_registry.get_all_adapters().items():
|
||||
models.append(
|
||||
|
||||
@@ -394,7 +394,7 @@ async def lifespan(fast_api_app: FastAPI):
|
||||
if (
|
||||
getattr(fast_api_app, "is_single_tokenizer_mode", False)
|
||||
and get_serving().grpc_port is not None
|
||||
and not (server_args.smg_grpc_mode or server_args.grpc_mode)
|
||||
and not (get_serving().smg_grpc_mode or server_args.grpc_mode)
|
||||
):
|
||||
grpc_handle = _start_native_grpc_server_for_runtime(
|
||||
server_args=server_args,
|
||||
@@ -484,6 +484,7 @@ from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_lora,
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_serving,
|
||||
@@ -744,9 +745,9 @@ async def model_info():
|
||||
# Manager-owned, and moved by a weight update alongside `model_path`:
|
||||
# this is where a client reads the identity the server answers under.
|
||||
"served_model_name": _global_state.tokenizer_manager.served_model_name,
|
||||
"tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path,
|
||||
"tokenizer_path": get_serving().tokenizer_path,
|
||||
"is_generation": _global_state.tokenizer_manager.is_generation,
|
||||
"preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params,
|
||||
"preferred_sampling_params": get_serving().preferred_sampling_params,
|
||||
"weight_version": _global_state.tokenizer_manager.config_value(
|
||||
"weight_version"
|
||||
),
|
||||
@@ -1857,7 +1858,7 @@ async def available_models():
|
||||
)
|
||||
|
||||
# Add loaded LoRA adapters
|
||||
if _global_state.tokenizer_manager.server_args.enable_lora:
|
||||
if get_lora().enable_lora:
|
||||
lora_registry = _global_state.tokenizer_manager.lora_registry
|
||||
for _, lora_ref in lora_registry.get_all_adapters().items():
|
||||
model_cards.append(
|
||||
|
||||
@@ -18,6 +18,9 @@ import logging
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
|
||||
from sglang.srt.runtime_context import (
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.utils.common import kill_itself_when_parent_died, kill_process_tree
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
from sglang.srt.utils.watchdog import SubprocessWatchdog
|
||||
@@ -36,12 +39,10 @@ def _loopback_host(host: str) -> str:
|
||||
return host
|
||||
|
||||
|
||||
def build_sidecar_endpoint(server_args) -> str:
|
||||
"""Both halves of the endpoint come from the argument: this is a helper
|
||||
over a config object, callable before anything is published."""
|
||||
return NetworkAddress(
|
||||
_loopback_host(server_args.host), server_args.grpc_port
|
||||
).to_url()
|
||||
def build_sidecar_endpoint(host: str, grpc_port: int) -> str:
|
||||
"""Both halves are passed in: this is a string helper, and the caller is
|
||||
the one that knows where the effective values live."""
|
||||
return NetworkAddress(_loopback_host(host), grpc_port).to_url()
|
||||
|
||||
|
||||
def _parse_sidecar_args(args: list[str] | None) -> tuple[list[str], float]:
|
||||
@@ -117,7 +118,7 @@ def start_sidecar(server_args) -> Sidecar:
|
||||
module_name = server_args.sidecar
|
||||
assert module_name is not None
|
||||
sidecar_args, shutdown_timeout = _parse_sidecar_args(server_args.sidecar_args)
|
||||
endpoint = build_sidecar_endpoint(server_args)
|
||||
endpoint = build_sidecar_endpoint(server_args.host, get_serving().grpc_port)
|
||||
proc = mp.get_context("spawn").Process(
|
||||
name=f"sglang_sidecar_{module_name}",
|
||||
target=_run_sidecar,
|
||||
|
||||
@@ -26,6 +26,10 @@ from typing import Optional
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi.responses import Response
|
||||
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.utils import get_device_name
|
||||
from sglang.version import __version__
|
||||
|
||||
@@ -144,9 +148,9 @@ async def get_loads(
|
||||
"accelerator": _accelerator_name(),
|
||||
"num_accelerators": _num_accelerators_per_dp_rank(
|
||||
tokenizer_manager.server_args.tp_size,
|
||||
tokenizer_manager.server_args.pp_size,
|
||||
tokenizer_manager.server_args.dp_size,
|
||||
tokenizer_manager.server_args.enable_dp_attention,
|
||||
configured_pp_size(),
|
||||
get_parallel().dp_size,
|
||||
get_parallel().enable_dp_attention,
|
||||
),
|
||||
"loads": loads,
|
||||
}
|
||||
|
||||
@@ -1,6 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_spec,
|
||||
)
|
||||
|
||||
"""
|
||||
end to end attention solution with aiter kernels
|
||||
@@ -152,7 +156,7 @@ class AiterAttnBackend(AttentionBackend):
|
||||
|
||||
self.input_dtype = model_runner.model_config.dtype
|
||||
|
||||
self.page_size = model_runner.server_args.page_size
|
||||
self.page_size = get_schedule().page_size
|
||||
|
||||
self.extend_attention_fwd = torch.compiler.disable(extend_attention_fwd)
|
||||
|
||||
@@ -2957,7 +2961,7 @@ class AiterMultiStepDraftBackend:
|
||||
# Cached variables for generate_draft_decode_kv_indices
|
||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
||||
self.page_size = model_runner.server_args.page_size
|
||||
self.page_size = get_schedule().page_size
|
||||
|
||||
def common_template(
|
||||
self, forward_batch: ForwardBatch, kv_indices_buffer: torch.Tensor, call_fn: int
|
||||
|
||||
@@ -13,7 +13,10 @@ from sglang.srt.configs.linear_attn_model_registry import (
|
||||
get_linear_attn_config,
|
||||
import_backend_class,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import get_device_capability, is_hip, is_musa, is_npu
|
||||
|
||||
_is_musa = is_musa()
|
||||
@@ -49,7 +52,7 @@ def create_flashinfer_backend(runner):
|
||||
)
|
||||
|
||||
# Init streams
|
||||
if runner.server_args.speculative_algorithm == "EAGLE":
|
||||
if get_spec().speculative_algorithm == "EAGLE":
|
||||
if (
|
||||
not hasattr(runner, "plan_stream_for_flashinfer")
|
||||
or not runner.plan_stream_for_flashinfer
|
||||
@@ -70,10 +73,7 @@ def create_flashinfer_backend(runner):
|
||||
def create_trtllm_mla_backend(runner):
|
||||
if not runner.use_mla_backend:
|
||||
raise ValueError("trtllm_mla backend can only be used with MLA models.")
|
||||
if (
|
||||
get_parallel().dcp_enabled
|
||||
and runner.server_args.speculative_algorithm is not None
|
||||
):
|
||||
if get_parallel().dcp_enabled and get_spec().speculative_algorithm is not None:
|
||||
_, decode_backend = runner.server_args.get_attention_backends()
|
||||
if decode_backend == "trtllm_mla":
|
||||
raise ValueError(
|
||||
@@ -265,7 +265,7 @@ def create_hpc_ops_backend(runner):
|
||||
raise ValueError(
|
||||
"Cross attention is not supported in the hpc_ops attention backend."
|
||||
)
|
||||
if runner.server_args.speculative_algorithm is not None:
|
||||
if get_spec().speculative_algorithm is not None:
|
||||
raise ValueError(
|
||||
"hpc_ops backend does not support speculative decoding for now."
|
||||
)
|
||||
|
||||
@@ -69,7 +69,10 @@ from sglang.srt.layers.attention.verify_mask import (
|
||||
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
||||
from sglang.srt.speculative.ragged_verify import (
|
||||
RaggedVerifyMode,
|
||||
@@ -576,7 +579,7 @@ class DeepseekV4AttnBackend(
|
||||
self._q8kv8_qpad_buf = None
|
||||
self._q8kv8_attn_sink_pad = None
|
||||
self._q8kv8_identity_scale = None
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||
self.topk = get_spec().speculative_eagle_topk or 0
|
||||
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
||||
self.mtp_enabled = self.topk > 0
|
||||
self.speculative_num_steps = speculative_num_steps
|
||||
|
||||
@@ -36,7 +36,10 @@ from sglang.srt.layers.attention.dsv4.metadata import (
|
||||
)
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
||||
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
||||
from sglang.srt.utils import ceil_align
|
||||
@@ -456,7 +459,7 @@ class DeepseekV4HipRadixBackend(
|
||||
self.enable_deepseek_v4_fp4_indexer: bool = (
|
||||
model_runner.server_args.enable_deepseek_v4_fp4_indexer
|
||||
)
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||
self.topk = get_spec().speculative_eagle_topk or 0
|
||||
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
||||
self.mtp_enabled = self.topk > 0
|
||||
self.speculative_num_steps = speculative_num_steps
|
||||
|
||||
@@ -15,7 +15,11 @@ from typing import (
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
|
||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
|
||||
@@ -309,7 +313,7 @@ class DeepseekSparseAttnBackend(
|
||||
assert isinstance(model_runner.page_size, int)
|
||||
self.real_page_size = model_runner.page_size
|
||||
self.num_splits = (
|
||||
1 if model_runner.server_args.enable_deterministic_inference else 0
|
||||
1 if get_exec().deterministic.enable_deterministic_inference else 0
|
||||
)
|
||||
self.use_dsa = is_deepseek_dsa(model_runner.model_config.hf_config)
|
||||
assert self.use_dsa, "DSA backend only supports DeepSeek DSA"
|
||||
@@ -334,10 +338,8 @@ class DeepseekSparseAttnBackend(
|
||||
|
||||
self.use_mha: bool = False
|
||||
self.supports_mha_one_shot: bool = True
|
||||
self.dsa_prefill_impl: _DSA_IMPL_T = (
|
||||
model_runner.server_args.dsa_prefill_backend
|
||||
)
|
||||
self.dsa_decode_impl: _DSA_IMPL_T = model_runner.server_args.dsa_decode_backend
|
||||
self.dsa_prefill_impl: _DSA_IMPL_T = get_exec().kernel.dsa_prefill_backend
|
||||
self.dsa_decode_impl: _DSA_IMPL_T = get_exec().kernel.dsa_decode_backend
|
||||
self.dsa_topk_backend: DSATopKBackend = DSATopKBackend(
|
||||
model_runner.server_args.dsa_topk_backend
|
||||
)
|
||||
@@ -390,7 +392,7 @@ class DeepseekSparseAttnBackend(
|
||||
)
|
||||
|
||||
# Speculative decoding
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||
self.topk = get_spec().speculative_eagle_topk or 0
|
||||
self.speculative_num_steps = speculative_num_steps
|
||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||
self.speculative_step_id = speculative_step_id
|
||||
|
||||
@@ -1,6 +1,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
)
|
||||
|
||||
"""
|
||||
Support different attention backends.
|
||||
@@ -403,7 +406,7 @@ class FlashInferAttnBackend(AttentionBackend):
|
||||
# Also set split tile sizes for prefill and decode from environment variables, and disable kv split for cuda graph
|
||||
# More information can be found here: https://github.com/flashinfer-ai/flashinfer/pull/1675
|
||||
self.enable_deterministic = (
|
||||
model_runner.server_args.enable_deterministic_inference
|
||||
get_exec().deterministic.enable_deterministic_inference
|
||||
)
|
||||
self.prefill_split_tile_size = None
|
||||
self.decode_split_tile_size = None
|
||||
|
||||
@@ -1,6 +1,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sglang.srt.runtime_context import get_disagg, get_exec, get_parallel, get_schedule
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
)
|
||||
|
||||
"""
|
||||
Support attention backend for flashinfer MLA.
|
||||
@@ -1101,7 +1106,7 @@ class FlashInferMLAMultiStepDraftBackend:
|
||||
# Cached variables for generate_draft_decode_kv_indices
|
||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
||||
self.page_size = model_runner.server_args.page_size
|
||||
self.page_size = get_schedule().page_size
|
||||
|
||||
def common_template(
|
||||
self,
|
||||
|
||||
@@ -12,6 +12,9 @@ from sglang.srt.layers.attention.dsa.dsa_indexer_metadata import BaseIndexerMeta
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.runtime_context import (
|
||||
get_spec,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.verify_mask import VerifyMask
|
||||
@@ -33,12 +36,8 @@ class HybridAttnBackend(AttentionBackend):
|
||||
self.data_type = model_runner.kv_cache_dtype
|
||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||
self.spec_attn_is_decode = (
|
||||
model_runner.server_args.speculative_attention_mode == "decode"
|
||||
)
|
||||
self.spec_attn_is_prefill = (
|
||||
model_runner.server_args.speculative_attention_mode == "prefill"
|
||||
)
|
||||
self.spec_attn_is_decode = get_spec().speculative_attention_mode == "decode"
|
||||
self.spec_attn_is_prefill = get_spec().speculative_attention_mode == "prefill"
|
||||
# Gates the FutureMap's per-step seq_lens D2H (decide_needs_cpu_seq_lens
|
||||
# ORs it across backends). Count only what runs in the spec decode loop:
|
||||
# decode always, prefill only when mode=prefill routes verify to it --
|
||||
@@ -121,10 +120,7 @@ class HybridAttnBackend(AttentionBackend):
|
||||
|
||||
def init_cuda_graph_state(self, max_bs: int, max_num_tokens: int):
|
||||
self.decode_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||
if (
|
||||
self.model_runner.server_args.speculative_algorithm is not None
|
||||
and self.spec_attn_is_prefill
|
||||
):
|
||||
if get_spec().speculative_algorithm is not None and self.spec_attn_is_prefill:
|
||||
# When speculative decoding is enabled, we need to initialize the backend
|
||||
# that will be used for target_verify.
|
||||
self.prefill_backend.init_cuda_graph_state(max_bs, max_num_tokens)
|
||||
|
||||
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_spec,
|
||||
mamba_cache_chunk_size,
|
||||
)
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
|
||||
@@ -47,7 +48,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
super().__init__()
|
||||
self.pad_slot_id = PAD_SLOT_ID
|
||||
self.device = model_runner.device
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||
self.topk = get_spec().speculative_eagle_topk or 0
|
||||
self.is_draft_worker = model_runner.is_draft_worker
|
||||
self.req_to_token_pool: HybridReqToTokenPool = model_runner.req_to_token_pool
|
||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
|
||||
@@ -31,6 +31,9 @@ elif is_cpu():
|
||||
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.runtime_context import (
|
||||
get_spec,
|
||||
)
|
||||
|
||||
|
||||
class KDAKernelDispatcher:
|
||||
@@ -384,7 +387,7 @@ class KDAAttnBackend(MambaAttnBackendBase):
|
||||
# traversal). Reject EAGLE tree verify (topk > 1) early at setup, keyed on the
|
||||
# verify backend (not decode). The kernel keeps a per-call
|
||||
# retrieve_parent_token backstop that also covers ngram tree.
|
||||
speculative_topk = model_runner.server_args.speculative_eagle_topk or 1
|
||||
speculative_topk = get_spec().speculative_eagle_topk or 1
|
||||
if verify_backend.is_flashinfer() and speculative_topk > 1:
|
||||
raise ValueError(
|
||||
"KDA FlashInfer speculative decoding only supports topk=1 "
|
||||
|
||||
@@ -38,7 +38,12 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
cuda_graph_fully_disabled,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.spec_utils import (
|
||||
draft_kv_indices_buffer_width,
|
||||
draft_kv_indices_used_len,
|
||||
@@ -254,7 +259,7 @@ class TritonAttnBackend(AttentionBackend):
|
||||
self.static_kv_splits = get_bool_env_var(
|
||||
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
|
||||
)
|
||||
self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits
|
||||
self.max_kv_splits = get_exec().kernel.triton_attention_num_kv_splits
|
||||
if self.use_mla and not _is_xpu:
|
||||
self.max_kv_splits = _mla_decode_kv_splits_cap(
|
||||
self.max_kv_splits,
|
||||
@@ -280,11 +285,11 @@ class TritonAttnBackend(AttentionBackend):
|
||||
cuda_graph_fully_disabled()
|
||||
or check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
|
||||
)
|
||||
and model_runner.server_args.chunked_prefill_size == -1
|
||||
and get_schedule().chunked_prefill_size == -1
|
||||
)
|
||||
|
||||
self.enable_deterministic = (
|
||||
model_runner.server_args.enable_deterministic_inference
|
||||
get_exec().deterministic.enable_deterministic_inference
|
||||
)
|
||||
|
||||
if self.enable_deterministic:
|
||||
@@ -1987,7 +1992,7 @@ class TritonMultiStepDraftBackend:
|
||||
# Cached variables for generate_draft_decode_kv_indices
|
||||
self.req_to_token_pool = model_runner.req_to_token_pool
|
||||
self.pool_len = model_runner.req_to_token_pool.req_to_token.shape[1]
|
||||
self.page_size = model_runner.server_args.page_size
|
||||
self.page_size = get_schedule().page_size
|
||||
|
||||
def common_template(
|
||||
self,
|
||||
|
||||
@@ -13,7 +13,11 @@ from sglang.kernels.ops.kvcache.kv_indices import (
|
||||
)
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import get_bool_env_var, get_device_core_count
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -109,7 +113,7 @@ class WaveAttnBackend(AttentionBackend):
|
||||
self.static_kv_splits = get_bool_env_var(
|
||||
"SGLANG_TRITON_DECODE_ATTN_STATIC_KV_SPLITS", "false"
|
||||
)
|
||||
self.max_kv_splits = model_runner.server_args.triton_attention_num_kv_splits
|
||||
self.max_kv_splits = get_exec().kernel.triton_attention_num_kv_splits
|
||||
self.v_head_dim = model_runner.token_to_kv_pool.get_value_buffer(0).shape[-1]
|
||||
|
||||
self.forward_metadata: ForwardMetadata = None
|
||||
|
||||
@@ -15,7 +15,10 @@ from sglang.srt.layers.attention.flashattention_backend import (
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||
from sglang.srt.runtime_context import get_schedule, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_schedule,
|
||||
get_spec,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
@@ -78,7 +81,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
isinstance(model_runner.token_to_kv_pool, SWAKVPool)
|
||||
and model_runner.token_to_kv_pool.swa_layer_nums > 0
|
||||
)
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||
self.topk = get_spec().speculative_eagle_topk or 0
|
||||
self.speculative_num_steps = speculative_num_steps
|
||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||
self.speculative_step_id = speculative_step_id
|
||||
|
||||
@@ -16,7 +16,11 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.deep_gemm_wrapper.configurer import ENABLE_JIT_DEEPGEMM
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import ceil_align, ceil_div, get_available_gpu_memory, is_musa
|
||||
|
||||
@@ -472,7 +476,7 @@ def pp_parallel_deep_gemm_warmup(runner) -> None:
|
||||
|
||||
# In PD, prefill-only nodes never decode (indexer would OOM at large
|
||||
# bs) and decode-only nodes never extend.
|
||||
disagg_mode = model_runner.server_args.disaggregation_mode
|
||||
disagg_mode = get_disagg().disaggregation_mode
|
||||
run_decode = model_runner.is_generation and disagg_mode != "prefill"
|
||||
run_extend = disagg_mode != "decode"
|
||||
|
||||
|
||||
@@ -13,7 +13,10 @@ from typing import TYPE_CHECKING, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.layers.quantization.base_config import FusedMoEMethodBase
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
)
|
||||
from sglang.srt.utils import get_compiler_backend
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -88,7 +91,7 @@ def create_kt_config_from_server_args(
|
||||
cpuinfer_threads=server_args.kt_cpuinfer,
|
||||
threadpool_count=server_args.kt_threadpool_count,
|
||||
weight_path=server_args.kt_weight_path,
|
||||
chunked_prefill_size=server_args.chunked_prefill_size,
|
||||
chunked_prefill_size=get_schedule().chunked_prefill_size,
|
||||
method=server_args.kt_method,
|
||||
max_deferred_experts_per_token=server_args.kt_max_deferred_experts_per_token,
|
||||
num_layers=num_layers,
|
||||
|
||||
@@ -31,7 +31,10 @@ from sglang.srt.layers.quantization.base_config import (
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from sglang.srt.layers.utils import copy_or_rebind_param
|
||||
from sglang.srt.runtime_context import get_exec, get_lora
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_lora,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -100,7 +103,9 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
|
||||
backend_str = server_args.bf16_gemm_backend
|
||||
if backend_str == "auto" and is_sm100_supported():
|
||||
backend_str = (
|
||||
"torch" if server_args.enable_deterministic_inference else "cutedsl"
|
||||
"torch"
|
||||
if get_exec().deterministic.enable_deterministic_inference
|
||||
else "cutedsl"
|
||||
)
|
||||
|
||||
backend = Bf16GemmBackend(backend_str)
|
||||
@@ -118,7 +123,7 @@ def initialize_bf16_gemm_config(server_args: ServerArgs) -> None:
|
||||
_hopper_bf16_gemv = hopper_bf16_gemv
|
||||
_use_hopper_bf16_gemv = use_hopper_bf16_gemv
|
||||
elif backend.is_cutedsl():
|
||||
if server_args.enable_deterministic_inference:
|
||||
if get_exec().deterministic.enable_deterministic_inference:
|
||||
raise ValueError(
|
||||
"--bf16-gemm-backend cutedsl is batch-size dependent and cannot "
|
||||
"be combined with --enable-deterministic-inference"
|
||||
|
||||
@@ -20,8 +20,8 @@ def validate_experimental_sgl_marlin_server_args(
|
||||
|
||||
# A provided adapter path implicitly enables LoRA later unless it was
|
||||
# explicitly disabled. No-LoRA delegates to the stock Marlin fused path.
|
||||
lora_enabled = bool(server_args.enable_lora) or (
|
||||
server_args.enable_lora is None and bool(server_args.lora_paths)
|
||||
lora_enabled = bool(resolved_args.enable_lora) or (
|
||||
resolved_args.enable_lora is None and bool(server_args.lora_paths)
|
||||
)
|
||||
if not lora_enabled:
|
||||
return
|
||||
|
||||
@@ -61,6 +61,9 @@ from sglang.srt.managers.load_snapshot import (
|
||||
zmq_reader_owner,
|
||||
)
|
||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
)
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
configure_logger,
|
||||
@@ -667,7 +670,7 @@ class TokenizerWorker(TokenizerManager):
|
||||
|
||||
# For PD disaggregation
|
||||
self.disaggregation_transfer_backend = TransferBackend(
|
||||
self.server_args.disaggregation_transfer_backend
|
||||
get_disagg().disaggregation_transfer_backend
|
||||
)
|
||||
|
||||
# Register this worker with the router for pause/continue broadcasting
|
||||
|
||||
@@ -8,7 +8,10 @@ from typing import TYPE_CHECKING, NamedTuple, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
)
|
||||
from sglang.srt.utils import get_bool_env_var
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -112,7 +115,7 @@ class PrefillDelayer:
|
||||
# env flag is on (or overlap scheduling is disabled), ride the NCCL
|
||||
# device group on `device` instead of gloo on CPU.
|
||||
use_nccl = (
|
||||
server_args.disable_overlap_schedule
|
||||
get_schedule().disable_overlap_schedule
|
||||
or envs.SGLANG_NCCL_ALL_GATHER_IN_OVERLAP_SCHEDULER_SYNC_BATCH.get()
|
||||
)
|
||||
if use_nccl:
|
||||
@@ -140,7 +143,7 @@ class PrefillDelayer:
|
||||
self.skip_first_delayer = True
|
||||
|
||||
assert (
|
||||
not server_args.disable_overlap_schedule
|
||||
not get_schedule().disable_overlap_schedule
|
||||
), "To use PrefillDelayer, disable_overlap_schedule must be False."
|
||||
|
||||
def _negotiate_should_allow_prefill(
|
||||
|
||||
@@ -27,6 +27,10 @@ from sglang.srt.managers.utils import (
|
||||
compute_num_reserved_tokens,
|
||||
msgpack_decode_explained,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_mm,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.utils.flatten import (
|
||||
FlatPairColumns,
|
||||
NestedRowColumns,
|
||||
@@ -216,9 +220,7 @@ class NativeMmHost:
|
||||
|
||||
# `--mm-process-config {"image": {...}}`: only pixel-limit overrides are
|
||||
# mirrored natively, anything else disables the pipeline.
|
||||
image_overrides = dict(
|
||||
(self.server_args.mm_process_config or {}).get("image", {})
|
||||
)
|
||||
image_overrides = dict((get_mm().mm_process_config or {}).get("image", {}))
|
||||
if not set(image_overrides) <= {"min_pixels", "max_pixels"}:
|
||||
return None
|
||||
|
||||
@@ -386,7 +388,7 @@ class RustServer:
|
||||
# Refuse rather than run: silently dropping it means generating with
|
||||
# sampling the operator did not configure, and `/get_model_info` would go on
|
||||
# advertising values no request ever receives.
|
||||
if server_args.preferred_sampling_params:
|
||||
if get_serving().preferred_sampling_params:
|
||||
raise ValueError(
|
||||
"SGLANG_RUST_SERVER does not yet apply --preferred-sampling-params "
|
||||
"(the Python TokenizerManager merges it into every request; the rust "
|
||||
|
||||
@@ -22,7 +22,13 @@ from sglang.srt.observability.metrics_collector import (
|
||||
SchedulerStats,
|
||||
compute_routing_key_stats,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_context, get_observability, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_context,
|
||||
get_disagg,
|
||||
get_observability,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils.device_timer import DeviceTimer
|
||||
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
||||
|
||||
@@ -621,9 +627,8 @@ class SchedulerMetricsReporter:
|
||||
msg += f"#optimistic-req: {num_optimistic}, "
|
||||
|
||||
if (
|
||||
self.scheduler.server_args.language_only
|
||||
and self.scheduler.server_args.encoder_transfer_backend
|
||||
== "zmq_to_scheduler"
|
||||
get_disagg().language_only
|
||||
and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
|
||||
):
|
||||
msg += (
|
||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||
@@ -875,9 +880,8 @@ class SchedulerMetricsReporter:
|
||||
msg += f"#retracted-req: {len(self.scheduler.disagg_decode_prealloc_queue.retracted_queue)}, "
|
||||
|
||||
if (
|
||||
self.scheduler.server_args.language_only
|
||||
and self.scheduler.server_args.encoder_transfer_backend
|
||||
== "zmq_to_scheduler"
|
||||
get_disagg().language_only
|
||||
and get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
|
||||
):
|
||||
msg += (
|
||||
f"waiting-image-req: {len(self.scheduler.mm_receiver.waiting_list)}, "
|
||||
@@ -1113,7 +1117,7 @@ class SchedulerMetricsReporter:
|
||||
active_lora_ids = set()
|
||||
|
||||
# For PP mode, check all running micro batches
|
||||
if self.scheduler.server_args.pp_size > 1:
|
||||
if configured_pp_size() > 1:
|
||||
for batch in self.scheduler.running_mbs:
|
||||
if batch and hasattr(batch, "reqs"):
|
||||
for req in batch.reqs:
|
||||
|
||||
@@ -74,7 +74,11 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightsFromTensorReqOutput,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_lora,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -186,7 +190,7 @@ class TokenizerControlMixin:
|
||||
self: TokenizerManager, obj: AddExternalCorpusReqInput
|
||||
) -> AddExternalCorpusReqOutput:
|
||||
self.auto_create_handle_loop()
|
||||
if self.server_args.speculative_algorithm != "NGRAM":
|
||||
if get_spec().speculative_algorithm != "NGRAM":
|
||||
return AddExternalCorpusReqOutput(
|
||||
success=False,
|
||||
message="Ngram speculative decoding is not enabled.",
|
||||
@@ -262,7 +266,7 @@ class TokenizerControlMixin:
|
||||
self: TokenizerManager, corpus_id: str
|
||||
) -> RemoveExternalCorpusReqOutput:
|
||||
self.auto_create_handle_loop()
|
||||
if self.server_args.speculative_algorithm != "NGRAM":
|
||||
if get_spec().speculative_algorithm != "NGRAM":
|
||||
return RemoveExternalCorpusReqOutput(
|
||||
success=False,
|
||||
message="Ngram speculative decoding is not enabled.",
|
||||
@@ -277,7 +281,7 @@ class TokenizerControlMixin:
|
||||
self: TokenizerManager,
|
||||
) -> ListExternalCorporaReqOutput:
|
||||
self.auto_create_handle_loop()
|
||||
if self.server_args.speculative_algorithm != "NGRAM":
|
||||
if get_spec().speculative_algorithm != "NGRAM":
|
||||
return ListExternalCorporaReqOutput(
|
||||
success=False,
|
||||
message="Ngram speculative decoding is not enabled.",
|
||||
@@ -602,7 +606,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not self.server_args.enable_lora:
|
||||
if not get_lora().enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
@@ -680,7 +684,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not self.server_args.enable_lora:
|
||||
if not get_lora().enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
@@ -756,7 +760,7 @@ class TokenizerControlMixin:
|
||||
self.auto_create_handle_loop()
|
||||
|
||||
try:
|
||||
if not self.server_args.enable_lora:
|
||||
if not get_lora().enable_lora:
|
||||
raise ValueError(
|
||||
"LoRA is not enabled. Please set `--enable-lora` to enable LoRA."
|
||||
)
|
||||
|
||||
@@ -121,7 +121,7 @@ def _register_legacy_hicache_draft(
|
||||
host_to_device_ratio=primary_host_pool.logical_size / pool.size,
|
||||
host_size=0,
|
||||
page_size=page_size,
|
||||
layout=server_args.hicache_mem_layout,
|
||||
layout=get_memory().hicache_mem_layout,
|
||||
allocator_type=server_args.hicache_storage_backend,
|
||||
pool_label="draft",
|
||||
)
|
||||
|
||||
@@ -70,6 +70,7 @@ from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_mm,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_spec,
|
||||
@@ -1253,7 +1254,7 @@ class KVCacheConfigurator:
|
||||
|
||||
token_to_kv_pool = NPUMiniMaxSparseKVPool(
|
||||
size=max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
index_dtype=self.model_dtype,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
@@ -1846,7 +1847,7 @@ class KVCacheConfigurator:
|
||||
)
|
||||
mm_reservation_gb = mm_runtime_reservation_gb(
|
||||
is_multimodal=self.model_config.is_multimodal,
|
||||
mm_feature_transport=self.server_args.mm_feature_transport,
|
||||
mm_feature_transport=get_mm().mm_feature_transport,
|
||||
)
|
||||
rest_memory = available_gpu_memory - slack_gb - mm_reservation_gb
|
||||
if self.mambaish_config is not None:
|
||||
@@ -2219,16 +2220,16 @@ def calculate_mla_kv_cache_dim(
|
||||
# since it is not compatible for trtllm and other mla attn backend due to the different
|
||||
# kv cache layout.
|
||||
if (
|
||||
server_args.dsa_prefill_backend == "trtllm"
|
||||
or server_args.dsa_decode_backend == "trtllm"
|
||||
get_exec().kernel.dsa_prefill_backend == "trtllm"
|
||||
or get_exec().kernel.dsa_decode_backend == "trtllm"
|
||||
):
|
||||
return kv_cache_dim
|
||||
|
||||
# On HIP, TileLang and AITER DSA kernels consume the raw MLA KV layout:
|
||||
# nope(512 fp8) + rope(64 fp8), without extra per-block scales.
|
||||
if _is_hip and (
|
||||
server_args.dsa_prefill_backend in ("tilelang", "aiter")
|
||||
or server_args.dsa_decode_backend in ("tilelang", "aiter")
|
||||
get_exec().kernel.dsa_prefill_backend in ("tilelang", "aiter")
|
||||
or get_exec().kernel.dsa_decode_backend in ("tilelang", "aiter")
|
||||
):
|
||||
return kv_cache_dim
|
||||
|
||||
|
||||
@@ -214,7 +214,7 @@ def create_tree_cache(ctx: TreeCacheBuildContext) -> BasePrefixCache:
|
||||
source = "default"
|
||||
|
||||
if (
|
||||
ctx.server_args.enable_hierarchical_cache
|
||||
get_memory().enable_hierarchical_cache
|
||||
and ctx.server_args.hicache_host_memory_mode == "buffer_only"
|
||||
):
|
||||
from sglang.srt.mem_cache.unified_radix_cache import UnifiedRadixCache
|
||||
|
||||
@@ -39,7 +39,14 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
|
||||
from sglang.srt.runtime_context import get_exec, get_flags, get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_lora,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
empty_context,
|
||||
log_info_on_rank0,
|
||||
@@ -585,22 +592,20 @@ class CPUGraphRunner:
|
||||
self.enable_two_batch_overlap = (
|
||||
model_runner.server_args.enable_two_batch_overlap
|
||||
)
|
||||
self.speculative_algorithm = model_runner.server_args.speculative_algorithm
|
||||
self.speculative_algorithm = get_spec().speculative_algorithm
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
self.tp_size = model_runner.server_args.tp_size
|
||||
self.dp_size = get_parallel().dp_size
|
||||
self.pp_size = model_runner.server_args.pp_size
|
||||
self.pp_size = configured_pp_size()
|
||||
|
||||
self.capture_forward_mode = ForwardMode.DECODE
|
||||
self.capture_hidden_mode = self.return_hidden_states_mode
|
||||
# Static capture width: CPU graphs are decode-only.
|
||||
self.captured_req_width = 1
|
||||
|
||||
assert (
|
||||
not self.model_runner.server_args.enable_lora
|
||||
), "CPUGraphRunner does not support LoRA yet."
|
||||
assert not get_lora().enable_lora, "CPUGraphRunner does not support LoRA yet."
|
||||
assert (
|
||||
not self.enable_two_batch_overlap
|
||||
), "CPUGraphRunner does not support two batch overlap yet."
|
||||
@@ -991,7 +996,7 @@ class CPUGraphRunner:
|
||||
retrieve_next_sibling=None,
|
||||
retrieve_cum_len=None,
|
||||
spec_steps=get_spec().speculative_num_steps,
|
||||
topk=self.model_runner.server_args.speculative_eagle_topk,
|
||||
topk=get_spec().speculative_eagle_topk,
|
||||
draft_token_num=get_spec().speculative_num_draft_tokens,
|
||||
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||
seq_lens_sum=None,
|
||||
|
||||
@@ -51,7 +51,11 @@ from sglang.srt.layers.dp_attention import (
|
||||
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
||||
ForwardBatchDeepSeekMHAMixin,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_lora,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
is_cpu,
|
||||
is_cuda,
|
||||
@@ -922,7 +926,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
if model_runner.lora_manager is not None:
|
||||
# In the non-LoRA overlap loading case, we fetch LoRA adapters into the memory pool
|
||||
# as a batch, right before running the batch
|
||||
if not model_runner.server_args.enable_lora_overlap_loading:
|
||||
if not get_lora().enable_lora_overlap_loading:
|
||||
model_runner.lora_manager.fetch_new_loras(set(ret.lora_ids))
|
||||
|
||||
model_runner.lora_manager.prepare_lora_batch(ret)
|
||||
@@ -1320,7 +1324,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# graph; larger prefills fall back to eager and keep the
|
||||
# memory-efficient SUM_LEN. global_num_tokens is identical across ranks
|
||||
# (all-gathered), so the decision is consistent cluster-wide.
|
||||
prefill_cg = model_runner.server_args.cuda_graph_config.prefill
|
||||
prefill_cg = get_exec().graph.cuda_graph_config.prefill
|
||||
if (
|
||||
self.can_run_dp_breakable_cuda_graph
|
||||
and self.is_extend_in_batch
|
||||
|
||||
@@ -5,6 +5,9 @@ from typing import TYPE_CHECKING, Dict, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
@@ -29,7 +32,7 @@ class GraphSharedOutput:
|
||||
def create_for_model_runner(
|
||||
cls, model_runner: ModelRunner
|
||||
) -> Optional[GraphSharedOutput]:
|
||||
cuda_graph_config = model_runner.server_args.cuda_graph_config
|
||||
cuda_graph_config = get_exec().graph.cuda_graph_config
|
||||
if cuda_graph_config is None:
|
||||
return None
|
||||
|
||||
|
||||
@@ -41,7 +41,13 @@ from sglang.srt.model_executor.runner import (
|
||||
)
|
||||
from sglang.srt.model_loader.utils import resolve_language_model
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_flags
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_schedule,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import get_available_gpu_memory, log_info_on_rank0
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -111,16 +117,15 @@ def capture_cuda_graphs(
|
||||
|
||||
if model_runner.is_draft_worker:
|
||||
moe_runner_backend = (
|
||||
model_runner.server_args.speculative_moe_runner_backend
|
||||
or model_runner.server_args.moe_runner_backend
|
||||
get_spec().speculative_moe_runner_backend
|
||||
or get_exec().moe.moe_runner_backend
|
||||
)
|
||||
moe_a2a_backend = (
|
||||
model_runner.server_args.speculative_moe_a2a_backend
|
||||
or model_runner.server_args.moe_a2a_backend
|
||||
get_spec().speculative_moe_a2a_backend or get_exec().moe.moe_a2a_backend
|
||||
)
|
||||
else:
|
||||
moe_runner_backend = model_runner.server_args.moe_runner_backend
|
||||
moe_a2a_backend = model_runner.server_args.moe_a2a_backend
|
||||
moe_runner_backend = get_exec().moe.moe_runner_backend
|
||||
moe_a2a_backend = get_exec().moe.moe_a2a_backend
|
||||
|
||||
uses_deep_gemm_moe_runner = moe_runner_backend == "deep_gemm"
|
||||
if moe_runner_backend == "auto" and model_runner.model_config.quantization in (
|
||||
@@ -200,7 +205,7 @@ def capture_cuda_graphs(
|
||||
|
||||
prealloc_symmetric_memory_pool(
|
||||
is_draft_worker=model_runner.is_draft_worker,
|
||||
enable_symm_mem=model_runner.server_args.enable_symm_mem,
|
||||
enable_symm_mem=get_exec().comm.enable_symm_mem,
|
||||
device=model_runner.device,
|
||||
forward_stream=model_runner.forward_stream,
|
||||
)
|
||||
@@ -292,19 +297,17 @@ def capture_prefill_graph(
|
||||
return result(None)
|
||||
|
||||
# Disable prefill CUDA graph for non capture size
|
||||
if not model_runner.server_args.cuda_graph_config.prefill.bs:
|
||||
if not get_exec().graph.cuda_graph_config.prefill.bs:
|
||||
logger.warning("Disable prefill CUDA graph because the capture size is not set")
|
||||
return result(None)
|
||||
|
||||
prefill_config = model_runner.server_args.cuda_graph_config.prefill
|
||||
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
||||
prefill_backend = prefill_config.backend
|
||||
context_length = model_runner.model_config.context_len
|
||||
if prefill_backend == Backend.FULL:
|
||||
max_capture_requests = prefill_config.full_prefill_max_req
|
||||
if max_capture_requests is None:
|
||||
max_capture_requests = max(
|
||||
model_runner.server_args.chunked_prefill_size // 512, 1
|
||||
)
|
||||
max_capture_requests = max(get_schedule().chunked_prefill_size // 512, 1)
|
||||
max_capture_requests = min(
|
||||
max_capture_requests, model_runner.req_to_token_pool.size
|
||||
)
|
||||
@@ -418,7 +421,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
|
||||
if (
|
||||
model_runner.spec_algorithm.is_speculative()
|
||||
and not model_runner.is_draft_worker
|
||||
and model_runner.server_args.disaggregation_mode == "prefill"
|
||||
and get_disagg().disaggregation_mode == "prefill"
|
||||
):
|
||||
return no_capture
|
||||
if not model_runner.is_generation:
|
||||
@@ -453,7 +456,7 @@ def capture_decode_graph(*, model_runner: ModelRunner) -> GraphCapture:
|
||||
capture_name = f"{role} decode"
|
||||
num_tokens_per_req = 1
|
||||
capture_bs, _ = get_batch_sizes_to_capture(model_runner, num_tokens_per_req)
|
||||
decode_backend = model_runner.server_args.cuda_graph_config.decode.backend
|
||||
decode_backend = get_exec().graph.cuda_graph_config.decode.backend
|
||||
logger.info(
|
||||
f"Capture {capture_name} {graph_backend[model_runner.device]} begin. "
|
||||
f"backend={decode_backend}, num_tokens_per_req={num_tokens_per_req}, "
|
||||
|
||||
@@ -11,7 +11,12 @@ from sglang.srt.distributed import get_world_group
|
||||
from sglang.srt.mem_cache.kv_cache_configurator import mm_runtime_reservation_gb
|
||||
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import pre_capture_activation_reserve_mb
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_mm,
|
||||
pre_capture_activation_reserve_mb,
|
||||
)
|
||||
from sglang.srt.utils.common import get_available_gpu_memory, get_device_memory_capacity
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -55,11 +60,11 @@ def compute_post_capture_kv_resize(
|
||||
headroom_gb = model_runner.pre_model_load_memory * (
|
||||
1 - model_runner.mem_fraction_static
|
||||
)
|
||||
decode_cuda_graph_config = model_runner.server_args.cuda_graph_config.decode
|
||||
decode_cuda_graph_config = get_exec().graph.cuda_graph_config.decode
|
||||
decode_max_bs = int(decode_cuda_graph_config.max_bs or 0)
|
||||
running_requests = int(model_runner.max_running_requests or decode_max_bs or 1)
|
||||
eager_decode_gap = (
|
||||
model_runner.server_args.disaggregation_mode != "prefill"
|
||||
get_disagg().disaggregation_mode != "prefill"
|
||||
and decode_cuda_graph_config.backend != Backend.DISABLED
|
||||
and decode_max_bs < running_requests
|
||||
)
|
||||
@@ -80,7 +85,7 @@ def compute_post_capture_kv_resize(
|
||||
)
|
||||
mm_reservation_gb = mm_runtime_reservation_gb(
|
||||
is_multimodal=model_runner.model_config.is_multimodal,
|
||||
mm_feature_transport=model_runner.server_args.mm_feature_transport,
|
||||
mm_feature_transport=get_mm().mm_feature_transport,
|
||||
)
|
||||
budget_bytes = (
|
||||
int(max(0.0, free_gb - headroom_gb - mm_reservation_gb) * (1 << 30))
|
||||
|
||||
@@ -451,7 +451,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
||||
self._swa_layers_num > 0
|
||||
), "Hybrid SWA model must have at least one SWA layer"
|
||||
|
||||
self._swa_full_tokens_ratio = kvc.server_args.swa_full_tokens_ratio
|
||||
self._swa_full_tokens_ratio = get_schedule().swa_full_tokens_ratio
|
||||
self._sliding_window_size = kvc.sliding_window_size
|
||||
self._page_size = kvc.page_size
|
||||
|
||||
@@ -779,19 +779,19 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
||||
f"local={len(self.compression_ratios)}/{len(cfg.compress_ratios)}"
|
||||
)
|
||||
self.swa_page_size = cfg.window_size
|
||||
self.swa_ratio = kvc.server_args.swa_full_tokens_ratio
|
||||
self.is_speculative = kvc.server_args.speculative_algorithm is not None
|
||||
self.swa_ratio = get_schedule().swa_full_tokens_ratio
|
||||
self.is_speculative = get_spec().speculative_algorithm is not None
|
||||
self.online_c128_mtp_max_draft_tokens = (
|
||||
kvc.server_args.max_speculative_num_draft_tokens or 0
|
||||
)
|
||||
self.requested_max_running_requests_per_worker = (
|
||||
kvc.server_args.max_running_requests // kvc.ps.attn_dp_size
|
||||
if kvc.server_args.max_running_requests is not None
|
||||
get_schedule().max_running_requests // kvc.ps.attn_dp_size
|
||||
if get_schedule().max_running_requests is not None
|
||||
else None
|
||||
)
|
||||
self.disaggregation_mode = kvc.server_args.disaggregation_mode
|
||||
self.disaggregation_mode = get_disagg().disaggregation_mode
|
||||
self.disaggregation_decode_extra_slots = (
|
||||
kvc.server_args.disaggregation_decode_extra_slots or 0
|
||||
get_disagg().disaggregation_decode_extra_slots or 0
|
||||
)
|
||||
if kvc.server_args.enable_hisparse:
|
||||
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||
|
||||
@@ -46,7 +46,13 @@ from sglang.srt.model_executor.runner.flashinfer_autotune import (
|
||||
run_flashinfer_autotune_forward,
|
||||
should_run_flashinfer_autotune,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_flags, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import create_dummy_verify_input
|
||||
from sglang.srt.utils import (
|
||||
empty_context,
|
||||
@@ -214,7 +220,7 @@ class BaseRunner(ABC):
|
||||
self.tp_size = model_runner.server_args.tp_size
|
||||
# elastic-EP scale-up rewrites dp_size on the published config
|
||||
self.dp_size = get_parallel().dp_size
|
||||
self.pp_size = model_runner.server_args.pp_size
|
||||
self.pp_size = configured_pp_size()
|
||||
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
||||
self.return_hidden_states_mode = (
|
||||
CaptureHiddenMode.NULL
|
||||
@@ -265,7 +271,7 @@ class BaseRunner(ABC):
|
||||
with custom_all_reduce.register_graph_buffers).
|
||||
"""
|
||||
mr = self.model_runner
|
||||
if mr.server_args.flashinfer_allreduce_fusion_backend is None:
|
||||
if get_exec().comm.flashinfer_allreduce_fusion_backend is None:
|
||||
return
|
||||
|
||||
from sglang.srt.layers.communicator import FUSE_ALLREDUCE_MAX_BATCH_SIZE
|
||||
@@ -344,7 +350,7 @@ class BaseRunner(ABC):
|
||||
vocab_size=mr.model_config.vocab_size,
|
||||
dtype=mr.model_config.dtype,
|
||||
dp_size=get_parallel().dp_size,
|
||||
pp_size=mr.server_args.pp_size,
|
||||
pp_size=configured_pp_size(),
|
||||
is_encoder_decoder=mr.model_config.is_encoder_decoder,
|
||||
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
|
||||
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
|
||||
@@ -410,7 +416,7 @@ class BaseRunner(ABC):
|
||||
# TARGET_VERIFY dummy forward would trip the linear-attn backend's
|
||||
# pool-type assert. Warm up in plain DECODE instead.
|
||||
_is_pd_prefill_target = (
|
||||
mr.server_args.disaggregation_mode == "prefill" and not mr.is_draft_worker
|
||||
get_disagg().disaggregation_mode == "prefill" and not mr.is_draft_worker
|
||||
)
|
||||
if mr.spec_algorithm.is_speculative() and not _is_pd_prefill_target:
|
||||
if mr.is_draft_worker:
|
||||
@@ -514,7 +520,7 @@ class BaseRunner(ABC):
|
||||
extend_prefix_lens = None
|
||||
extend_start_loc = None
|
||||
|
||||
if mr.server_args.pp_size > 1:
|
||||
if configured_pp_size() > 1:
|
||||
# PP0 already cp-split hidden_states before send.
|
||||
pp_hidden_tokens = num_tokens
|
||||
if (
|
||||
@@ -638,7 +644,7 @@ class BaseRunner(ABC):
|
||||
|
||||
kwargs = {}
|
||||
if (
|
||||
mr.server_args.pp_size > 1
|
||||
configured_pp_size() > 1
|
||||
and "pp_proxy_tensors" in inspect.signature(mr.model.forward).parameters
|
||||
):
|
||||
kwargs["pp_proxy_tensors"] = PPProxyTensors(
|
||||
|
||||
@@ -99,7 +99,12 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import (
|
||||
)
|
||||
from sglang.srt.model_executor.runner_utils.shared_read_event import make_external_event
|
||||
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
||||
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
||||
from sglang.srt.utils import (
|
||||
empty_context,
|
||||
@@ -233,7 +238,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.require_mlp_tp_gather or self.require_attn_tp_gather
|
||||
)
|
||||
self.require_mlp_sync = (
|
||||
model_runner.server_args.enable_dp_attention or self.require_gathered_buffer
|
||||
get_parallel().enable_dp_attention or self.require_gathered_buffer
|
||||
)
|
||||
self.enable_two_batch_overlap = (
|
||||
model_runner.server_args.enable_two_batch_overlap
|
||||
@@ -243,7 +248,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
hf_config = model_runner.model_config.hf_config
|
||||
self.ngram_embedding_n = hf_config.ngram_embedding_n
|
||||
self.ngram_embedding_k = hf_config.ngram_embedding_k
|
||||
self.speculative_algorithm = model_runner.server_args.speculative_algorithm
|
||||
self.speculative_algorithm = get_spec().speculative_algorithm
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
@@ -1127,7 +1132,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
bs = self._ragged_capture_slots(num_tokens) if self.ragged_verify_mode else size
|
||||
|
||||
# Sanity-check: --debug-cuda-graph requires breakable backend.
|
||||
if self.model_runner.server_args.debug_cuda_graph:
|
||||
if get_exec().graph.debug_cuda_graph:
|
||||
assert isinstance(
|
||||
self.backend, BreakableCudaGraphBackend
|
||||
), "Breakable CUDA graph is required for --debug-cuda-graph"
|
||||
@@ -1477,7 +1482,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
||||
retrieve_next_sibling=None,
|
||||
retrieve_cum_len=None,
|
||||
spec_steps=self.speculative_num_steps,
|
||||
topk=self.model_runner.server_args.speculative_eagle_topk,
|
||||
topk=get_spec().speculative_eagle_topk,
|
||||
draft_token_num=self.speculative_num_draft_tokens,
|
||||
capture_hidden_mode=capture_mode,
|
||||
seq_lens_sum=None,
|
||||
|
||||
@@ -116,7 +116,12 @@ from sglang.srt.model_executor.runner_utils.buffers import (
|
||||
PrefillInputBuffers,
|
||||
)
|
||||
from sglang.srt.model_loader.utils import resolve_language_model
|
||||
from sglang.srt.runtime_context import get_parallel, get_schedule
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
)
|
||||
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||
from sglang.srt.utils import (
|
||||
get_available_gpu_memory,
|
||||
@@ -262,7 +267,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
self.capture_return_pooled_hidden_states = not model_runner.is_generation
|
||||
|
||||
# --- prefill graph config -------------------------------------
|
||||
prefill_config = model_runner.server_args.cuda_graph_config.prefill
|
||||
prefill_config = get_exec().graph.cuda_graph_config.prefill
|
||||
self.prefill_backend_name = prefill_config.backend
|
||||
# bs in prefill carries the captured shape (token count for
|
||||
# tc_piecewise) — one shape knob per phase.
|
||||
@@ -424,7 +429,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
f"{type(attn_backend).__name__} does not support chunked-prefix "
|
||||
"Full prefill CUDA graphs"
|
||||
)
|
||||
prefix_config = model_runner.server_args.cuda_graph_config.prefill
|
||||
prefix_config = get_exec().graph.cuda_graph_config.prefill
|
||||
(
|
||||
self._prefix_chunk_len,
|
||||
self._prefix_chunk_capacity,
|
||||
@@ -531,7 +536,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
|
||||
def _is_mamba_track_enabled(self) -> bool:
|
||||
return self.model_runner.server_args.enable_mamba_extra_buffer() and (
|
||||
not self.model_runner.server_args.disable_radix_cache
|
||||
not get_memory().disable_radix_cache
|
||||
)
|
||||
|
||||
def _cache_loc_dtype(self):
|
||||
@@ -802,10 +807,10 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
|
||||
model_runner, capture_req_slots: int
|
||||
) -> tuple[int, int]:
|
||||
"""Resolve per-request length and aggregate capacity of one chunk."""
|
||||
prefix_config = model_runner.server_args.cuda_graph_config.prefill
|
||||
prefix_config = get_exec().graph.cuda_graph_config.prefill
|
||||
requested_capacity = prefix_config.full_prefill_prefix_chunk_tokens
|
||||
if requested_capacity is None:
|
||||
requested_capacity = model_runner.server_args.chunked_prefill_size
|
||||
requested_capacity = get_schedule().chunked_prefill_size
|
||||
if requested_capacity is None or requested_capacity <= 0:
|
||||
requested_capacity = prefix_config.max_bs
|
||||
if requested_capacity is None or requested_capacity <= 0:
|
||||
|
||||
@@ -49,7 +49,10 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
||||
from sglang.srt.model_executor.runner_utils.pool import (
|
||||
get_or_create_global_graph_memory_pool,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.utils import is_hip
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -110,7 +113,7 @@ class TcPiecewiseCudaGraphBackend(BaseCudaGraphBackend):
|
||||
def build_compilation_config(server_args: ServerArgs) -> CompilationConfig:
|
||||
"""Construct a CompilationConfig from ServerArgs and
|
||||
register the MoE A2A split-op when DeepEP / Mooncake is in use."""
|
||||
prefill = server_args.cuda_graph_config.prefill
|
||||
prefill = get_exec().graph.cuda_graph_config.prefill
|
||||
num_tokens = prefill.bs
|
||||
compiler = prefill.tc_compiler
|
||||
assert num_tokens is not None, "cuda_graph_config[prefill].bs is not set"
|
||||
|
||||
@@ -37,6 +37,9 @@ from sglang.srt.model_executor.runner_backend.full_cuda_graph_backend import (
|
||||
from sglang.srt.model_executor.runner_backend.tc_piecewise_cuda_graph_backend import (
|
||||
TcPiecewiseCudaGraphBackend,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.runner.base_cuda_graph_runner import (
|
||||
@@ -58,7 +61,7 @@ def resolve_decode_backend(
|
||||
the Full-style backend is wired for NPU today).
|
||||
"""
|
||||
model_runner = cuda_graph_runner.model_runner
|
||||
cfg = model_runner.server_args.cuda_graph_config
|
||||
cfg = get_exec().graph.cuda_graph_config
|
||||
backend_name = cfg.decode.backend if cfg is not None else Backend.FULL
|
||||
|
||||
enable_memory_saver = model_runner.server_args.enable_memory_saver
|
||||
@@ -86,7 +89,7 @@ def resolve_decode_backend(
|
||||
return BreakableCudaGraphBackend(
|
||||
cuda_graph_runner,
|
||||
enable_memory_saver=enable_memory_saver,
|
||||
debug_eager=model_runner.server_args.debug_cuda_graph,
|
||||
debug_eager=get_exec().graph.debug_cuda_graph,
|
||||
)
|
||||
if backend_name == Backend.TC_PIECEWISE:
|
||||
global _TC_PIECEWISE_DECODE_FALLBACK_LOGGED
|
||||
@@ -106,14 +109,14 @@ def resolve_prefill_backend(
|
||||
) -> BaseCudaGraphBackend:
|
||||
"""Pick a backend instance from cuda_graph_config['prefill']['backend']."""
|
||||
model_runner = cuda_graph_runner.model_runner
|
||||
cfg = model_runner.server_args.cuda_graph_config
|
||||
cfg = get_exec().graph.cuda_graph_config
|
||||
backend_name = cfg.prefill.backend if cfg is not None else Backend.TC_PIECEWISE
|
||||
|
||||
if backend_name == Backend.BREAKABLE:
|
||||
return BreakableCudaGraphBackend(
|
||||
cuda_graph_runner,
|
||||
enable_memory_saver=model_runner.server_args.enable_memory_saver,
|
||||
debug_eager=model_runner.server_args.debug_cuda_graph,
|
||||
debug_eager=get_exec().graph.debug_cuda_graph,
|
||||
)
|
||||
if backend_name == Backend.FULL:
|
||||
return FullCudaGraphBackend(
|
||||
|
||||
@@ -34,7 +34,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
||||
from sglang.srt.model_executor.runner_backend_utils import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_flags, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_flags,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput
|
||||
from sglang.srt.speculative.eagle_utils import get_draft_recurrent_hidden_state_spec
|
||||
@@ -109,7 +113,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.device_module = torch.get_device_module(self.device)
|
||||
self.tp_size = model_runner.ps.tp_size
|
||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
||||
self.pp_size = model_runner.server_args.pp_size
|
||||
self.pp_size = configured_pp_size()
|
||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
@@ -124,7 +128,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
if speculative_num_steps is None
|
||||
else speculative_num_steps
|
||||
)
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk
|
||||
self.topk = get_spec().speculative_eagle_topk
|
||||
self.draft_attn_backend = draft_attn_backend or model_runner.draft_attn_backend
|
||||
|
||||
# Patch_model in parent's capture() needs an attn_backend reference.
|
||||
|
||||
@@ -35,7 +35,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
||||
from sglang.srt.model_executor.runner_backend_utils import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_flags, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_flags,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
||||
@@ -95,7 +99,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.device_module = torch.get_device_module(self.device)
|
||||
self.tp_size = model_runner.ps.tp_size
|
||||
self.attn_dp_size = model_runner.ps.attn_dp_size
|
||||
self.pp_size = model_runner.server_args.pp_size
|
||||
self.pp_size = configured_pp_size()
|
||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
|
||||
@@ -32,7 +32,11 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
||||
from sglang.srt.model_executor.runner_backend_utils import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_flags, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_flags,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPDraftInput
|
||||
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
||||
from sglang.srt.utils import (
|
||||
@@ -95,9 +99,9 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
|
||||
self.tp_size = self.model_runner.ps.tp_size
|
||||
self.attn_dp_size = self.model_runner.ps.attn_dp_size
|
||||
self.pp_size = model_runner.server_args.pp_size
|
||||
self.pp_size = configured_pp_size()
|
||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk
|
||||
self.topk = get_spec().speculative_eagle_topk
|
||||
self.draft_attn_backend = frozen_kv_mtp_worker.draft_attn_backend
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
|
||||
@@ -59,7 +59,12 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
|
||||
from sglang.srt.model_executor.runner_backend_utils import (
|
||||
CUDA_GRAPH_CAPTURE_FAILED_MSG,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_pp_size,
|
||||
get_flags,
|
||||
get_parallel,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||
from sglang.srt.speculative.eagle_utils import get_draft_input_from_target_hidden_dim
|
||||
from sglang.srt.speculative.multi_layer_eagle_utils import (
|
||||
@@ -151,7 +156,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.device_module = torch.get_device_module(self.device)
|
||||
self.tp_size = model_runner.ps.tp_size
|
||||
self.dp_size = get_parallel().dp_size
|
||||
self.pp_size = model_runner.server_args.pp_size
|
||||
self.pp_size = configured_pp_size()
|
||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
|
||||
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
|
||||
@@ -161,7 +166,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
||||
self.enable_pdmux = model_runner.server_args.enable_pdmux
|
||||
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||
self.topk = model_runner.server_args.speculative_eagle_topk
|
||||
self.topk = get_spec().speculative_eagle_topk
|
||||
self.enable_profile_cuda_graph = (
|
||||
model_runner.server_args.enable_profile_cuda_graph
|
||||
)
|
||||
|
||||
@@ -33,6 +33,7 @@ from sglang.srt.managers.schedule_batch import (
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
configured_tp_size,
|
||||
get_mm,
|
||||
get_parallel,
|
||||
)
|
||||
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
||||
@@ -933,7 +934,7 @@ class CudaVmmFeatureTransport:
|
||||
|
||||
def __init__(self, server_args, mm_processor) -> None:
|
||||
self.pool: CudaVmmMemoryPool | None = None
|
||||
if server_args.mm_feature_transport != "cuda_vmm":
|
||||
if get_mm().mm_feature_transport != "cuda_vmm":
|
||||
return
|
||||
if mm_processor is None:
|
||||
raise RuntimeError(
|
||||
|
||||
Reference in New Issue
Block a user