config: read resolved config via namespace accessors (#31814)
This commit is contained in:
@@ -39,7 +39,12 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
compute_position,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpecInput
|
||||
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
|
||||
|
||||
@@ -183,7 +188,7 @@ def _update_device_and_sum_field_from_cpu_field(
|
||||
cpu_value
|
||||
if isinstance(cpu_value, torch.Tensor)
|
||||
else torch.tensor(cpu_value, dtype=old_device_value.dtype)
|
||||
).to(device=get_server_args().device, non_blocking=True)
|
||||
).to(device=get_device().device, non_blocking=True)
|
||||
setattr(batch, device_field, new_device_value)
|
||||
|
||||
if sum_field is not None:
|
||||
@@ -335,7 +340,7 @@ def compute_split_indices_for_cuda_graph_replay(
|
||||
class TboCudaGraphRunnerPlugin:
|
||||
def __init__(self):
|
||||
self._tbo_children_num_token_non_padded = torch.zeros(
|
||||
(2,), dtype=torch.int32, device=get_server_args().device
|
||||
(2,), dtype=torch.int32, device=get_device().device
|
||||
)
|
||||
|
||||
def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int):
|
||||
@@ -633,7 +638,7 @@ class TboForwardBatchPreparer:
|
||||
sum_field=None,
|
||||
)
|
||||
_, child_b.extend_start_loc = compute_position(
|
||||
get_server_args().attention_backend,
|
||||
get_exec().kernel.attention_backend,
|
||||
child_b.extend_prefix_lens,
|
||||
child_b.extend_seq_lens,
|
||||
child_b.extend_num_tokens,
|
||||
@@ -832,7 +837,7 @@ class TboForwardBatchPreparer:
|
||||
value_a = min(tbo_split_token_index, num_token_non_padded)
|
||||
value_b = max(0, num_token_non_padded - tbo_split_token_index)
|
||||
return torch.tensor([value_a, value_b], dtype=torch.int32).to(
|
||||
device=get_server_args().device, non_blocking=True
|
||||
device=get_device().device, non_blocking=True
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -8,6 +8,7 @@ from transformers import CONFIG_MAPPING
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
|
||||
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
|
||||
|
||||
class InklingModelConfig(PretrainedConfig):
|
||||
@@ -224,9 +225,8 @@ class InklingModelConfig(PretrainedConfig):
|
||||
self.swa_num_key_value_heads, self.swa_head_dim
|
||||
)
|
||||
stream_dim = self.hidden_size
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
if get_server_args().enable_scattered_sconv:
|
||||
if get_exec().comm.enable_scattered_sconv:
|
||||
# Scattered sconv: the attn/mlp output sconvs run on the [T, H/P]
|
||||
# hidden shard, so their conv-state caches shard with them.
|
||||
assert (
|
||||
|
||||
@@ -14,6 +14,7 @@ from sglang.srt.constrained.base_grammar_backend import (
|
||||
from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
|
||||
from sglang.srt.distributed.communication_tags import P2PTag
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.managers.io_struct import AbortReq
|
||||
@@ -28,7 +29,7 @@ class GrammarManager:
|
||||
self.scheduler = scheduler
|
||||
self.server_args = scheduler.server_args
|
||||
self.grammar_queue: List[Req] = []
|
||||
if not self.server_args.skip_tokenizer_init:
|
||||
if not get_serving().skip_tokenizer_init:
|
||||
self.grammar_backend = create_grammar_backend(
|
||||
self.server_args,
|
||||
scheduler.tokenizer,
|
||||
|
||||
@@ -32,11 +32,8 @@ from sglang.srt.disaggregation.utils import (
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group, get_world_group
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
get_attention_dp_rank,
|
||||
get_attention_dp_size,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_model, get_parallel
|
||||
from sglang.srt.layers.dp_attention import get_attention_dp_rank, get_attention_dp_size
|
||||
from sglang.srt.runtime_context import get_model, get_parallel, get_serving
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.network import (
|
||||
NetworkAddress,
|
||||
@@ -634,7 +631,7 @@ class CommonKVManager(BaseKVManager):
|
||||
# Self-register the HTTP API port so the decode can derive the PD
|
||||
# retract rebootstrap /generate URL from bootstrap info instead of a
|
||||
# router-injected pd_rebootstrap_prefill_url.
|
||||
"prefill_http_port": self.server_args.port,
|
||||
"prefill_http_port": get_serving().port,
|
||||
}
|
||||
|
||||
max_retries, initial_delay, max_delay = 5, 1.0, 30.0
|
||||
|
||||
@@ -86,7 +86,7 @@ from sglang.srt.observability.req_time_stats import (
|
||||
set_schedule_time_batch,
|
||||
set_time_batch,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_disagg, get_parallel
|
||||
from sglang.srt.utils import get_num_new_pages
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
@@ -2151,7 +2151,7 @@ class SchedulerDisaggregationDecodeMixin:
|
||||
# Decode-radix path: new requests already matched in
|
||||
# `pop_preallocated`. Retracted requests reset `last_node`,
|
||||
# so re-match only when that state is missing.
|
||||
if self.server_args.disaggregation_decode_enable_radix_cache:
|
||||
if get_disagg().disaggregation_decode_enable_radix_cache:
|
||||
tree_cache = self.tree_cache if req.last_node is None else None
|
||||
else:
|
||||
tree_cache = self.tree_cache
|
||||
@@ -2191,7 +2191,7 @@ class SchedulerDisaggregationDecodeMixin:
|
||||
if self.enable_decode_hicache:
|
||||
self.tree_cache.check_hicache_events()
|
||||
|
||||
if self.server_args.disaggregation_decode_enable_offload_kvcache:
|
||||
if get_disagg().disaggregation_decode_enable_offload_kvcache:
|
||||
self.decode_offload_manager.check_offload_progress()
|
||||
|
||||
# try to resume retracted requests if there are enough space for another `num_reserved_decode_tokens` decode steps
|
||||
@@ -2203,9 +2203,7 @@ class SchedulerDisaggregationDecodeMixin:
|
||||
|
||||
if not hasattr(self, "polling_count"):
|
||||
self.polling_count = 0
|
||||
self.polling_interval = (
|
||||
self.server_args.disaggregation_decode_polling_interval
|
||||
)
|
||||
self.polling_interval = get_disagg().disaggregation_decode_polling_interval
|
||||
|
||||
self.polling_count = (self.polling_count + 1) % self.polling_interval
|
||||
|
||||
|
||||
@@ -28,6 +28,7 @@ from sglang.srt.disaggregation.encode_server import (
|
||||
)
|
||||
from sglang.srt.managers.io_struct import async_sock_send, wrap_as_pickle
|
||||
from sglang.srt.managers.schedule_batch import Modality
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import random_uuid
|
||||
from sglang.srt.utils.network import NetworkAddress, get_zmq_socket
|
||||
@@ -117,13 +118,13 @@ class SGLangEncoderServer(SGLangEncoderServicer):
|
||||
context.set_details(error_msg)
|
||||
return sglang_encoder_pb2.EncodeResponse()
|
||||
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
return sglang_encoder_pb2.EncodeResponse(
|
||||
embedding_size=nbytes,
|
||||
embedding_len=embedding_len,
|
||||
embedding_dim=embedding_dim,
|
||||
)
|
||||
elif self.server_args.encoder_transfer_backend == "zmq_to_scheduler":
|
||||
elif get_disagg().encoder_transfer_backend == "zmq_to_scheduler":
|
||||
embedding_ports = list(request.embedding_port)
|
||||
logger.info(f"embedding_port = {embedding_ports}")
|
||||
if not embedding_ports:
|
||||
@@ -141,7 +142,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
|
||||
await asyncio.gather(*tasks)
|
||||
self.encoder.embedding_to_send.pop(request.req_id, None)
|
||||
return sglang_encoder_pb2.EncodeResponse()
|
||||
elif self.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
|
||||
elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer":
|
||||
embedding_port = (
|
||||
request.embedding_port[0] if request.embedding_port else 0
|
||||
)
|
||||
|
||||
@@ -59,14 +59,9 @@ from sglang.srt.model_loader import get_model
|
||||
from sglang.srt.multimodal.processors.qwen_vl import preprocess_video
|
||||
from sglang.srt.observability.metrics_collector import EncoderMetricsCollector
|
||||
from sglang.srt.observability.req_time_stats import EncoderReqTimeStats
|
||||
from sglang.srt.observability.trace import (
|
||||
process_tracing_init,
|
||||
trace_set_thread_info,
|
||||
)
|
||||
from sglang.srt.server_args import (
|
||||
PortArgs,
|
||||
ServerArgs,
|
||||
)
|
||||
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
||||
from sglang.srt.runtime_context import get_disagg, get_exec, get_mm
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
add_prometheus_middleware,
|
||||
configure_logger,
|
||||
@@ -350,7 +345,7 @@ class MMEncoder:
|
||||
[], dtype=self._embedding_dtype
|
||||
).element_size()
|
||||
|
||||
if self.server_args.enable_mm_global_cache:
|
||||
if get_mm().enable_mm_global_cache:
|
||||
from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import (
|
||||
EmbeddingCacheController,
|
||||
)
|
||||
@@ -368,15 +363,15 @@ class MMEncoder:
|
||||
self.mm_global_cache = None
|
||||
|
||||
# Pre-compute embedding metadata (needed by all ranks for mooncake)
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
self._embedding_dims = self._infer_embedding_dims()
|
||||
|
||||
if self.rank == 0:
|
||||
logger.info(
|
||||
f"Using transfer backend: {self.server_args.encoder_transfer_backend}"
|
||||
f"Using transfer backend: {get_disagg().encoder_transfer_backend}"
|
||||
)
|
||||
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
self.local_ip = get_local_ip_auto()
|
||||
|
||||
self.engine = get_mooncake_transfer_engine()
|
||||
@@ -389,8 +384,8 @@ class MMEncoder:
|
||||
hostname=self.local_ip,
|
||||
gpu_id=self.gpu_id,
|
||||
ib_device=(
|
||||
self.server_args.disaggregation_ib_device
|
||||
or self.server_args.mooncake_ib_device
|
||||
get_disagg().disaggregation_ib_device
|
||||
or get_exec().moe.mooncake_ib_device
|
||||
),
|
||||
)
|
||||
|
||||
@@ -399,7 +394,7 @@ class MMEncoder:
|
||||
self.encode_dispatch_lock = asyncio.Lock()
|
||||
|
||||
# Async mooncake state: track background VIT forward completion
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
self._forward_ready_events: Dict[str, asyncio.Event] = {}
|
||||
self._forward_results: Dict[str, dict] = {}
|
||||
# when multiple decoder TP ranks call
|
||||
@@ -413,12 +408,12 @@ class MMEncoder:
|
||||
|
||||
# Bind unified encode entry point based on backend and cache config
|
||||
if self.mm_global_cache is not None:
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
self._encode_fn = self.encode_with_global_cache_mooncake
|
||||
else:
|
||||
self._encode_fn = self.encode_with_global_cache
|
||||
else:
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
self._encode_fn = self.encode_with_mooncake
|
||||
else:
|
||||
self._encode_fn = self.encode
|
||||
@@ -1688,7 +1683,7 @@ class MMEncoder:
|
||||
mm_item.set(k, _convert(v))
|
||||
|
||||
cache_hit = False
|
||||
use_mm_cache = self.server_args.enable_prefix_mm_cache and log_metrics
|
||||
use_mm_cache = get_mm().enable_prefix_mm_cache and log_metrics
|
||||
if use_mm_cache:
|
||||
mm_item.set_pad_value()
|
||||
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
|
||||
@@ -1784,7 +1779,7 @@ class MMEncoder:
|
||||
embedding_port=None,
|
||||
url=None,
|
||||
):
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
# Wait for async VIT forward completion if needed
|
||||
req_id = mm_data.req_id
|
||||
if req_id in self._forward_ready_events:
|
||||
@@ -1855,7 +1850,7 @@ class MMEncoder:
|
||||
logger.info(f"{endpoint = }")
|
||||
|
||||
# Serialize data
|
||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
||||
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||
# Mooncake already pushed the embedding via RDMA;
|
||||
new_mm_data = mm_data.copy_without_embedding()
|
||||
serialized_data = pickle.dumps(new_mm_data)
|
||||
@@ -1887,11 +1882,11 @@ class MMEncoder:
|
||||
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
|
||||
if (
|
||||
encoder_metrics_collector is not None
|
||||
and self.server_args.encoder_transfer_backend != "mooncake"
|
||||
and get_disagg().encoder_transfer_backend != "mooncake"
|
||||
):
|
||||
encoder_metrics_collector.observe_transfer(
|
||||
time.perf_counter() - _zmq_xfer_start,
|
||||
backend=self.server_args.encoder_transfer_backend,
|
||||
backend=get_disagg().encoder_transfer_backend,
|
||||
)
|
||||
|
||||
async def encode(
|
||||
|
||||
@@ -55,6 +55,7 @@ from sglang.srt.observability.trace import (
|
||||
TraceReqContext,
|
||||
trace_set_thread_info,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.network import NetworkAddress
|
||||
|
||||
@@ -314,9 +315,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
self.kv_buffer_tensors = None
|
||||
|
||||
def _handle_staging_req(self, msg):
|
||||
from sglang.srt.disaggregation.common.staging_handler import (
|
||||
handle_staging_req,
|
||||
)
|
||||
from sglang.srt.disaggregation.common.staging_handler import handle_staging_req
|
||||
|
||||
room = int(msg[1].decode("ascii"))
|
||||
session_id = msg[4].decode("ascii")
|
||||
@@ -350,9 +349,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
def _is_watermark_ready(
|
||||
self, session_id: str, alloc_round: int, alloc_end: int
|
||||
) -> bool:
|
||||
from sglang.srt.disaggregation.common.staging_handler import (
|
||||
is_watermark_ready,
|
||||
)
|
||||
from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready
|
||||
|
||||
return is_watermark_ready(self._staging_ctx, session_id, alloc_round, alloc_end)
|
||||
|
||||
@@ -469,7 +466,7 @@ class MooncakeKVManager(CommonKVManager):
|
||||
room,
|
||||
self.transfer_infos,
|
||||
self.kv_buffer_tensors,
|
||||
self.server_args.chunked_prefill_size,
|
||||
get_schedule().chunked_prefill_size,
|
||||
self._staging_ctx.prefetch_requested,
|
||||
self._staging_ctx.prefetch_sockets,
|
||||
)
|
||||
|
||||
@@ -13,6 +13,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set, Tuple
|
||||
import numpy as np
|
||||
import numpy.typing as npt
|
||||
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.disaggregation.common.staging_handler import StagingTransferInfo
|
||||
|
||||
@@ -535,9 +537,7 @@ class NixlKVManager(CommonKVManager):
|
||||
def _is_watermark_ready(
|
||||
self, agent_name: str, alloc_round: int, alloc_end: int
|
||||
) -> bool:
|
||||
from sglang.srt.disaggregation.common.staging_handler import (
|
||||
is_watermark_ready,
|
||||
)
|
||||
from sglang.srt.disaggregation.common.staging_handler import is_watermark_ready
|
||||
|
||||
return is_watermark_ready(self._staging_ctx, agent_name, alloc_round, alloc_end)
|
||||
|
||||
@@ -558,9 +558,7 @@ class NixlKVManager(CommonKVManager):
|
||||
threading.Thread(target=decode_staging_thread, daemon=True).start()
|
||||
|
||||
def _handle_staging_req(self, msg):
|
||||
from sglang.srt.disaggregation.common.staging_handler import (
|
||||
handle_staging_req,
|
||||
)
|
||||
from sglang.srt.disaggregation.common.staging_handler import handle_staging_req
|
||||
|
||||
room = int(msg[1].decode("ascii"))
|
||||
session_id = msg[4].decode("ascii")
|
||||
@@ -625,7 +623,7 @@ class NixlKVManager(CommonKVManager):
|
||||
room,
|
||||
self.transfer_infos,
|
||||
self.kv_buffer_tensors,
|
||||
self.server_args.chunked_prefill_size,
|
||||
get_schedule().chunked_prefill_size,
|
||||
self._staging_ctx.prefetch_requested,
|
||||
self._staging_ctx.prefetch_sockets,
|
||||
)
|
||||
@@ -1739,9 +1737,7 @@ class NixlKVManager(CommonKVManager):
|
||||
req, page_start, num_pages, session_id=req.agent_name
|
||||
)
|
||||
if not ready:
|
||||
from sglang.srt.disaggregation.common.staging_buffer import (
|
||||
StagingAllocator,
|
||||
)
|
||||
from sglang.srt.disaggregation.common.staging_buffer import StagingAllocator
|
||||
|
||||
if c_offset == StagingAllocator.ALLOC_OVERSIZED:
|
||||
raise RuntimeError(
|
||||
|
||||
@@ -64,6 +64,7 @@ from sglang.srt.mem_cache.common import (
|
||||
)
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.observability.req_time_stats import set_schedule_time_batch
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -1181,7 +1182,7 @@ class SchedulerDisaggregationPrefillMixin:
|
||||
|
||||
def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None:
|
||||
"""Release KV cache and requeue an optimistic prefill request."""
|
||||
max_attempts = self.server_args.optimistic_prefill_attempts
|
||||
max_attempts = get_disagg().optimistic_prefill_attempts
|
||||
maybe_cache_unfinished_req(req, self.tree_cache)
|
||||
release_kv_cache(req, self.tree_cache)
|
||||
req.reset_for_retract()
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.srt.compilation.compile_phase import (
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -25,7 +25,7 @@ class PyMscclppCommunicator:
|
||||
|
||||
def _is_symm_mem_enabled(self) -> bool:
|
||||
try:
|
||||
return get_server_args().enable_symm_mem
|
||||
return get_exec().comm.enable_symm_mem
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
@@ -15,7 +15,7 @@ from torch.cuda.memory import (
|
||||
|
||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils.common import torch_release
|
||||
|
||||
after_2_8_0 = torch_release >= (2, 8)
|
||||
@@ -159,7 +159,7 @@ _register_func = None
|
||||
|
||||
def is_symmetric_memory_enabled():
|
||||
try:
|
||||
return get_server_args().enable_symm_mem
|
||||
return get_exec().comm.enable_symm_mem
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
@@ -12,6 +12,7 @@ from sglang.srt.distributed.device_communicators.all_reduce_utils import (
|
||||
TORCH_SYMM_MEM_ALL_REDUCE_MAX_SIZES,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import is_cuda, is_hip
|
||||
|
||||
try:
|
||||
@@ -98,10 +99,9 @@ class TorchSymmMemCommunicator:
|
||||
# ([16384, 6144] bf16 = 192 MiB), including room for tail regions.
|
||||
if envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get():
|
||||
self.max_size = max(self.max_size, 256 * 1024 * 1024)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
if (
|
||||
get_server_args().enable_scattered_sconv
|
||||
get_exec().comm.enable_scattered_sconv
|
||||
or envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get()
|
||||
):
|
||||
# Fused extend kernels are out-of-place, so OUT must hold the
|
||||
|
||||
@@ -11,6 +11,7 @@ from sglang.srt.managers.schedule_policy import AddReqResult, PrefillAdder
|
||||
from sglang.srt.mem_cache.common import release_kv_cache
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
||||
from sglang.srt.runtime_context import get_exec, get_schedule
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -22,7 +23,7 @@ class SchedulerDllmMixin:
|
||||
def init_diffusion_llm(self: Scheduler):
|
||||
self.dllm_config = (
|
||||
DllmConfig.from_server_args(self.server_args)
|
||||
if self.server_args.dllm_algorithm is not None
|
||||
if get_exec().dllm.dllm_algorithm is not None
|
||||
else None
|
||||
)
|
||||
self.dllm_manager = DllmManager(dllm_config=self.dllm_config)
|
||||
@@ -200,7 +201,7 @@ class SchedulerDllmMixin:
|
||||
self.chunked_prefill_size,
|
||||
running_bs if self.is_mixed_chunk else 0,
|
||||
self.priority_scheduling_preemption_threshold,
|
||||
prefill_max_requests=self.server_args.prefill_max_requests,
|
||||
prefill_max_requests=get_schedule().prefill_max_requests,
|
||||
dllm_config=self.dllm_config,
|
||||
)
|
||||
|
||||
|
||||
@@ -442,7 +442,7 @@ def get_healthy_expert_location_src_rank(
|
||||
*, invoked_in_elastic_ep_rejoin_path: bool
|
||||
) -> int:
|
||||
world_group = get_world_group()
|
||||
# NOTE: do not key off `self.server_args.elastic_ep_rejoin` here.
|
||||
# NOTE: do not key off `get_exec().moe.elastic_ep_rejoin` here.
|
||||
# A rank that was started as a rejoin rank may later act as a healthy
|
||||
# rank in a subsequent recovery cycle.
|
||||
local_rejoin_flag = bool(invoked_in_elastic_ep_rejoin_path)
|
||||
|
||||
@@ -7,13 +7,11 @@ from typing import Any, Callable
|
||||
import torch
|
||||
import zmq
|
||||
|
||||
from sglang.srt.distributed.parallel_state import (
|
||||
get_world_group,
|
||||
get_world_size,
|
||||
)
|
||||
from sglang.srt.distributed.parallel_state import get_world_group, get_world_size
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
|
||||
from sglang.srt.managers.io_struct import UpdateExpertBackupReq, sock_recv, sock_send
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.network import get_local_ip_auto
|
||||
|
||||
@@ -111,7 +109,7 @@ class ExpertBackupClient:
|
||||
global_expert_location_metadata = get_global_expert_location_metadata()
|
||||
num_experts = (
|
||||
self.model_config.hf_config.n_routed_experts
|
||||
+ self.server_args.ep_num_redundant_experts
|
||||
+ get_exec().moe.ep_num_redundant_experts
|
||||
)
|
||||
num_local_experts = num_experts // self.moe_ep_size
|
||||
for i in range(self.engine_num):
|
||||
|
||||
@@ -877,7 +877,14 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
server_args, port_args
|
||||
)
|
||||
else:
|
||||
# Launch multi-tokenizer router
|
||||
# Launch multi-tokenizer router. Unlike TokenizerManager, the router
|
||||
# does not publish; but it runs in this parent process and reads
|
||||
# resolved config through the namespace accessors (e.g. get_parallel()
|
||||
# for routed_dp_rank), so publish here. The child TokenizerWorkers
|
||||
# publish independently in their own processes.
|
||||
from sglang.srt.runtime_context import publish
|
||||
|
||||
publish(server_args, role="tokenizer")
|
||||
tokenizer_manager = MultiTokenizerRouter(server_args, port_args)
|
||||
template_manager = None
|
||||
|
||||
@@ -997,12 +1004,18 @@ class Engine(EngineScoreMixin, EngineBase):
|
||||
)
|
||||
|
||||
def get_server_info(self):
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
internal_states = self.loop.run_until_complete(
|
||||
self.tokenizer_manager.get_internal_state()
|
||||
)
|
||||
return msgspec_to_builtins(
|
||||
{
|
||||
**dataclasses.asdict(self.tokenizer_manager.server_args),
|
||||
# Overlay post-publish overrides so the report reflects current
|
||||
# config (weight version, model path, runtime tunables).
|
||||
**get_context().resolved_server_args_dict(
|
||||
base=dataclasses.asdict(self.tokenizer_manager.server_args)
|
||||
),
|
||||
**self._scheduler_init_result.scheduler_infos[0],
|
||||
"internal_states": internal_states,
|
||||
"version": __version__,
|
||||
|
||||
@@ -15,6 +15,7 @@ from typing import Any, Awaitable, Callable, Dict, List, Optional
|
||||
|
||||
from pydantic import ValidationError
|
||||
|
||||
from sglang.srt.runtime_context import get_context, get_lora, get_serving
|
||||
from sglang.srt.utils.msgspec_utils import msgspec_to_builtins
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -229,9 +230,7 @@ class RuntimeHandle:
|
||||
return self._openai_serving_classes
|
||||
|
||||
from sglang.srt.entrypoints.openai.serving_chat import OpenAIServingChat
|
||||
from sglang.srt.entrypoints.openai.serving_classify import (
|
||||
OpenAIServingClassify,
|
||||
)
|
||||
from sglang.srt.entrypoints.openai.serving_classify import OpenAIServingClassify
|
||||
from sglang.srt.entrypoints.openai.serving_completions import (
|
||||
OpenAIServingCompletion,
|
||||
)
|
||||
@@ -376,16 +375,20 @@ class RuntimeHandle:
|
||||
model_config = self.tokenizer_manager.model_config
|
||||
result = {
|
||||
"model_path": self.tokenizer_manager.model_path,
|
||||
"tokenizer_path": self.server_args.tokenizer_path,
|
||||
"tokenizer_path": get_serving().tokenizer_path,
|
||||
"is_generation": self.tokenizer_manager.is_generation,
|
||||
"weight_version": self.server_args.weight_version,
|
||||
"weight_version": get_serving().weight_version,
|
||||
"model_type": getattr(model_config.hf_config, "model_type", None),
|
||||
"architectures": getattr(model_config.hf_config, "architectures", None),
|
||||
}
|
||||
return json.dumps(result, default=str)
|
||||
|
||||
def get_server_info(self) -> str:
|
||||
result: Dict[str, Any] = dataclasses.asdict(self.server_args)
|
||||
# Overlay post-publish overrides (weight version, model path, runtime
|
||||
# tunables) so the report reflects current config, not the startup record.
|
||||
result: Dict[str, Any] = get_context().resolved_server_args_dict(
|
||||
base=dataclasses.asdict(self.server_args)
|
||||
)
|
||||
result.update(self.scheduler_info)
|
||||
return json.dumps(msgspec_to_builtins(result), default=str)
|
||||
|
||||
@@ -424,9 +427,7 @@ class RuntimeHandle:
|
||||
"max_model_len": self.tokenizer_manager.model_config.context_len,
|
||||
}
|
||||
]
|
||||
if self.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(
|
||||
|
||||
@@ -693,18 +693,19 @@ async def get_model_info():
|
||||
@app.get("/model_info")
|
||||
async def model_info():
|
||||
"""Get the model information."""
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
|
||||
model_config = _global_state.tokenizer_manager.model_config
|
||||
result = {
|
||||
"model_path": _global_state.tokenizer_manager.model_path,
|
||||
"tokenizer_path": _global_state.tokenizer_manager.server_args.tokenizer_path,
|
||||
"is_generation": _global_state.tokenizer_manager.is_generation,
|
||||
"preferred_sampling_params": _global_state.tokenizer_manager.server_args.preferred_sampling_params,
|
||||
"weight_version": _global_state.tokenizer_manager.server_args.weight_version,
|
||||
"weight_version": get_serving().weight_version,
|
||||
"has_image_understanding": model_config.is_image_understandable_model,
|
||||
"has_audio_understanding": model_config.is_audio_understandable_model,
|
||||
"model_type": getattr(model_config.hf_config, "model_type", None),
|
||||
"architectures": getattr(model_config.hf_config, "architectures", None),
|
||||
"weight_version": _global_state.tokenizer_manager.server_args.weight_version,
|
||||
# "hf_config": model_config.hf_config.to_dict(),
|
||||
}
|
||||
return result
|
||||
@@ -738,12 +739,18 @@ async def server_info():
|
||||
await _global_state.tokenizer_manager.get_internal_state()
|
||||
)
|
||||
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
server_args = _global_state.tokenizer_manager.server_args
|
||||
|
||||
# server_args.model_config is not serializable but should be excluded by asdict.
|
||||
# Overlay post-publish overrides so runtime updates (weight version, model
|
||||
# path/load format) are reported, not the startup record.
|
||||
return msgspec_to_builtins(
|
||||
{
|
||||
**dataclasses.asdict(server_args),
|
||||
**get_context().resolved_server_args_dict(
|
||||
base=dataclasses.asdict(server_args)
|
||||
),
|
||||
**_global_state.scheduler_info,
|
||||
"internal_states": internal_states,
|
||||
"version": __version__,
|
||||
@@ -1368,7 +1375,9 @@ async def update_weight_version(
|
||||
# since weight_version update is a simple operation that doesn't affect model weights
|
||||
try:
|
||||
# Update the weight version in server args (the single source of truth)
|
||||
_global_state.tokenizer_manager.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"http.update_weight_version", weight_version=obj.new_version
|
||||
)
|
||||
|
||||
|
||||
@@ -55,6 +55,8 @@ class HttpServerEngineAdapter(EngineBase):
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.server_args = ServerArgs(**kwargs)
|
||||
# Read host/port from the adapter's own args: no config is published yet
|
||||
# in this process (publish happens in the child from launch_server_process).
|
||||
print(
|
||||
f"Launch HttpServerEngineAdapter at: {self.server_args.host}:{self.server_args.port}"
|
||||
)
|
||||
|
||||
@@ -72,6 +72,7 @@ from sglang.srt.entrypoints.openai.transcription_adapters.base import (
|
||||
TranscriptionAdapter,
|
||||
)
|
||||
from sglang.srt.managers.tokenizer_manager import TokenizerManager
|
||||
from sglang.srt.runtime_context import get_serving
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import random_uuid
|
||||
|
||||
@@ -338,12 +339,12 @@ class RealtimeConnection:
|
||||
if (
|
||||
transcription is not None
|
||||
and transcription.model
|
||||
and transcription.model != self.server_args.served_model_name
|
||||
and transcription.model != get_serving().served_model_name
|
||||
):
|
||||
await self._send_error(
|
||||
"not_supported",
|
||||
f"Model {transcription.model!r} is not served by this endpoint "
|
||||
f"(serving {self.server_args.served_model_name!r}); set "
|
||||
f"(serving {get_serving().served_model_name!r}); set "
|
||||
f"transcription.model to null or to the server's model name.",
|
||||
param="session.audio.input.transcription.model",
|
||||
)
|
||||
|
||||
@@ -16,7 +16,7 @@ from sglang.srt.eplb.expert_location import (
|
||||
get_global_expert_location_metadata,
|
||||
)
|
||||
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_model
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
@@ -274,8 +274,8 @@ def update_expert_location_with_recovery(
|
||||
else:
|
||||
# Load the missing weights from disk
|
||||
update_weights_from_disk_callable(
|
||||
get_server_args().model_path,
|
||||
get_server_args().load_format,
|
||||
get_model().model_path,
|
||||
get_model().load_format,
|
||||
weight_name_filter=weight_name_filter,
|
||||
)
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from typing import Literal, Optional
|
||||
import torch
|
||||
|
||||
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -34,7 +34,7 @@ class ExpertLocationDispatchInfo:
|
||||
|
||||
@classmethod
|
||||
def init_new(cls, layer_id: int):
|
||||
ep_dispatch_algorithm = get_server_args().ep_dispatch_algorithm
|
||||
ep_dispatch_algorithm = get_exec().moe.ep_dispatch_algorithm
|
||||
expert_location_metadata = get_global_expert_location_metadata()
|
||||
assert expert_location_metadata is not None
|
||||
|
||||
|
||||
@@ -26,7 +26,7 @@ from sglang.srt.eplb.expert_location import (
|
||||
ExpertLocationMetadata,
|
||||
get_global_expert_location_metadata,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_device
|
||||
from sglang.srt.utils import get_bool_env_var
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -107,7 +107,7 @@ def _update_expert_weights_with_canary(
|
||||
canary_tensor = (
|
||||
_get_canary_value(old_expert_location_metadata, layer_id)
|
||||
.clone()
|
||||
.to(device=get_server_args().device, non_blocking=True)
|
||||
.to(device=get_device().device, non_blocking=True)
|
||||
)
|
||||
routed_experts_weights_of_layer[layer_id].append(canary_tensor)
|
||||
|
||||
|
||||
@@ -16,9 +16,8 @@ from sglang.srt.hardware_backend.mlx.kv_cache.auxiliary_state import (
|
||||
from sglang.srt.mem_cache.allocator import TokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.memory_pool import KVCache, ReqToTokenPool
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
from sglang.srt.model_executor.model_runner_components.layer_setup import (
|
||||
ModelLayerInfo,
|
||||
)
|
||||
from sglang.srt.model_executor.model_runner_components.layer_setup import ModelLayerInfo
|
||||
from sglang.srt.runtime_context import get_exec, get_memory, get_schedule
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -144,7 +143,7 @@ class MlxModelRunnerStub(ModelRunner):
|
||||
(``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for
|
||||
the mode.
|
||||
"""
|
||||
if self.server_args.disable_radix_cache:
|
||||
if get_memory().disable_radix_cache:
|
||||
return 1
|
||||
return MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO
|
||||
|
||||
@@ -165,7 +164,7 @@ class MlxModelRunnerStub(ModelRunner):
|
||||
Requires ``self.max_total_num_tokens`` to already be set.
|
||||
"""
|
||||
capacity_cap = self.max_total_num_tokens // 2
|
||||
requested = self.server_args.max_running_requests
|
||||
requested = get_schedule().max_running_requests
|
||||
if requested is None:
|
||||
requested_per_worker = None
|
||||
resolved = min(capacity_cap, 4096)
|
||||
@@ -173,7 +172,7 @@ class MlxModelRunnerStub(ModelRunner):
|
||||
requested_per_worker = requested // self.dp_size
|
||||
resolved = min(requested_per_worker, capacity_cap)
|
||||
|
||||
aux_state_size = self.server_args.max_mamba_cache_size
|
||||
aux_state_size = get_schedule().max_mamba_cache_size
|
||||
if (
|
||||
mambaish_config(self.model_config) is not None
|
||||
and aux_state_size is not None
|
||||
@@ -209,7 +208,7 @@ class MlxModelRunnerStub(ModelRunner):
|
||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||
|
||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=self.server_args.enable_memory_saver
|
||||
enable=get_exec().features.enable_memory_saver
|
||||
)
|
||||
|
||||
# Load model (sets metadata only)
|
||||
@@ -241,7 +240,7 @@ class MlxModelRunnerStub(ModelRunner):
|
||||
|
||||
# Create minimal pools
|
||||
if mambaish_config(self.model_config) is not None:
|
||||
auxiliary_state_size = self.server_args.max_mamba_cache_size
|
||||
auxiliary_state_size = get_schedule().max_mamba_cache_size
|
||||
if auxiliary_state_size is None:
|
||||
auxiliary_state_size = (
|
||||
self.max_running_requests * self._aux_state_slots_per_request()
|
||||
@@ -255,7 +254,7 @@ class MlxModelRunnerStub(ModelRunner):
|
||||
# With the radix cache disabled no tree component exists to
|
||||
# release auxiliary slots, so the pool owns their release
|
||||
# (see MlxAuxiliaryStateReqToTokenPool docstring).
|
||||
owns_auxiliary_state_release=self.server_args.disable_radix_cache,
|
||||
owns_auxiliary_state_release=get_memory().disable_radix_cache,
|
||||
)
|
||||
else:
|
||||
self.req_to_token_pool = ReqToTokenPool(
|
||||
|
||||
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_memory, get_model, get_schedule
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -47,25 +48,23 @@ class MlxTpModelWorker(TpModelWorker):
|
||||
def _init_model_runner(self):
|
||||
"""Create MLX runner first (auto-sizes pool), then stub with matching size."""
|
||||
from sglang.srt.hardware_backend.mlx.model_runner import MlxModelRunner
|
||||
from sglang.srt.hardware_backend.mlx.model_runner_stub import (
|
||||
MlxModelRunnerStub,
|
||||
)
|
||||
from sglang.srt.hardware_backend.mlx.model_runner_stub import MlxModelRunnerStub
|
||||
|
||||
logger.info("Initializing MlxModelRunner for end-to-end MLX inference")
|
||||
init_kwargs = dict(
|
||||
model_path=self.server_args.model_path,
|
||||
trust_remote_code=self.server_args.trust_remote_code,
|
||||
disable_radix_cache=self.server_args.disable_radix_cache,
|
||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
||||
quantization=self.server_args.quantization,
|
||||
model_path=get_model().model_path,
|
||||
trust_remote_code=get_model().trust_remote_code,
|
||||
disable_radix_cache=get_memory().disable_radix_cache,
|
||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||
quantization=get_model().quantization,
|
||||
)
|
||||
if self.server_args.max_total_tokens is not None:
|
||||
init_kwargs["pool_size"] = self.server_args.max_total_tokens
|
||||
if get_schedule().max_total_tokens is not None:
|
||||
init_kwargs["pool_size"] = get_schedule().max_total_tokens
|
||||
self._mlx_runner = MlxModelRunner(**init_kwargs)
|
||||
|
||||
self._model_runner = MlxModelRunnerStub(
|
||||
model_config=self.model_config,
|
||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||
gpu_id=self.gpu_id,
|
||||
ps=self.ps,
|
||||
nccl_port=self.nccl_port,
|
||||
|
||||
@@ -19,11 +19,9 @@ from sglang.srt.layers.attention.flashattention_backend import (
|
||||
merge_state_v2_wrapper,
|
||||
)
|
||||
from sglang.srt.layers.radix_attention import AttentionType
|
||||
from sglang.srt.layers.utils.cp_utils import (
|
||||
cp_allgather_and_save_kv_cache,
|
||||
)
|
||||
from sglang.srt.layers.utils.cp_utils import cp_allgather_and_save_kv_cache
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
@@ -515,7 +513,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
||||
and not forward_batch.forward_mode.is_draft_extend_v2()
|
||||
):
|
||||
if forward_batch.attn_attend_prefix_cache:
|
||||
assert not get_server_args().disable_chunked_prefix_cache
|
||||
assert not get_schedule().disable_chunked_prefix_cache
|
||||
assert forward_batch.prefix_chunk_idx is not None
|
||||
assert forward_batch.prefix_chunk_cu_seq_lens is not None
|
||||
assert forward_batch.prefix_chunk_max_seq_lens is not None
|
||||
|
||||
@@ -12,7 +12,7 @@ from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin
|
||||
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
|
||||
from sglang.srt.model_executor.forward_batch_info import DSV4OutCacheLoc, ForwardMode
|
||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
@@ -1362,9 +1362,8 @@ class DeepseekV4AscendAttnBackend(
|
||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||
):
|
||||
B = forward_batch.batch_size
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
n_draft = get_server_args().speculative_num_draft_tokens or 1
|
||||
n_draft = get_spec().speculative_num_draft_tokens or 1
|
||||
actual_q = torch.arange(
|
||||
n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device
|
||||
)
|
||||
@@ -1409,9 +1408,8 @@ class DeepseekV4AscendAttnBackend(
|
||||
forward_batch.forward_mode.is_target_verify()
|
||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||
):
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
max_seqlen_q = get_server_args().speculative_num_draft_tokens or 1
|
||||
max_seqlen_q = get_spec().speculative_num_draft_tokens or 1
|
||||
else:
|
||||
max_seqlen_q = 1
|
||||
return self._kernel_metadata_from_parts(
|
||||
|
||||
@@ -27,7 +27,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
)
|
||||
from sglang.srt.layers.attention.vision import VisionAttention
|
||||
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_mm
|
||||
|
||||
|
||||
class ViTNpuGraphRunner(ViTCudaGraphRunner):
|
||||
@@ -70,7 +70,7 @@ class ViTNpuGraphRunner(ViTCudaGraphRunner):
|
||||
graph = torch_npu.npu.NPUGraph()
|
||||
vit = self.vit
|
||||
|
||||
override_backend = get_server_args().mm_attention_backend
|
||||
override_backend = get_mm().mm_attention_backend
|
||||
with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool):
|
||||
y = None
|
||||
deepstack_outs: List[torch.Tensor] = []
|
||||
|
||||
@@ -17,7 +17,7 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.hardware_backend.npu.utils import npu_format_cast
|
||||
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
|
||||
from sglang.srt.layers.moe.utils import DeepEPMode
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
||||
@@ -57,7 +57,7 @@ def forward_fuseep(
|
||||
envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
|
||||
),
|
||||
num_experts=layer.num_experts,
|
||||
fuse_mode=get_server_args().fuseep_mode,
|
||||
fuse_mode=get_exec().moe.fuseep_mode,
|
||||
)
|
||||
return hidden_states
|
||||
|
||||
@@ -126,7 +126,7 @@ def process_fuseep_weights(layer: torch.nn.Module, weight_prefix: str) -> None:
|
||||
|
||||
Invoked by ``maybe_apply_fuseep_weights`` for both ``"w13"`` and ``"w2"``.
|
||||
"""
|
||||
if get_server_args().fuseep_mode == 1:
|
||||
if get_exec().moe.fuseep_mode == 1:
|
||||
# -- The fused MoE optimization mode "1": dispatch_gmm_combine_decode --
|
||||
if weight_prefix == "w13":
|
||||
cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu()
|
||||
@@ -143,7 +143,7 @@ def process_fuseep_weights(layer: torch.nn.Module, weight_prefix: str) -> None:
|
||||
layer.w2_weight_scale = torch.nn.Parameter(
|
||||
w2_scale.to(torch.float32), requires_grad=False
|
||||
)
|
||||
elif get_server_args().fuseep_mode == 2:
|
||||
elif get_exec().moe.fuseep_mode == 2:
|
||||
# -- The fused MoE optimization mode "2": dispatch_ffn_combine --
|
||||
if weight_prefix == "w13":
|
||||
w13_weight = _release_weight_cache(layer.w13_weight)
|
||||
|
||||
@@ -22,9 +22,7 @@ import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
divide,
|
||||
)
|
||||
from sglang.srt.distributed import divide
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.layers.utils import MultiPlatformOp
|
||||
@@ -33,7 +31,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Phase,
|
||||
check_cuda_graph_backend,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -89,7 +87,7 @@ logger = logging.getLogger(__name__)
|
||||
class SiluAndMul(MultiPlatformOp):
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
if get_server_args().rl_on_policy_target is not None:
|
||||
if get_exec().deterministic.rl_on_policy_target is not None:
|
||||
self._forward_method = self.forward_native
|
||||
elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get():
|
||||
self._forward_method = self.forward_aiter
|
||||
|
||||
@@ -37,10 +37,14 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
||||
get_tc_piecewise_forward_context,
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.state_capturer.indexer_topk import (
|
||||
maybe_capture_indexer_topk,
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk
|
||||
from sglang.srt.utils import (
|
||||
add_prefix,
|
||||
ceil_align,
|
||||
@@ -105,9 +109,7 @@ if is_npu():
|
||||
import torch_npu
|
||||
from sglang.srt.hardware_backend.npu.utils import get_indexer_weight_stream
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_attn_tp_group,
|
||||
)
|
||||
from sglang.srt.distributed import get_attn_tp_group
|
||||
from sglang.srt.distributed.parallel_state import get_pp_group
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.communicator import ScatterMode
|
||||
@@ -458,7 +460,7 @@ class Indexer(MultiPlatformOp):
|
||||
base=rope_theta, # type: ignore
|
||||
rope_scaling=rope_scaling,
|
||||
is_neox_style=is_neox_style,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
self.block_size = block_size
|
||||
self.scale_fmt = scale_fmt
|
||||
@@ -469,7 +471,7 @@ class Indexer(MultiPlatformOp):
|
||||
self.num_local_tokens = getattr(config, "index_local_tokens", 0)
|
||||
|
||||
self.paged_mqa_logits_backend = DSAPagedMQALogitsBackend.resolve(
|
||||
get_server_args().dsa_paged_mqa_logits_backend
|
||||
get_exec().kernel.dsa_paged_mqa_logits_backend
|
||||
)
|
||||
|
||||
@contextlib.contextmanager
|
||||
@@ -1055,7 +1057,7 @@ class Indexer(MultiPlatformOp):
|
||||
total_mem = torch.cuda.get_device_properties(device_index).total_memory
|
||||
|
||||
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION)
|
||||
mem_fraction_static = get_server_args().mem_fraction_static
|
||||
mem_fraction_static = get_schedule().mem_fraction_static
|
||||
if mem_fraction_static is None:
|
||||
static_budget = total_mem_budget
|
||||
else:
|
||||
|
||||
@@ -28,16 +28,14 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph.context
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
|
||||
from sglang.srt.utils import add_prefix, is_cuda, is_hip, is_xpu
|
||||
from sglang.srt.utils.common import is_sm120_supported
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.attention.dsv4.compressor import (
|
||||
CompressorBackendMixin,
|
||||
)
|
||||
from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin
|
||||
from sglang.srt.layers.quantization import QuantizationConfig
|
||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
@@ -129,9 +127,7 @@ def _aiter_fp8_paged_mqa_logits(
|
||||
clean_logits: bool = False,
|
||||
) -> torch.Tensor:
|
||||
"""Wrapper adapting aiter's deepgemm_fp8_paged_mqa_logits to SGLang's interface."""
|
||||
from aiter.ops.triton.attention.pa_mqa_logits import (
|
||||
deepgemm_fp8_paged_mqa_logits,
|
||||
)
|
||||
from aiter.ops.triton.attention.pa_mqa_logits import deepgemm_fp8_paged_mqa_logits
|
||||
|
||||
batch_size = q_fp8.shape[0]
|
||||
next_n = q_fp8.shape[1]
|
||||
@@ -838,9 +834,8 @@ class C4Indexer(nn.Module):
|
||||
self.rotary_emb = rotary_emb
|
||||
self.freqs_cis = freqs_cis
|
||||
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer
|
||||
self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer
|
||||
self.alt_streams = alt_streams
|
||||
|
||||
def compute_q(
|
||||
|
||||
@@ -13,9 +13,7 @@ from sglang.kernels.ops.attention.metadata import (
|
||||
)
|
||||
from sglang.kernels.ops.attention.pa_page_table import _build_pa_page_table
|
||||
from sglang.kernels.ops.attention.utils import assert_buffer_fits
|
||||
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import (
|
||||
build_trtllm_mha_page_table,
|
||||
)
|
||||
from sglang.kernels.ops.kvcache.trtllm_mha_page_table import build_trtllm_mha_page_table
|
||||
from sglang.srt.configs.model_config import AttentionArch
|
||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||
from sglang.srt.layers.cp.base import CPAttentionBackendKind, get_cp_strategy
|
||||
@@ -28,7 +26,7 @@ from sglang.srt.layers.utils.cp_utils 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_server_args
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
from sglang.srt.speculative.ragged_verify import build_ragged_target_verify_geometry
|
||||
from sglang.srt.speculative.spec_info import SpecInput, SpeculativeAlgorithm
|
||||
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
||||
@@ -166,9 +164,12 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
||||
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
||||
from sglang.srt.runtime_context import get_model
|
||||
|
||||
self.kv_cache_dtype_str = get_model().kv_cache_dtype
|
||||
self.kv_cache_dtype_str = getattr(
|
||||
model_runner,
|
||||
"kv_cache_dtype_str",
|
||||
model_runner.server_args.kv_cache_dtype,
|
||||
)
|
||||
self.kv_cache_is_mxfp8 = self.kv_cache_dtype_str == "mxfp8"
|
||||
self.page_size = model_runner.page_size
|
||||
# Static page-table width (upper bound). The device-side page-table build
|
||||
@@ -1479,7 +1480,7 @@ class FlashAttentionBackend(AttentionBackend):
|
||||
):
|
||||
# Do multi-head attention with chunked prefix cache
|
||||
if forward_batch.attn_attend_prefix_cache:
|
||||
assert not get_server_args().disable_chunked_prefix_cache
|
||||
assert not get_schedule().disable_chunked_prefix_cache
|
||||
# MHA for chunked prefix kv cache when running model with MLA
|
||||
assert forward_batch.prefix_chunk_idx is not None
|
||||
assert forward_batch.prefix_chunk_cu_seq_lens is not None
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_disagg, get_exec, get_parallel, get_schedule
|
||||
|
||||
"""
|
||||
Support attention backend for flashinfer MLA.
|
||||
@@ -32,7 +32,7 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_buffer, get_server_args
|
||||
from sglang.srt.runtime_context import get_buffer
|
||||
from sglang.srt.speculative.spec_info import SpecInput
|
||||
from sglang.srt.speculative.spec_utils import (
|
||||
draft_kv_indices_buffer_width,
|
||||
@@ -223,9 +223,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
self.enable_chunk_kv = (
|
||||
not skip_prefill
|
||||
and get_server_args().disaggregation_mode != "decode"
|
||||
and not get_server_args().disable_chunked_prefix_cache
|
||||
and not get_server_args().flashinfer_mla_disable_ragged
|
||||
and get_disagg().disaggregation_mode != "decode"
|
||||
and not get_schedule().disable_chunked_prefix_cache
|
||||
and not get_exec().kernel.flashinfer_mla_disable_ragged
|
||||
)
|
||||
self.page_size = model_runner.page_size
|
||||
|
||||
@@ -401,7 +401,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
||||
prefix_lens = forward_batch.extend_prefix_lens
|
||||
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
||||
use_ragged = (
|
||||
not get_server_args().flashinfer_mla_disable_ragged
|
||||
not get_exec().kernel.flashinfer_mla_disable_ragged
|
||||
and extend_no_prefix
|
||||
# Piecewise cuda graph should use paged prefill to be compatible with prefix cache
|
||||
and not is_in_tc_piecewise_cuda_graph()
|
||||
|
||||
@@ -19,7 +19,7 @@ from sglang.srt.layers.radix_attention import RadixAttention
|
||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
|
||||
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_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_memory, get_server_args
|
||||
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
|
||||
from sglang.srt.speculative.spec_info import SpecInput
|
||||
|
||||
@@ -350,7 +350,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
||||
"""Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu %
|
||||
mamba_track_interval == 0, so force-flush and snapshot fire on the same
|
||||
steps (no off-by-one)."""
|
||||
interval = get_server_args().mamba_track_interval
|
||||
interval = get_exec().mamba.mamba_track_interval
|
||||
if seq_lens_cpu is None:
|
||||
# Should not happen for the supported config; stay safe and never flush.
|
||||
return torch.zeros((bs,), dtype=torch.bool)
|
||||
@@ -764,7 +764,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
||||
# Page-major stores state strided; only the stride-aware Triton causal-conv
|
||||
# reads it (CUDA causal_conv1d garbles it). A model may also force Triton.
|
||||
use_triton_causal_conv = (
|
||||
use_triton_causal_conv or get_server_args().enable_page_major_kv_layout
|
||||
use_triton_causal_conv or get_memory().enable_page_major_kv_layout
|
||||
)
|
||||
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
||||
mixer_out, intermediate_states = mixer.forward(
|
||||
|
||||
@@ -38,7 +38,11 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
|
||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||
is_in_tc_piecewise_cuda_graph,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_buffer, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_buffer,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
)
|
||||
from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2
|
||||
|
||||
if is_flashinfer_available():
|
||||
@@ -197,9 +201,7 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
||||
self.forward_prefill_metadata: Optional[TRTLLMMLAPrefillMetadata] = None
|
||||
self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None
|
||||
|
||||
self.disable_chunked_prefix_cache = (
|
||||
get_server_args().disable_chunked_prefix_cache
|
||||
)
|
||||
self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache
|
||||
|
||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
||||
self.cuda_graph_custom_mask = None
|
||||
|
||||
@@ -15,7 +15,7 @@ from einops import rearrange
|
||||
from sglang.jit_kernel.norm import can_use_fused_inplace_qknorm as can_use_jit_qk_norm
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.models.utils import apply_qk_norm
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_exec, get_mm, get_parallel
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -69,9 +69,7 @@ if _is_npu:
|
||||
if _is_xpu:
|
||||
from sgl_kernel.flash_attn import flash_attn_varlen_func
|
||||
|
||||
from sglang.kernels.ops.attention.prefill_attention import (
|
||||
context_attention_fwd,
|
||||
)
|
||||
from sglang.kernels.ops.attention.prefill_attention import context_attention_fwd
|
||||
from sglang.srt.distributed import (
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
@@ -86,7 +84,6 @@ from sglang.srt.layers.linear import (
|
||||
from sglang.srt.layers.quantization import QuantizationConfig
|
||||
from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb
|
||||
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_pos_emb_native_eager
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||
@@ -1045,7 +1042,7 @@ class VisionAttention(nn.Module):
|
||||
# Select attention backend via a unified method
|
||||
_passed_backend = qkv_backend
|
||||
qkv_backend = self._determine_attention_backend(_passed_backend)
|
||||
if get_server_args().mm_attention_backend is None and _passed_backend is None:
|
||||
if get_mm().mm_attention_backend is None and _passed_backend is None:
|
||||
print_info_once(f"Multimodal attention backend not set. Use {qkv_backend}.")
|
||||
print_info_once(f"Using {qkv_backend} as multimodal attention backend.")
|
||||
|
||||
@@ -1124,7 +1121,7 @@ class VisionAttention(nn.Module):
|
||||
weight_dtype=torch.float32,
|
||||
cast_x_before_out_mul=True,
|
||||
)
|
||||
if get_server_args().rl_on_policy_target is not None
|
||||
if get_exec().deterministic.rl_on_policy_target is not None
|
||||
else {}
|
||||
)
|
||||
q_norm = RMSNorm(
|
||||
@@ -1152,7 +1149,7 @@ class VisionAttention(nn.Module):
|
||||
- CUDA (other): "triton_attn"
|
||||
- Non-CUDA: "sdpa"
|
||||
"""
|
||||
override_backend = get_server_args().mm_attention_backend
|
||||
override_backend = get_mm().mm_attention_backend
|
||||
if override_backend is not None:
|
||||
backend = override_backend
|
||||
elif passed_backend is not None:
|
||||
@@ -1257,7 +1254,7 @@ class VisionAttention(nn.Module):
|
||||
x = x.unsqueeze(0)
|
||||
assert x.dim() == 3, x.shape
|
||||
if (
|
||||
get_server_args().rl_on_policy_target is not None
|
||||
get_exec().deterministic.rl_on_policy_target is not None
|
||||
and position_embeddings is not None
|
||||
):
|
||||
assert isinstance(position_embeddings, tuple), (
|
||||
|
||||
@@ -15,7 +15,7 @@ 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_server_args
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.radix_attention import RadixAttention
|
||||
@@ -69,9 +69,12 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||
self.req_to_token = model_runner.req_to_token_pool.req_to_token
|
||||
self.kv_cache_dtype = model_runner.kv_cache_dtype
|
||||
from sglang.srt.runtime_context import get_model
|
||||
|
||||
self.kv_cache_dtype_str = get_model().kv_cache_dtype
|
||||
self.kv_cache_dtype_str = getattr(
|
||||
model_runner,
|
||||
"kv_cache_dtype_str",
|
||||
model_runner.server_args.kv_cache_dtype,
|
||||
)
|
||||
self.page_size = model_runner.page_size
|
||||
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
|
||||
self.skip_prefill = skip_prefill
|
||||
@@ -640,7 +643,7 @@ class XPUAttentionBackend(AttentionBackend):
|
||||
):
|
||||
# Do multi-head attention with chunked prefix cache
|
||||
if forward_batch.attn_attend_prefix_cache:
|
||||
assert not get_server_args().disable_chunked_prefix_cache
|
||||
assert not get_schedule().disable_chunked_prefix_cache
|
||||
# MHA for chunked prefix kv cache when running model with MLA
|
||||
assert forward_batch.prefix_chunk_idx is not None
|
||||
assert forward_batch.prefix_chunk_cu_seq_lens is not None
|
||||
|
||||
@@ -72,7 +72,13 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
check_cuda_graph_backend,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -170,7 +176,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
|
||||
and batch_size > 0
|
||||
and batch_size <= FUSE_ALLREDUCE_MAX_BATCH_SIZE
|
||||
and not is_dp_attention_enabled()
|
||||
and get_server_args().flashinfer_allreduce_fusion_backend is not None
|
||||
and get_exec().comm.flashinfer_allreduce_fusion_backend is not None
|
||||
and not is_flashinfer_allreduce_unavailable()
|
||||
)
|
||||
|
||||
@@ -186,7 +192,7 @@ def apply_aiter_all_reduce_fusion(input_tensor: torch.Tensor):
|
||||
and total_bytes <= 8 * 1024 * 8192
|
||||
and get_parallel().tp_size != 6
|
||||
and not is_dp_attention_enabled()
|
||||
and get_server_args().enable_aiter_allreduce_fusion
|
||||
and get_exec().comm.enable_aiter_allreduce_fusion
|
||||
)
|
||||
|
||||
|
||||
@@ -274,7 +280,7 @@ class AttnTpContext:
|
||||
and get_moe_a2a_backend().is_none()
|
||||
and not enable_moe_dense_fully_dp()
|
||||
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE)
|
||||
and get_server_args().speculative_algorithm != "EAGLE3"
|
||||
and get_spec().speculative_algorithm != "EAGLE3"
|
||||
)
|
||||
if get_server_args().enable_attn_tp_input_scattered:
|
||||
if not self.allow_input_scattered:
|
||||
@@ -407,7 +413,7 @@ class LayerScatterModes:
|
||||
not context.is_layer_sparse
|
||||
and context.is_next_layer_sparse
|
||||
and enable_moe_dense_fully_dp()
|
||||
and get_server_args().enable_two_batch_overlap
|
||||
and get_exec().overlap.enable_two_batch_overlap
|
||||
)
|
||||
|
||||
@classmethod
|
||||
@@ -467,7 +473,7 @@ class LayerCommunicator:
|
||||
)
|
||||
self._post_init_communicate()
|
||||
self._speculative_algo = SpeculativeAlgorithm.from_string(
|
||||
get_server_args().speculative_algorithm
|
||||
get_spec().speculative_algorithm
|
||||
)
|
||||
|
||||
def _post_init_communicate(self):
|
||||
@@ -815,7 +821,7 @@ class LayerCommunicator:
|
||||
and get_parallel().tp_size != 6
|
||||
and not is_dp_attention_enabled()
|
||||
and get_moe_a2a_backend().is_none()
|
||||
and get_server_args().enable_aiter_allreduce_fusion
|
||||
and get_exec().comm.enable_aiter_allreduce_fusion
|
||||
)
|
||||
)
|
||||
and (not self.is_last_layer)
|
||||
@@ -1120,7 +1126,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
|
||||
if not handled:
|
||||
quantize_communications = (
|
||||
not forward_batch.forward_mode.is_decode_or_idle()
|
||||
and get_server_args().enable_quant_communications
|
||||
and get_exec().comm.enable_quant_communications
|
||||
)
|
||||
if quantize_communications:
|
||||
hidden_states = attention_tensor_model_parallel_quant_all_reduce(
|
||||
|
||||
@@ -48,12 +48,10 @@ from sglang.srt.layers.cp.base import (
|
||||
CPAttentionBackendKind,
|
||||
)
|
||||
from sglang.srt.layers.cp.padding import pad_local_rows
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_allocation_symmetric,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||
from sglang.srt.model_executor.forward_context import get_token_to_kv_pool
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_device, get_parallel
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -208,10 +206,8 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
||||
actual_seq_q_prev_list.append(block_sizes[cp_rank])
|
||||
actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1])
|
||||
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
try:
|
||||
device = torch.device(get_server_args().device)
|
||||
device = torch.device(get_device().device)
|
||||
except Exception:
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
cu_prev = [0] + list(accumulate(actual_seq_q_prev_list))
|
||||
|
||||
@@ -26,7 +26,7 @@ from sglang.kernels.ops.attention.dcp_kernels import (
|
||||
)
|
||||
from sglang.srt.layers.dcp.layout import update_local_kv_lens_for_dcp
|
||||
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_device, get_parallel
|
||||
|
||||
|
||||
def prepare_decode_context_parallel_metadata(
|
||||
@@ -53,12 +53,12 @@ def prepare_decode_context_parallel_metadata(
|
||||
extend_prefix_starts = torch.zeros(
|
||||
len(seq_lens),
|
||||
dtype=torch.int32,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
extend_cu_prefix_lens = torch.zeros(
|
||||
len(seq_lens) + 1,
|
||||
dtype=torch.int32,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
extend_cu_prefix_lens[1:] = torch.cumsum(extend_prefix_lens, dim=0)
|
||||
extend_cu_prefix_lens = extend_cu_prefix_lens[:-1]
|
||||
@@ -67,7 +67,7 @@ def prepare_decode_context_parallel_metadata(
|
||||
dcp_prefix_kv_indices = torch.empty(
|
||||
sum(extend_prefix_lens_cpu),
|
||||
dtype=torch.int32,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)](
|
||||
req_to_token,
|
||||
@@ -81,20 +81,20 @@ def prepare_decode_context_parallel_metadata(
|
||||
dcp_kv_indptr = torch.zeros(
|
||||
len(seq_lens) + 1,
|
||||
dtype=torch.int32,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0)
|
||||
dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)]
|
||||
dcp_kv_indices = torch.zeros(
|
||||
seq_lens_sum,
|
||||
dtype=torch.int32,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
|
||||
extend_cu_lens = torch.zeros(
|
||||
len(seq_lens) + 1,
|
||||
dtype=torch.int32,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
extend_cu_lens[1:] = torch.cumsum(extend_seq_lens, dim=0)
|
||||
extend_cu_lens = extend_cu_lens[:-1]
|
||||
|
||||
@@ -31,7 +31,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
Phase,
|
||||
check_cuda_graph_backend,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -130,9 +130,7 @@ if _is_cuda:
|
||||
# BEFORE the weight multiply, so the multiply is done in the narrow dtype.
|
||||
_jit_rmsnorm_hf_available = False
|
||||
try:
|
||||
from sglang.jit_kernel.rmsnorm_hf import (
|
||||
is_supported_rmsnorm_hf_hidden_size,
|
||||
)
|
||||
from sglang.jit_kernel.rmsnorm_hf import is_supported_rmsnorm_hf_hidden_size
|
||||
from sglang.jit_kernel.rmsnorm_hf import rmsnorm_hf as _jit_rmsnorm_hf
|
||||
|
||||
_jit_rmsnorm_hf_available = True
|
||||
@@ -144,9 +142,7 @@ if _is_cuda:
|
||||
_jit_rmsnorm_hf = None
|
||||
|
||||
from sglang.jit_kernel.norm import fused_add_rmsnorm as _jit_fused_add_rmsnorm
|
||||
from sglang.jit_kernel.norm import (
|
||||
is_supported_jit_fused_add_rmsnorm_hidden_size,
|
||||
)
|
||||
from sglang.jit_kernel.norm import is_supported_jit_fused_add_rmsnorm_hidden_size
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -206,7 +202,7 @@ def _forward_with_allreduce_fusion(
|
||||
return fused_result
|
||||
|
||||
# For AITER route, preserve correctness when fused path is unavailable.
|
||||
if _use_aiter and get_server_args().enable_aiter_allreduce_fusion:
|
||||
if _use_aiter and get_exec().comm.enable_aiter_allreduce_fusion:
|
||||
x = tensor_model_parallel_all_reduce(x)
|
||||
return norm_module.forward(x, residual, None)
|
||||
|
||||
@@ -284,7 +280,7 @@ class RMSNorm(MultiPlatformOp):
|
||||
if (
|
||||
residual is not None
|
||||
or self.cast_x_before_out_mul
|
||||
or get_server_args().rl_on_policy_target == "fsdp"
|
||||
or get_exec().deterministic.rl_on_policy_target == "fsdp"
|
||||
):
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
out = rms_norm_batch_invariant(
|
||||
@@ -391,7 +387,7 @@ class RMSNorm(MultiPlatformOp):
|
||||
if (
|
||||
residual is not None
|
||||
or self.cast_x_before_out_mul
|
||||
or get_server_args().rl_on_policy_target == "fsdp"
|
||||
or get_exec().deterministic.rl_on_policy_target == "fsdp"
|
||||
or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0)
|
||||
):
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
@@ -452,7 +448,7 @@ class RMSNorm(MultiPlatformOp):
|
||||
if (
|
||||
residual is not None
|
||||
or self.cast_x_before_out_mul
|
||||
or get_server_args().rl_on_policy_target == "fsdp"
|
||||
or get_exec().deterministic.rl_on_policy_target == "fsdp"
|
||||
):
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
return rms_norm_batch_invariant(
|
||||
@@ -579,7 +575,10 @@ class RMSNorm(MultiPlatformOp):
|
||||
if self.variance_size_override is not None:
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
if is_batch_invariant_mode_enabled():
|
||||
if residual is not None or get_server_args().rl_on_policy_target == "fsdp":
|
||||
if (
|
||||
residual is not None
|
||||
or get_exec().deterministic.rl_on_policy_target == "fsdp"
|
||||
):
|
||||
return self.forward_native(x, residual, post_residual_addition)
|
||||
return rms_norm_batch_invariant(
|
||||
x,
|
||||
|
||||
@@ -25,9 +25,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_allocation_symmetric,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||
from sglang.srt.layers.moe.utils import should_skip_mlp_all_reduce
|
||||
from sglang.srt.layers.parameter import (
|
||||
BasevLLMParameter,
|
||||
@@ -39,7 +37,7 @@ from sglang.srt.layers.parameter import (
|
||||
_ColumnvLLMParameter,
|
||||
)
|
||||
from sglang.srt.layers.utils import pad_or_narrow_weight
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -759,9 +757,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
shard_offsets.append((i, current_shard_offset, output_size))
|
||||
current_shard_offset += output_size
|
||||
if _is_cpu:
|
||||
from sglang.srt.model_loader.weight_utils import (
|
||||
pad_loaded_weight,
|
||||
)
|
||||
from sglang.srt.model_loader.weight_utils import pad_loaded_weight
|
||||
|
||||
loaded_weight = pad_loaded_weight(
|
||||
loaded_weight, param.output_dim, output_sizes
|
||||
@@ -805,9 +801,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
current_block_offset += shard_block_size
|
||||
|
||||
if _is_cpu:
|
||||
from sglang.srt.model_loader.weight_utils import (
|
||||
pad_loaded_weight,
|
||||
)
|
||||
from sglang.srt.model_loader.weight_utils import pad_loaded_weight
|
||||
|
||||
loaded_weight = pad_loaded_weight(
|
||||
loaded_weight, param.output_dim, shard_block_sizes
|
||||
@@ -1596,7 +1590,7 @@ class RowParallelLinear(LinearBase):
|
||||
quantize_communications = (
|
||||
(
|
||||
not forward_batch.forward_mode.is_decode_or_idle()
|
||||
and get_server_args().enable_quant_communications
|
||||
and get_exec().comm.enable_quant_communications
|
||||
)
|
||||
if forward_batch is not None
|
||||
else False
|
||||
|
||||
@@ -47,7 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
ForwardMode,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.utils.common import (
|
||||
is_cpu,
|
||||
is_npu,
|
||||
@@ -346,7 +346,7 @@ class LogitsProcessor(nn.Module):
|
||||
self.vocab_size = config.vocab_size
|
||||
self.logit_scale = logit_scale
|
||||
self.use_attn_tp_group = get_server_args().enable_dp_lm_head
|
||||
self.use_fp32_lm_head = get_server_args().enable_fp32_lm_head
|
||||
self.use_fp32_lm_head = get_exec().features.enable_fp32_lm_head
|
||||
if self.use_attn_tp_group:
|
||||
self.attn_tp_size = get_parallel().attn_tp_size
|
||||
self.do_tensor_parallel_all_gather = (
|
||||
@@ -370,8 +370,8 @@ class LogitsProcessor(nn.Module):
|
||||
self.final_logit_softcapping = None
|
||||
|
||||
self.return_full_logits = return_full_logits
|
||||
self.enable_mis = get_server_args().enable_mis
|
||||
self.rl_on_policy_target = get_server_args().rl_on_policy_target
|
||||
self.enable_mis = get_exec().features.enable_mis
|
||||
self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target
|
||||
|
||||
self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer(
|
||||
max_tokens=triton_symm_mem_ag.recommended_max_tokens(
|
||||
|
||||
@@ -7,9 +7,7 @@ import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.eplb.expert_distribution import (
|
||||
get_global_expert_distribution_recorder,
|
||||
)
|
||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||
from sglang.srt.eplb.expert_location_dispatch import (
|
||||
ExpertLocationDispatchInfo,
|
||||
topk_ids_logical_to_physical,
|
||||
@@ -22,6 +20,7 @@ from sglang.srt.layers.moe.topk import (
|
||||
remap_topk_for_per_rank_shared_slots,
|
||||
)
|
||||
from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import is_hip, is_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -44,10 +43,9 @@ class HashTopK(nn.Module):
|
||||
):
|
||||
super().__init__()
|
||||
self.layer_id = layer_id
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
self.enable_waterfill = (
|
||||
num_fused_shared_experts > 0 and get_server_args().enable_waterfill
|
||||
num_fused_shared_experts > 0 and get_exec().moe.enable_waterfill
|
||||
)
|
||||
self.waterfill_balancer = None
|
||||
|
||||
|
||||
@@ -28,7 +28,7 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||
from sglang.srt.layers.moe.utils import get_moe_padding_size
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -506,7 +506,7 @@ def _fused_moe_kernel_sequence(
|
||||
out_hidden_states = torch.empty_like(hidden_states)
|
||||
|
||||
use_fused_moe_sum_all_reduce = (
|
||||
get_server_args().enable_fused_moe_sum_all_reduce
|
||||
get_exec().moe.enable_fused_moe_sum_all_reduce
|
||||
and (not no_combine)
|
||||
and (topk > 2)
|
||||
and (not use_int8_w8a16)
|
||||
|
||||
@@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
import torch
|
||||
import triton
|
||||
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import get_device_name, is_hip
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -69,7 +69,7 @@ def get_moe_configs(
|
||||
kernel on a given batch size bs, the closest batch size in the grid should
|
||||
be picked and the associated configuration chosen to invoke the kernel.
|
||||
"""
|
||||
if get_server_args().enable_deterministic_inference:
|
||||
if get_exec().deterministic.enable_deterministic_inference:
|
||||
logger.warning(
|
||||
"Deterministic inference is enabled, using default MoE kernel config."
|
||||
)
|
||||
@@ -187,7 +187,7 @@ def get_default_config(
|
||||
is_marlin: bool,
|
||||
block_shape: Optional[List[int]] = None,
|
||||
) -> Dict[str, int]:
|
||||
if get_server_args().enable_deterministic_inference:
|
||||
if get_exec().deterministic.enable_deterministic_inference:
|
||||
config = {
|
||||
"BLOCK_SIZE_M": 64,
|
||||
"BLOCK_SIZE_N": 64,
|
||||
|
||||
@@ -21,13 +21,9 @@ from sglang.srt.layers.moe.token_dispatcher import (
|
||||
from sglang.srt.layers.moe.token_dispatcher.flashinfer_utils import (
|
||||
TorchDistributedCommBackend,
|
||||
)
|
||||
from sglang.srt.layers.moe.topk import (
|
||||
StandardTopKOutput,
|
||||
TopKOutput,
|
||||
TopKOutputChecker,
|
||||
)
|
||||
from sglang.srt.layers.moe.topk import StandardTopKOutput, TopKOutput, TopKOutputChecker
|
||||
from sglang.srt.layers.moe.utils import get_moe_runner_backend
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_schedule, get_spec
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.utils import get_int_env_var
|
||||
|
||||
@@ -123,7 +119,7 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
# max_running_requests is not yet resolved at model-construction time,
|
||||
# so we use 4096 as a floor to cover decode batches and _dummy_run
|
||||
# (which warms up at batch_size = req_to_token_pool.size).
|
||||
cps = get_server_args().chunked_prefill_size
|
||||
cps = get_schedule().chunked_prefill_size
|
||||
default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096)
|
||||
self.max_num_tokens = get_int_env_var(
|
||||
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK",
|
||||
@@ -132,7 +128,7 @@ class FlashinferDispatcher(BaseDispatcher):
|
||||
|
||||
# Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized.
|
||||
speculative_algo = SpeculativeAlgorithm.from_string(
|
||||
get_server_args().speculative_algorithm
|
||||
get_spec().speculative_algorithm
|
||||
)
|
||||
if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle():
|
||||
total_dispatch_payload_size_per_token = (
|
||||
|
||||
@@ -32,7 +32,7 @@ from typing import (
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_exec, get_lora, get_parallel
|
||||
|
||||
try:
|
||||
from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx
|
||||
@@ -83,9 +83,7 @@ except ImportError:
|
||||
pass
|
||||
|
||||
from sglang.jit_kernel.dsv4 import mask_topk_ids
|
||||
from sglang.srt.distributed import (
|
||||
get_tp_group,
|
||||
)
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
)
|
||||
@@ -98,9 +96,7 @@ from sglang.srt.eplb.expert_location_dispatch import (
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||
from sglang.srt.layers.moe import get_moe_runner_backend
|
||||
from sglang.srt.layers.moe.utils import (
|
||||
has_per_rank_fused_shared_slots,
|
||||
)
|
||||
from sglang.srt.layers.moe.utils import has_per_rank_fused_shared_slots
|
||||
from sglang.srt.layers.utils import MultiPlatformOp
|
||||
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
|
||||
from sglang.srt.utils import (
|
||||
@@ -419,10 +415,9 @@ class TopK(MultiPlatformOp):
|
||||
assert num_expert_group is not None and topk_group is not None
|
||||
|
||||
self.layer_id = layer_id
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
self.enable_waterfill = (
|
||||
num_fused_shared_experts > 0 and get_server_args().enable_waterfill
|
||||
num_fused_shared_experts > 0 and get_exec().moe.enable_waterfill
|
||||
)
|
||||
|
||||
self.waterfill_balancer = None
|
||||
@@ -496,9 +491,8 @@ class TopK(MultiPlatformOp):
|
||||
# ===== TO BE REFACTORED ====
|
||||
elif get_moe_runner_backend().is_experimental_sgl_trtllm():
|
||||
try:
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
use_standard_for_lora = bool(get_server_args().enable_lora)
|
||||
use_standard_for_lora = bool(get_lora().enable_lora)
|
||||
except ValueError:
|
||||
use_standard_for_lora = False
|
||||
output_format = (
|
||||
|
||||
@@ -13,7 +13,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
)
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils.common import torch_release
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -34,7 +34,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
w8a8_block_fp8_matmul_deepgemm,
|
||||
w8a8_block_fp8_matmul_triton,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.utils import (
|
||||
ceil_align,
|
||||
ceil_div,
|
||||
@@ -1470,9 +1469,7 @@ def requant_block_scale_ue8m0_for_deepgemm(
|
||||
scales are not already UE8M0, and DeepGEMM can run the layer (bf16 output,
|
||||
aligned shape). Returns True when it requantizes.
|
||||
"""
|
||||
from sglang.srt.model_loader.utils import (
|
||||
should_deepgemm_weight_requant_ue8m0,
|
||||
)
|
||||
from sglang.srt.model_loader.utils import should_deepgemm_weight_requant_ue8m0
|
||||
|
||||
if (
|
||||
not use_deepgemm_runner
|
||||
@@ -1794,7 +1791,7 @@ def apply_fp8_linear(
|
||||
if (
|
||||
input_scale is not None
|
||||
and input_scale.numel() == 1
|
||||
and get_server_args().cuda_graph_config.prefill.tc_compiler == "inductor"
|
||||
and get_exec().graph.cuda_graph_config.prefill.tc_compiler == "inductor"
|
||||
):
|
||||
qinput = (
|
||||
(input_2d * input_scale.reciprocal())
|
||||
|
||||
@@ -48,7 +48,7 @@ from sglang.srt.layers.quantization.base_config import (
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from sglang.srt.layers.quantization.utils import is_layer_skipped
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
is_cpu,
|
||||
@@ -77,9 +77,7 @@ if is_flashinfer_available():
|
||||
nvfp4_block_scale_interleave,
|
||||
trtllm_fp4_block_scale_moe,
|
||||
)
|
||||
from flashinfer.fused_moe.core import (
|
||||
get_w2_permute_indices_with_cache,
|
||||
)
|
||||
from flashinfer.fused_moe.core import get_w2_permute_indices_with_cache
|
||||
|
||||
# SM90 mixed-input helpers landed in FlashInfer #3084 (post-0.6.10). Older
|
||||
# versions don't ship them; gate at import so unrelated code paths still load.
|
||||
@@ -334,7 +332,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
||||
self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4()
|
||||
self.use_marlin = get_moe_runner_backend().is_marlin()
|
||||
self.flashinfer_mxfp4_moe_precision = (
|
||||
get_server_args().flashinfer_mxfp4_moe_precision
|
||||
get_exec().moe.flashinfer_mxfp4_moe_precision
|
||||
)
|
||||
# When `flashinfer_mxfp4` is enabled, dispatch to one of two FlashInfer
|
||||
# entry points depending on the GPU:
|
||||
|
||||
@@ -14,7 +14,7 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_allocation_symmetric
|
||||
from sglang.srt.layers.moe.utils import RoutingMethodType
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import (
|
||||
is_flashinfer_available,
|
||||
log_info_on_rank0,
|
||||
@@ -51,7 +51,7 @@ class Mxfp4FlashinferTrtllmMoEMethod:
|
||||
self._fp8 = fp8_method
|
||||
self.prefix = prefix
|
||||
self.flashinfer_mxfp4_moe_precision = (
|
||||
get_server_args().flashinfer_mxfp4_moe_precision
|
||||
get_exec().moe.flashinfer_mxfp4_moe_precision
|
||||
)
|
||||
|
||||
def create_moe_runner(self, layer, moe_runner_config):
|
||||
@@ -376,9 +376,7 @@ def maybe_fuse_routed_scale_and_shared_add(
|
||||
from sglang.srt.layers.quantization.mxfp4_flashinfer_cutlass_moe import (
|
||||
Mxfp4FlashinferCutlassMoEMethod,
|
||||
)
|
||||
from sglang.srt.layers.quantization.mxfp4_marlin_moe import (
|
||||
Mxfp4MarlinMoEMethod,
|
||||
)
|
||||
from sglang.srt.layers.quantization.mxfp4_marlin_moe import Mxfp4MarlinMoEMethod
|
||||
|
||||
fused = isinstance(
|
||||
experts.quant_method,
|
||||
|
||||
@@ -11,7 +11,7 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.rotary_embedding.utils import apply_rotary_emb
|
||||
from sglang.srt.layers.utils import MultiPlatformOp
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
get_bool_env_var,
|
||||
@@ -65,9 +65,7 @@ if _is_npu:
|
||||
)
|
||||
|
||||
if _is_hip:
|
||||
from sglang.kernels.ops.attention.utils import (
|
||||
fused_qk_rope_reshape_and_cache,
|
||||
)
|
||||
from sglang.kernels.ops.attention.utils import fused_qk_rope_reshape_and_cache
|
||||
|
||||
if _is_xpu:
|
||||
from sgl_kernel import fused_qk_rope_with_cos_sin_cache_inplace
|
||||
@@ -127,7 +125,7 @@ class RotaryEmbedding(MultiPlatformOp):
|
||||
self._apply_rotary_emb_wrapped = apply_rotary_emb
|
||||
|
||||
# XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend
|
||||
if get_server_args().rl_on_policy_target is not None or _is_musa:
|
||||
if get_exec().deterministic.rl_on_policy_target is not None or _is_musa:
|
||||
self._forward_method = self.forward_native
|
||||
self._apply_rotary_emb_wrapped = torch.compile(
|
||||
dynamic=True,
|
||||
@@ -151,7 +149,7 @@ class RotaryEmbedding(MultiPlatformOp):
|
||||
# create the cache on GPU for faster initialization. This may cause
|
||||
# a slight numerical difference between the HF implementation and ours.
|
||||
init_device = (
|
||||
"cpu" if get_server_args().rl_on_policy_target is not None else None
|
||||
"cpu" if get_exec().deterministic.rl_on_policy_target is not None else None
|
||||
)
|
||||
inv_freq = 1.0 / (
|
||||
base
|
||||
@@ -162,7 +160,7 @@ class RotaryEmbedding(MultiPlatformOp):
|
||||
/ self.rotary_dim
|
||||
)
|
||||
)
|
||||
if get_server_args().rl_on_policy_target is not None:
|
||||
if get_exec().deterministic.rl_on_policy_target is not None:
|
||||
inv_freq = inv_freq.cuda()
|
||||
return inv_freq
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import (
|
||||
yarn_get_mscale_simple,
|
||||
yarn_linear_ramp_mask,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import (
|
||||
cpu_has_amx_support,
|
||||
is_cuda,
|
||||
@@ -42,7 +42,6 @@ if _is_xpu:
|
||||
from sgl_kernel import multimodal_rotary_embedding
|
||||
|
||||
from sglang.kernels.ops.attention.mrope import apply_interleaved_rope_triton
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
|
||||
def apply_interleaved_rope(x: torch.Tensor, mrope_section: list) -> torch.Tensor:
|
||||
@@ -132,7 +131,7 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
self.register_buffer("axis_map", axis_map, persistent=False)
|
||||
else:
|
||||
self.axis_map = None
|
||||
if get_server_args().rl_on_policy_target is not None:
|
||||
if get_exec().deterministic.rl_on_policy_target is not None:
|
||||
self._forward_method = self.forward_native
|
||||
|
||||
def get_cos_sin_with_position(self, positions):
|
||||
@@ -144,7 +143,7 @@ class MRotaryEmbedding(RotaryEmbedding):
|
||||
last_dim = cos_sin.size()[-1]
|
||||
cos, sin = cos_sin.chunk(2, dim=-1)
|
||||
if self.mrope_interleaved:
|
||||
if support_triton(get_server_args().attention_backend):
|
||||
if support_triton(get_exec().kernel.attention_backend):
|
||||
cos = apply_interleaved_rope_triton(cos, self.mrope_section)
|
||||
sin = apply_interleaved_rope_triton(sin, self.mrope_section)
|
||||
else:
|
||||
|
||||
@@ -8,34 +8,21 @@ from torch import nn
|
||||
|
||||
from sglang.kernels.ops.sampling.murmur_hash import murmur_hash32
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.logprob_processor import (
|
||||
OutputLogprobProcessor,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.layers.logprob_processor import OutputLogprobProcessor
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||
from sglang.srt.utils.async_probe import sanitize_nan_logits
|
||||
from sglang.srt.utils.common import (
|
||||
get_bool_env_var,
|
||||
is_cuda,
|
||||
is_hip,
|
||||
is_musa,
|
||||
is_npu,
|
||||
)
|
||||
from sglang.srt.utils.common import get_bool_env_var, is_cuda, is_hip, is_musa, is_npu
|
||||
|
||||
if is_cuda():
|
||||
from flashinfer.sampling import (
|
||||
min_p_sampling_from_probs,
|
||||
top_k_top_p_sampling_from_probs,
|
||||
)
|
||||
from sgl_kernel import (
|
||||
top_k_renorm_prob,
|
||||
top_p_renorm_prob,
|
||||
)
|
||||
from sgl_kernel import top_k_renorm_prob, top_p_renorm_prob
|
||||
|
||||
if is_musa():
|
||||
from sgl_kernel import (
|
||||
@@ -74,12 +61,14 @@ class Sampler(nn.Module):
|
||||
if is_dp_attention_enabled():
|
||||
self.tp_sync_group = get_parallel().attn_tp_group.device_group
|
||||
|
||||
self.rl_on_policy_target = get_server_args().rl_on_policy_target
|
||||
self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target
|
||||
# In RL on-policy mode, deterministic inference is automatically enabled.
|
||||
self.enable_deterministic = get_server_args().enable_deterministic_inference
|
||||
self.enable_deterministic = (
|
||||
get_exec().deterministic.enable_deterministic_inference
|
||||
)
|
||||
# In RL on-policy mode, we use log_softmax to compute logprobs to match the trainer.
|
||||
self.use_log_softmax_logprob = self.rl_on_policy_target is not None
|
||||
self.use_ascend_backend = get_server_args().sampling_backend == "ascend"
|
||||
self.use_ascend_backend = get_exec().kernel.sampling_backend == "ascend"
|
||||
|
||||
self.output_logprob_processor = OutputLogprobProcessor()
|
||||
|
||||
@@ -245,7 +234,7 @@ class Sampler(nn.Module):
|
||||
positions=positions,
|
||||
)
|
||||
else:
|
||||
backend = get_server_args().sampling_backend
|
||||
backend = get_exec().kernel.sampling_backend
|
||||
if backend == "flashinfer":
|
||||
assert (
|
||||
sampling_info.sampling_seed is None
|
||||
|
||||
@@ -48,6 +48,7 @@ from sglang.srt.managers.scheduler import run_scheduler_process
|
||||
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
|
||||
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
|
||||
from sglang.srt.observability.trace import process_tracing_init, trace_set_thread_info
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.server_args import (
|
||||
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||
PortArgs,
|
||||
@@ -231,7 +232,7 @@ class DataParallelController:
|
||||
sock_send(worker, obj)
|
||||
|
||||
def update_active_ranks(self, ranks: ActiveRanksOutput):
|
||||
if self.server_args.elastic_ep_backend is not None:
|
||||
if get_exec().moe.elastic_ep_backend is not None:
|
||||
if len(ranks.status) != self.max_dp_size:
|
||||
logger.warning(
|
||||
"[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; "
|
||||
@@ -484,7 +485,7 @@ class DataParallelController:
|
||||
logger.debug("Worker port broadcast completed")
|
||||
return worker_ports
|
||||
finally:
|
||||
if self.server_args.elastic_ep_backend is None:
|
||||
if get_exec().moe.elastic_ep_backend is None:
|
||||
rep_socket.close()
|
||||
else:
|
||||
threading.Thread(
|
||||
|
||||
@@ -33,7 +33,13 @@ from sglang.srt.managers.schedule_batch import (
|
||||
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.multimodal.evs import EVSEmbeddingResult
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_server_args,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once
|
||||
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
|
||||
from sglang.utils import logger
|
||||
@@ -878,7 +884,7 @@ def _adjust_embedding_length(
|
||||
f"tokens from multimodal embeddings."
|
||||
)
|
||||
if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding:
|
||||
chunked_prefill_size = get_server_args().chunked_prefill_size
|
||||
chunked_prefill_size = get_schedule().chunked_prefill_size
|
||||
if chunked_prefill_size != -1:
|
||||
logger.warning(
|
||||
"You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill"
|
||||
@@ -1287,7 +1293,7 @@ def general_mm_embed_routine(
|
||||
feature = getattr(mm_item, "feature", None)
|
||||
if isinstance(feature, torch.Tensor) and feature.is_cuda:
|
||||
mm_item.feature = feature.to("cpu", non_blocking=True)
|
||||
if get_server_args().language_only:
|
||||
if get_disagg().language_only:
|
||||
precomputed_embeddings = getattr(
|
||||
mm_item, "precomputed_embeddings", None
|
||||
)
|
||||
@@ -1967,7 +1973,7 @@ def wrap_shm_features(obj):
|
||||
"""
|
||||
Scan the object for multimodal tensors and wrap them in SHM pointers.
|
||||
"""
|
||||
if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
|
||||
if _get_is_default_transport() or get_serving().skip_tokenizer_init:
|
||||
return obj
|
||||
|
||||
if obj.mm_inputs:
|
||||
@@ -2028,7 +2034,7 @@ def unwrap_shm_features(obj):
|
||||
Restore ShmPointerMMData wrappers back into standard torch.Tensors.
|
||||
Handles both single requests and batch requests.
|
||||
"""
|
||||
if _get_is_default_transport() or get_server_args().skip_tokenizer_init:
|
||||
if _get_is_default_transport() or get_serving().skip_tokenizer_init:
|
||||
return obj
|
||||
# Handle batch requests
|
||||
if isinstance(obj, BaseBatchReq):
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
@@ -645,15 +647,15 @@ class TokenizerWorker(TokenizerManager):
|
||||
self.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
|
||||
# For PD disaggregtion
|
||||
self.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"tokenizer_worker.restore_disaggregation_mode",
|
||||
disaggregation_mode=disaggregation_mode,
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
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
|
||||
|
||||
@@ -77,10 +77,7 @@ from sglang.srt.managers.embed_types import PositionalEmbeds
|
||||
from sglang.srt.managers.scheduler_components.new_token_ratio_tracker import (
|
||||
NewTokenRatioTracker,
|
||||
)
|
||||
from sglang.srt.mem_cache.allocation import (
|
||||
alloc_for_decode,
|
||||
alloc_for_extend,
|
||||
)
|
||||
from sglang.srt.mem_cache.allocation import alloc_for_decode, alloc_for_extend
|
||||
from sglang.srt.mem_cache.allocation_sizing import get_alloc_reserve_per_decode
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
@@ -105,7 +102,12 @@ from sglang.srt.observability.req_time_stats import (
|
||||
DPControllerReqTimeStats,
|
||||
SchedulerReqTimeStats,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
@@ -1094,7 +1096,7 @@ class Req(ReqDllmMixin):
|
||||
"""Check if this request is prefill-only (no token generation needed)."""
|
||||
# NOTE: when spec is enabled, prefill_only optimizations are disabled
|
||||
|
||||
spec_alg = get_server_args().speculative_algorithm
|
||||
spec_alg = get_spec().speculative_algorithm
|
||||
return self.sampling_params.max_new_tokens == 0 and spec_alg is None
|
||||
|
||||
@property
|
||||
@@ -1115,7 +1117,7 @@ class Req(ReqDllmMixin):
|
||||
def effective_kv_committed_len(self) -> int:
|
||||
# Report only the prompt prefix so thinking + answer fall into the
|
||||
# overallocated range and are reclaimed by release_kv_cache. #22373.
|
||||
if get_server_args().strip_thinking_cache and self.reasoning_tokens > 0:
|
||||
if get_serving().strip_thinking_cache and self.reasoning_tokens > 0:
|
||||
return min(self.kv_committed_len, len(self.origin_input_ids))
|
||||
return self.kv_committed_len
|
||||
|
||||
|
||||
@@ -56,7 +56,7 @@ from sglang.srt.mem_cache.multi_ended_allocator import (
|
||||
UnifiedMambaTokenToKVPoolAllocator,
|
||||
)
|
||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -193,7 +193,7 @@ class SchedulePolicy:
|
||||
if (
|
||||
not isinstance(policy, CacheAwarePolicy)
|
||||
and self.tree_cache.supports_fast_match_prefix()
|
||||
and get_server_args().disaggregation_mode != "decode"
|
||||
and get_disagg().disaggregation_mode != "decode"
|
||||
):
|
||||
for r in waiting_queue:
|
||||
match_prefix_for_req(self.tree_cache, r, include_req=True)
|
||||
|
||||
@@ -210,9 +210,7 @@ from sglang.srt.managers.scheduler_components.pool_stats_observer import (
|
||||
from sglang.srt.managers.scheduler_components.profiler_manager import (
|
||||
SchedulerProfilerManager,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import (
|
||||
SchedulerRecvSkipper,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper
|
||||
from sglang.srt.managers.scheduler_components.request_receiver import (
|
||||
SchedulerRequestReceiver,
|
||||
)
|
||||
@@ -241,7 +239,20 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
|
||||
from sglang.srt.parser.reasoning_parser import ReasoningParser
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.plugins import load_plugins
|
||||
from sglang.srt.runtime_context import get_context, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_context,
|
||||
get_device,
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_lora,
|
||||
get_memory,
|
||||
get_mm,
|
||||
get_observability,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_serving,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
@@ -443,9 +454,9 @@ class Scheduler(
|
||||
attn_tp_cpu_group=self.attn_tp_cpu_group,
|
||||
tp_cpu_group=self.tp_cpu_group,
|
||||
attn_cp_cpu_group=self.attn_cp_cpu_group,
|
||||
enable_metrics=self.server_args.enable_metrics,
|
||||
enable_metrics=get_observability().enable_metrics,
|
||||
enable_kv_cache_events=bool(
|
||||
self.server_args.kv_events_config
|
||||
get_observability().kv_events_config
|
||||
and self.ps.pp_rank == 0
|
||||
and self.ps.attn_tp_rank == 0
|
||||
and self.ps.attn_cp_rank == 0
|
||||
@@ -471,8 +482,8 @@ class Scheduler(
|
||||
self.init_hisparse_coordinator()
|
||||
|
||||
if (
|
||||
self.server_args.disaggregation_mode == "decode"
|
||||
and self.server_args.disaggregation_decode_enable_offload_kvcache
|
||||
get_disagg().disaggregation_mode == "decode"
|
||||
and get_disagg().disaggregation_decode_enable_offload_kvcache
|
||||
):
|
||||
self.decode_offload_manager = DecodeKVCacheOffloadManager(
|
||||
req_to_token_pool=self.req_to_token_pool,
|
||||
@@ -583,7 +594,7 @@ class Scheduler(
|
||||
|
||||
self.dllm_config = ( # For diffusion LLM
|
||||
DllmConfig.from_server_args(self.server_args)
|
||||
if self.server_args.dllm_algorithm is not None
|
||||
if get_exec().dllm.dllm_algorithm is not None
|
||||
else None
|
||||
)
|
||||
|
||||
@@ -611,11 +622,11 @@ class Scheduler(
|
||||
self.ipc_channels = SchedulerIpcChannels.create(
|
||||
port_args=port_args,
|
||||
is_rank_zero=is_rank_zero,
|
||||
skip_tokenizer_init=self.server_args.skip_tokenizer_init,
|
||||
metrics_enabled=self.server_args.enable_metrics
|
||||
skip_tokenizer_init=get_serving().skip_tokenizer_init,
|
||||
metrics_enabled=get_observability().enable_metrics
|
||||
and (
|
||||
self.ps.attn_tp_rank == 0
|
||||
or self.server_args.enable_metrics_for_all_schedulers
|
||||
or get_observability().enable_metrics_for_all_schedulers
|
||||
),
|
||||
enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(),
|
||||
)
|
||||
@@ -631,7 +642,7 @@ class Scheduler(
|
||||
port_args,
|
||||
self.ps.dp_size,
|
||||
dp_rank,
|
||||
publish_interval=self.server_args.load_snapshot_publish_interval,
|
||||
publish_interval=get_observability().load_snapshot_publish_interval,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning("load snapshot writer init failed: %s", e)
|
||||
@@ -641,7 +652,7 @@ class Scheduler(
|
||||
self.ps.pp_rank == 0
|
||||
and self.ps.attn_tp_rank == 0
|
||||
and self.ps.attn_cp_rank == 0
|
||||
and self.server_args.sleep_on_idle
|
||||
and get_device().sleep_on_idle
|
||||
):
|
||||
self.idle_sleeper = IdleSleeper(
|
||||
sockets=[
|
||||
@@ -712,9 +723,9 @@ class Scheduler(
|
||||
)
|
||||
|
||||
# Set reasoning_parser and think_end_id if --reasoning_parser is enabled
|
||||
if self.server_args.reasoning_parser and self.tokenizer:
|
||||
if get_serving().reasoning_parser and self.tokenizer:
|
||||
reasoning_parser = ReasoningParser(
|
||||
model_type=self.server_args.reasoning_parser,
|
||||
model_type=get_serving().reasoning_parser,
|
||||
stream_reasoning=False,
|
||||
tokenizer=self.tokenizer,
|
||||
)
|
||||
@@ -785,7 +796,7 @@ class Scheduler(
|
||||
target_worker=self.tp_worker,
|
||||
)
|
||||
|
||||
if self.server_args.speculative_draft_load_format is not None:
|
||||
if get_spec().speculative_draft_load_format is not None:
|
||||
# Write the draft load_format onto server_args (not just the bag):
|
||||
# the draft worker is built from a copy of self.server_args and
|
||||
# build_load_config reads server_args.load_format, so a bag-only
|
||||
@@ -793,10 +804,10 @@ class Scheduler(
|
||||
# format.
|
||||
self.server_args.override(
|
||||
"scheduler.draft_load_format",
|
||||
load_format=self.server_args.speculative_draft_load_format,
|
||||
load_format=get_spec().speculative_draft_load_format,
|
||||
)
|
||||
logger.info(
|
||||
f"Using draft model load_format: '{self.server_args.speculative_draft_load_format}'"
|
||||
f"Using draft model load_format: '{get_spec().speculative_draft_load_format}'"
|
||||
)
|
||||
|
||||
DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
|
||||
@@ -887,7 +898,7 @@ class Scheduler(
|
||||
# --min-free-slots-delay. Built independently of the prefill delayer.
|
||||
self.min_free_slots_delayer: Optional[MinFreeSlotsDelayer] = None
|
||||
min_free_slots = resolve_min_free_slots(
|
||||
self.server_args.min_free_slots_delay,
|
||||
get_schedule().min_free_slots_delay,
|
||||
self.max_running_requests,
|
||||
is_dflash_family=self.spec_algorithm.is_dflash_family(),
|
||||
)
|
||||
@@ -933,14 +944,14 @@ class Scheduler(
|
||||
if self.ps.tp_rank == 0:
|
||||
logger.info(
|
||||
f"max_total_num_tokens={self.max_total_num_tokens}, "
|
||||
f"chunked_prefill_size={self.server_args.chunked_prefill_size}, "
|
||||
f"chunked_prefill_size={get_schedule().chunked_prefill_size}, "
|
||||
f"max_prefill_tokens={self.max_prefill_tokens}, "
|
||||
f"max_running_requests={self.max_running_requests}, "
|
||||
f"context_len={self.model_config.context_len}, "
|
||||
f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB"
|
||||
)
|
||||
|
||||
if self.server_args.enable_metrics:
|
||||
if get_observability().enable_metrics:
|
||||
self.metrics_collector.emit_constants(
|
||||
max_total_num_tokens=self.max_total_num_tokens,
|
||||
# TODO: max_running_requests_under_SLO has no setter — dead chain.
|
||||
@@ -987,7 +998,7 @@ class Scheduler(
|
||||
self._engine_paused = False
|
||||
|
||||
def init_chunked_prefill(self):
|
||||
self.chunked_prefill_size = self.server_args.chunked_prefill_size
|
||||
self.chunked_prefill_size = get_schedule().chunked_prefill_size
|
||||
uses_transformers_backend = (
|
||||
get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS
|
||||
)
|
||||
@@ -1007,13 +1018,12 @@ class Scheduler(
|
||||
self.chunked_req = None
|
||||
self._pending_chunked_abort_req = None
|
||||
self.is_mixed_chunk = (
|
||||
self.chunked_prefill_size is not None
|
||||
and self.server_args.enable_mixed_chunk
|
||||
self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk
|
||||
)
|
||||
|
||||
# Init the dynamic chunking predictor for PP
|
||||
self.enable_dynamic_chunking = (
|
||||
self.server_args.enable_dynamic_chunking and self.ps.pp_size > 1
|
||||
get_schedule().enable_dynamic_chunking and self.ps.pp_size > 1
|
||||
)
|
||||
if self.enable_dynamic_chunking:
|
||||
try:
|
||||
@@ -1049,8 +1059,8 @@ class Scheduler(
|
||||
)
|
||||
self.prefill_delayer: Optional[PrefillDelayer] = None
|
||||
self.max_prefill_bs: int = 0
|
||||
if self.server_args.enable_prefill_delayer:
|
||||
if self.server_args.disaggregation_mode == "decode":
|
||||
if get_schedule().enable_prefill_delayer:
|
||||
if get_disagg().disaggregation_mode == "decode":
|
||||
logger.info(
|
||||
"Ignoring --enable-prefill-delayer on decode engine "
|
||||
"(no prefill scheduling path; delayer would be a no-op)."
|
||||
@@ -1067,15 +1077,15 @@ class Scheduler(
|
||||
if self.metrics_reporter.enable_metrics
|
||||
else None
|
||||
),
|
||||
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
|
||||
token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark,
|
||||
max_delay_passes=get_schedule().prefill_delayer_max_delay_passes,
|
||||
token_usage_low_watermark=get_schedule().prefill_delayer_token_usage_low_watermark,
|
||||
device=self.tp_group.device,
|
||||
)
|
||||
|
||||
# NOTE: preemption is enabled by default for priority scheduling.
|
||||
self.enable_priority_preemption = (
|
||||
self.enable_priority_scheduling
|
||||
and not self.server_args.disable_priority_preemption
|
||||
and not get_schedule().disable_priority_preemption
|
||||
)
|
||||
|
||||
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
|
||||
@@ -1091,12 +1101,12 @@ class Scheduler(
|
||||
def init_watch_dog_memory_saver_input_blocker(self):
|
||||
# Start watchdog thread
|
||||
self.watchdog = create_scheduler_watchdog(
|
||||
self, watchdog_timeout=self.server_args.watchdog_timeout
|
||||
self, watchdog_timeout=get_device().watchdog_timeout
|
||||
)
|
||||
|
||||
# Init memory saver, profiler and metric stats
|
||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=self.server_args.enable_memory_saver
|
||||
enable=get_exec().features.enable_memory_saver
|
||||
)
|
||||
|
||||
# Init recv skipper and input blocker
|
||||
@@ -1118,11 +1128,9 @@ class Scheduler(
|
||||
self.disagg_decode_prealloc_queue = None
|
||||
self.disagg_decode_transfer_queue = None
|
||||
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
self.transfer_backend = TransferBackend(
|
||||
self.server_args.disaggregation_transfer_backend
|
||||
get_disagg().disaggregation_transfer_backend
|
||||
)
|
||||
|
||||
# todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D?
|
||||
@@ -1192,10 +1200,10 @@ class Scheduler(
|
||||
tp_size=self.ps.tp_size,
|
||||
dp_size=self.server_args.dp_size,
|
||||
gpu_id=self.ps.gpu_id,
|
||||
bootstrap_port=self.server_args.disaggregation_bootstrap_port,
|
||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||
max_total_num_tokens=self.max_total_num_tokens,
|
||||
pp_rank=self.ps.pp_rank,
|
||||
num_reserved_decode_tokens=self.server_args.num_reserved_decode_tokens,
|
||||
num_reserved_decode_tokens=get_disagg().num_reserved_decode_tokens,
|
||||
transfer_backend=self.transfer_backend,
|
||||
)
|
||||
|
||||
@@ -1221,7 +1229,7 @@ class Scheduler(
|
||||
tp_rank=self.ps.tp_rank,
|
||||
tp_size=self.ps.tp_size,
|
||||
gpu_id=self.ps.gpu_id,
|
||||
bootstrap_port=self.server_args.disaggregation_bootstrap_port,
|
||||
bootstrap_port=get_disagg().disaggregation_bootstrap_port,
|
||||
gloo_group=self.attn_tp_cpu_group,
|
||||
max_total_num_tokens=self.max_total_num_tokens,
|
||||
scheduler=self,
|
||||
@@ -1235,11 +1243,10 @@ class Scheduler(
|
||||
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||
|
||||
# Init mm receiver for EPD disaggregation mode
|
||||
if (
|
||||
self.server_args.language_only
|
||||
and self.server_args.encoder_transfer_backend
|
||||
in ["zmq_to_scheduler", "mooncake"]
|
||||
):
|
||||
if get_disagg().language_only and get_disagg().encoder_transfer_backend in [
|
||||
"zmq_to_scheduler",
|
||||
"mooncake",
|
||||
]:
|
||||
self.mm_receiver = create_mm_receiver(
|
||||
self.server_args,
|
||||
dtype=self.model_config.dtype,
|
||||
@@ -1320,7 +1327,7 @@ class Scheduler(
|
||||
|
||||
def init_deterministic_inference_config(self):
|
||||
"""Initialize deterministic inference configuration for different attention backends."""
|
||||
if not self.server_args.enable_deterministic_inference:
|
||||
if not get_exec().deterministic.enable_deterministic_inference:
|
||||
self.truncation_align_size = None
|
||||
return
|
||||
|
||||
@@ -1329,7 +1336,7 @@ class Scheduler(
|
||||
"triton": ("SGLANG_TRITON_PREFILL_TRUNCATION_ALIGN_SIZE", 4096),
|
||||
}
|
||||
env_var, default_size = backend_sizes.get(
|
||||
self.server_args.attention_backend, (None, None)
|
||||
get_exec().kernel.attention_backend, (None, None)
|
||||
)
|
||||
self.truncation_align_size = (
|
||||
get_int_env_var(env_var, default_size) if env_var else None
|
||||
@@ -1725,10 +1732,10 @@ class Scheduler(
|
||||
)
|
||||
|
||||
def init_lora_drainer(self) -> None:
|
||||
if self.server_args.lora_drain_wait_threshold > 0.0:
|
||||
if get_lora().lora_drain_wait_threshold > 0.0:
|
||||
self.lora_drainer = LoRADrainer(
|
||||
self.server_args.max_loras_per_batch,
|
||||
self.server_args.lora_drain_wait_threshold,
|
||||
get_lora().max_loras_per_batch,
|
||||
get_lora().lora_drain_wait_threshold,
|
||||
)
|
||||
else:
|
||||
self.lora_drainer = None
|
||||
@@ -1830,7 +1837,7 @@ class Scheduler(
|
||||
|
||||
def init_kv_events_publisher(self) -> None:
|
||||
self.kv_events_publisher = SchedulerKvEventsPublisher(
|
||||
kv_events_config=self.server_args.kv_events_config,
|
||||
kv_events_config=get_observability().kv_events_config,
|
||||
ps=self.ps,
|
||||
attn_tp_rank=self.ps.attn_tp_rank,
|
||||
attn_cp_rank=self.ps.attn_cp_rank,
|
||||
@@ -2006,7 +2013,7 @@ class Scheduler(
|
||||
return image_inputs
|
||||
|
||||
def _get_multimodal_inputs(self, mm_inputs_dict):
|
||||
if self.server_args.enable_broadcast_mm_inputs_process:
|
||||
if get_mm().enable_broadcast_mm_inputs_process:
|
||||
return self._process_and_broadcast_mm_inputs(mm_inputs_dict)
|
||||
else:
|
||||
return MultimodalInputs.from_processor_output(mm_inputs_dict)
|
||||
@@ -2053,7 +2060,7 @@ class Scheduler(
|
||||
|
||||
def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None:
|
||||
if (
|
||||
self.server_args.elastic_ep_backend is None
|
||||
get_exec().moe.elastic_ep_backend is None
|
||||
or self.disable_radix_cache
|
||||
or not self.tree_cache.is_tree_cache()
|
||||
):
|
||||
@@ -2099,8 +2106,7 @@ class Scheduler(
|
||||
)
|
||||
# Radix-native sessions use only the top-level session_id.
|
||||
radix_native_session = (
|
||||
recv_req.session_id is not None
|
||||
and self.server_args.enable_session_radix_cache
|
||||
recv_req.session_id is not None and get_memory().enable_session_radix_cache
|
||||
)
|
||||
|
||||
if session_id is None or radix_native_session:
|
||||
@@ -2112,7 +2118,7 @@ class Scheduler(
|
||||
|
||||
if recv_req.bootstrap_port is None:
|
||||
# Use default bootstrap port
|
||||
recv_req.bootstrap_port = self.server_args.disaggregation_bootstrap_port
|
||||
recv_req.bootstrap_port = get_disagg().disaggregation_bootstrap_port
|
||||
|
||||
req = Req(
|
||||
recv_req.rid,
|
||||
@@ -2265,7 +2271,7 @@ class Scheduler(
|
||||
self._add_request_to_queue(req)
|
||||
return
|
||||
|
||||
if req.return_sampling_mask and self.server_args.sampling_backend == "ascend":
|
||||
if req.return_sampling_mask and get_exec().kernel.sampling_backend == "ascend":
|
||||
# The ascend backend samples from logits directly and never builds the
|
||||
# top-k/top-p support, so it cannot produce a sampling mask.
|
||||
error_msg = (
|
||||
@@ -2314,7 +2320,7 @@ class Scheduler(
|
||||
error_msg = validate_input_length(
|
||||
req,
|
||||
self.max_req_input_len,
|
||||
self.server_args.allow_auto_truncate,
|
||||
get_serving().allow_auto_truncate,
|
||||
)
|
||||
if error_msg:
|
||||
req.set_finish_with_abort(error_msg)
|
||||
@@ -2592,7 +2598,7 @@ class Scheduler(
|
||||
error_msg = validate_input_length(
|
||||
req,
|
||||
self.max_req_input_len,
|
||||
self.server_args.allow_auto_truncate,
|
||||
get_serving().allow_auto_truncate,
|
||||
)
|
||||
if error_msg:
|
||||
self._add_request_to_queue(req)
|
||||
@@ -2804,7 +2810,7 @@ class Scheduler(
|
||||
if (
|
||||
need_mlp_sync
|
||||
and not self.spec_algorithm.is_none()
|
||||
and not self.server_args.speculative_skip_dp_mlp_sync
|
||||
and not get_spec().speculative_skip_dp_mlp_sync
|
||||
):
|
||||
# NOTE: This branch makes sure prefill and decode batches will not be mixed when spec and dp-attn is enabled.
|
||||
# Before merging the new batch into running batch:
|
||||
@@ -2878,7 +2884,7 @@ class Scheduler(
|
||||
for req in ready_grammar_requests:
|
||||
self._add_request_to_queue(req)
|
||||
|
||||
if self.enable_hierarchical_cache or self.server_args.enable_flexkv:
|
||||
if self.enable_hierarchical_cache or get_memory().enable_flexkv:
|
||||
self.tree_cache.check_hicache_events()
|
||||
|
||||
if self.enable_priority_preemption or self.is_hybrid_swa:
|
||||
@@ -2945,7 +2951,7 @@ class Scheduler(
|
||||
self.priority_scheduling_preemption_threshold,
|
||||
max_prefill_bs=self.max_prefill_bs,
|
||||
max_running_requests=self.max_running_requests,
|
||||
prefill_max_requests=self.server_args.prefill_max_requests,
|
||||
prefill_max_requests=get_schedule().prefill_max_requests,
|
||||
prefill_delayer_single_pass=prefill_delayer_single_pass,
|
||||
dllm_config=self.dllm_config,
|
||||
waiting_queue_len=len(self.waiting_queue),
|
||||
@@ -3516,7 +3522,7 @@ class Scheduler(
|
||||
|
||||
def _maybe_report_active_ranks(self) -> None:
|
||||
if not (
|
||||
self.enable_dp_attention and self.server_args.elastic_ep_backend is not None
|
||||
self.enable_dp_attention and get_exec().moe.elastic_ep_backend is not None
|
||||
):
|
||||
return
|
||||
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||
@@ -3792,7 +3798,7 @@ class Scheduler(
|
||||
ok, msg = self.tree_cache.attach_storage_backend(
|
||||
storage_backend=recv_req.hicache_storage_backend,
|
||||
storage_backend_extra_config_json=recv_req.hicache_storage_backend_extra_config_json,
|
||||
served_model_name=self.server_args.served_model_name,
|
||||
served_model_name=get_serving().served_model_name,
|
||||
hicache_storage_prefetch_policy=recv_req.hicache_storage_prefetch_policy,
|
||||
hicache_write_policy=recv_req.hicache_write_policy,
|
||||
)
|
||||
@@ -3912,7 +3918,7 @@ class Scheduler(
|
||||
}
|
||||
ret["effective_max_running_requests_per_dp"] = self.max_running_requests
|
||||
|
||||
if self.server_args.elastic_ep_backend is not None:
|
||||
if get_exec().moe.elastic_ep_backend is not None:
|
||||
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||
|
||||
ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling()
|
||||
@@ -4445,10 +4451,10 @@ class Scheduler(
|
||||
return None
|
||||
|
||||
def close_session(self, recv_req: CloseSessionReqInput):
|
||||
if self.server_args.enable_session_radix_cache:
|
||||
if get_memory().enable_session_radix_cache:
|
||||
self.tree_cache.release_radix_session(recv_req.session_id)
|
||||
if recv_req.session_id in self.session_controller or not (
|
||||
self.server_args.enable_session_radix_cache
|
||||
get_memory().enable_session_radix_cache
|
||||
):
|
||||
self.session_controller.close(recv_req)
|
||||
|
||||
|
||||
@@ -2,14 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
|
||||
@@ -23,11 +16,14 @@ from sglang.srt.managers.schedule_batch import (
|
||||
ScheduleBatch,
|
||||
mamba_lazy_spec_in_window,
|
||||
)
|
||||
from sglang.srt.mem_cache.common import (
|
||||
maybe_cache_unfinished_req,
|
||||
release_kv_cache,
|
||||
from sglang.srt.mem_cache.common import maybe_cache_unfinished_req, release_kv_cache
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_observability,
|
||||
get_server_args,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
||||
from sglang.srt.state_capturer.indexer_topk import get_global_indexer_capturer
|
||||
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
|
||||
@@ -48,10 +44,7 @@ if TYPE_CHECKING:
|
||||
SchedulerOutputStreamer,
|
||||
)
|
||||
from sglang.srt.managers.tp_worker import BaseTpWorker
|
||||
from sglang.srt.managers.utils import (
|
||||
EmbeddingBatchResult,
|
||||
GenerationBatchResult,
|
||||
)
|
||||
from sglang.srt.managers.utils import EmbeddingBatchResult, GenerationBatchResult
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
@@ -84,7 +77,7 @@ class SchedulerBatchResultProcessor:
|
||||
|
||||
def process_batch_result_prebuilt(self, batch: ScheduleBatch):
|
||||
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
||||
use_free_group = self.server_args.disaggregation_decode_enable_radix_cache
|
||||
use_free_group = get_disagg().disaggregation_decode_enable_radix_cache
|
||||
if use_free_group:
|
||||
self.token_to_kv_pool_allocator.free_group_begin()
|
||||
for req in batch.reqs:
|
||||
@@ -92,7 +85,7 @@ class SchedulerBatchResultProcessor:
|
||||
req.update_finish_state()
|
||||
if req.finished():
|
||||
req.time_stats.set_quick_finish_time()
|
||||
if self.server_args.enable_hisparse:
|
||||
if get_memory().enable_hisparse:
|
||||
self.hisparse_coordinator.request_finished(req)
|
||||
release_kv_cache(req, self.tree_cache)
|
||||
|
||||
@@ -243,7 +236,7 @@ class SchedulerBatchResultProcessor:
|
||||
req.time_stats.set_completion_time()
|
||||
elif not batch.decoding_reqs or req not in batch.decoding_reqs:
|
||||
maybe_cache_unfinished_req(req, self.tree_cache)
|
||||
if self.server_args.enable_hisparse:
|
||||
if get_memory().enable_hisparse:
|
||||
self.hisparse_coordinator.admit_request_into_staging(req)
|
||||
|
||||
self._maybe_collect_customized_info(i, req, logits_output)
|
||||
@@ -756,7 +749,7 @@ class SchedulerBatchResultProcessor:
|
||||
num_block_accept_tokens=result.num_block_accept_tokens,
|
||||
num_cap_tokens=result.num_cap_tokens,
|
||||
)
|
||||
if self.server_args.enable_metrics:
|
||||
if get_observability().enable_metrics:
|
||||
self.metrics_collector.increment_decode_cuda_graph_pass(
|
||||
value=can_run_cuda_graph
|
||||
)
|
||||
@@ -939,7 +932,7 @@ class SchedulerBatchResultProcessor:
|
||||
self._mamba_prefix_cache_update(req, batch, result, i)
|
||||
|
||||
if (
|
||||
self.server_args.disaggregation_decode_enable_offload_kvcache
|
||||
get_disagg().disaggregation_decode_enable_offload_kvcache
|
||||
and not req.finished()
|
||||
):
|
||||
self.decode_offload_manager.offload_kv_cache(req)
|
||||
@@ -959,12 +952,12 @@ class SchedulerBatchResultProcessor:
|
||||
self._maybe_collect_routed_experts(req)
|
||||
self._maybe_collect_indexer_topk(req)
|
||||
|
||||
if self.server_args.disaggregation_decode_enable_offload_kvcache:
|
||||
if get_disagg().disaggregation_decode_enable_offload_kvcache:
|
||||
# Asynchronously offload KV cache; release_kv_cache will be called after Device->Host transfer completes
|
||||
if not self.decode_offload_manager.offload_kv_cache(req):
|
||||
self.decode_offload_manager.finalize_release_on_finish(req)
|
||||
else:
|
||||
if self.server_args.enable_hisparse:
|
||||
if get_memory().enable_hisparse:
|
||||
self.hisparse_coordinator.request_finished(req)
|
||||
prepare_release = getattr(
|
||||
self.model_worker, "prepare_for_kv_cache_release", None
|
||||
@@ -1102,7 +1095,7 @@ class SchedulerBatchResultProcessor:
|
||||
For spec decode, the boundary is detected by comparing the
|
||||
accepted seq_len range against interval boundaries.
|
||||
"""
|
||||
interval = get_server_args().mamba_track_interval
|
||||
interval = get_exec().mamba.mamba_track_interval
|
||||
|
||||
if batch.spec_algorithm.is_none():
|
||||
if req.kv_committed_len % interval == 0:
|
||||
|
||||
@@ -12,9 +12,7 @@ from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import world_dp_gather_enabled
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import (
|
||||
SchedulerRecvSkipper,
|
||||
)
|
||||
from sglang.srt.managers.scheduler_components.recv_skipper import SchedulerRecvSkipper
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||
@@ -26,6 +24,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.observability.metrics_collector import DPCooperationInfo
|
||||
from sglang.srt.runtime_context import get_schedule
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||
@@ -385,7 +384,7 @@ class SchedulerDPAttnAdapter:
|
||||
get_idle_batch=self.get_idle_batch,
|
||||
disable_cuda_graph=cuda_graph_fully_disabled(),
|
||||
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
|
||||
disable_overlap_schedule=self.server_args.disable_overlap_schedule,
|
||||
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||
offload_tags=self.offload_tags,
|
||||
dwdp=self.server_args.dwdp_size > 1,
|
||||
)
|
||||
|
||||
@@ -14,6 +14,7 @@ from sglang.srt.managers.load_snapshot import (
|
||||
QueueMetrics,
|
||||
SpeculativeMetrics,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_lora
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
@@ -144,7 +145,7 @@ class SchedulerLoadInquirer:
|
||||
)
|
||||
|
||||
lora = None
|
||||
if self.server_args.enable_lora:
|
||||
if get_lora().enable_lora:
|
||||
lora = LoRAMetrics(
|
||||
slots_used=stats.lora_pool_slots_used,
|
||||
slots_total=stats.lora_pool_slots_total,
|
||||
|
||||
@@ -1,20 +1,15 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import (
|
||||
List,
|
||||
Tuple,
|
||||
)
|
||||
from typing import List, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.configs.model_config import ModelConfig
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.managers.schedule_batch import Req
|
||||
from sglang.srt.server_args import (
|
||||
MIS_DELIMITER_TOKEN_ID,
|
||||
ServerArgs,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.server_args import MIS_DELIMITER_TOKEN_ID, ServerArgs
|
||||
|
||||
|
||||
@dataclass(kw_only=True, slots=True, frozen=True)
|
||||
@@ -164,7 +159,7 @@ class SchedulerLogprobResultProcessor:
|
||||
delimiter token receive logprobs.
|
||||
"""
|
||||
return (
|
||||
self.server_args.enable_mis
|
||||
get_exec().features.enable_mis
|
||||
and req.is_prefill_only
|
||||
and req.multi_item_delimiter_indices is not None
|
||||
)
|
||||
|
||||
@@ -2,12 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import (
|
||||
Any,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
)
|
||||
from typing import Any, Callable, List, Optional
|
||||
|
||||
import torch
|
||||
import zmq
|
||||
@@ -21,11 +16,9 @@ from sglang.srt.managers.io_struct import (
|
||||
CachedTokensDetails,
|
||||
wrap_as_pickle,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
BaseFinishReason,
|
||||
Req,
|
||||
)
|
||||
from sglang.srt.managers.schedule_batch import BaseFinishReason, Req
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
||||
from sglang.srt.runtime_context import get_observability, get_serving
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
|
||||
@@ -144,7 +137,7 @@ class SchedulerOutputStreamer:
|
||||
return_sampling_mask=return_sampling_mask,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
disaggregation_mode=self.disaggregation_mode,
|
||||
default_stream_interval=self.server_args.stream_interval,
|
||||
default_stream_interval=get_serving().stream_interval,
|
||||
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
|
||||
get_cached_tokens_details=self.get_cached_tokens_details,
|
||||
)
|
||||
@@ -171,7 +164,7 @@ class SchedulerOutputStreamer:
|
||||
if (
|
||||
req.finished()
|
||||
and self.ps.attn_tp_rank == 0
|
||||
and self.server_args.enable_request_time_stats_logging
|
||||
and get_observability().enable_request_time_stats_logging
|
||||
):
|
||||
req.log_time_stats()
|
||||
|
||||
|
||||
@@ -5,13 +5,7 @@ import os
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Callable, List, Optional
|
||||
|
||||
import torch
|
||||
|
||||
@@ -19,7 +13,7 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_device
|
||||
from sglang.srt.utils import is_mps, is_npu
|
||||
from sglang.srt.utils.profile_merger import ProfileMerger
|
||||
from sglang.srt.utils.profile_utils import ProfileManager
|
||||
@@ -255,7 +249,7 @@ class SchedulerProfilerManager:
|
||||
self.profile_in_progress = True
|
||||
|
||||
if "CUDA_PROFILER" in activities:
|
||||
if self.ps.gpu_id == get_server_args().base_gpu_id:
|
||||
if self.ps.gpu_id == get_device().base_gpu_id:
|
||||
torch.cuda.cudart().cudaProfilerStart()
|
||||
self.profile_in_progress = True
|
||||
|
||||
@@ -365,7 +359,7 @@ class SchedulerProfilerManager:
|
||||
torch.cuda.memory._record_memory_history(enabled=None)
|
||||
|
||||
if "CUDA_PROFILER" in self.profiler_activities:
|
||||
if self.ps.gpu_id == get_server_args().base_gpu_id:
|
||||
if self.ps.gpu_id == get_device().base_gpu_id:
|
||||
torch.cuda.cudart().cudaProfilerStop()
|
||||
|
||||
merge_message = self._merge_profile_traces()
|
||||
|
||||
@@ -2,14 +2,7 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from http import HTTPStatus
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
Callable,
|
||||
List,
|
||||
Optional,
|
||||
Union,
|
||||
)
|
||||
from typing import TYPE_CHECKING, Any, Callable, List, Optional, Union
|
||||
|
||||
import zmq
|
||||
from torch.distributed import barrier
|
||||
@@ -22,14 +15,9 @@ from sglang.srt.managers.io_struct import (
|
||||
TokenizedGenerateReqInput,
|
||||
sock_recv,
|
||||
)
|
||||
from sglang.srt.managers.mm_utils import (
|
||||
has_shm_features,
|
||||
unwrap_shm_features,
|
||||
)
|
||||
from sglang.srt.utils import (
|
||||
broadcast_pyobj,
|
||||
point_to_point_pyobj,
|
||||
)
|
||||
from sglang.srt.managers.mm_utils import has_shm_features, unwrap_shm_features
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
from sglang.srt.utils import broadcast_pyobj, point_to_point_pyobj
|
||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -220,8 +208,8 @@ class SchedulerRequestReceiver:
|
||||
# Process MM requests under EPD-disaggregation mode
|
||||
if (
|
||||
self.ps.pp_rank == 0
|
||||
and self.server_args.language_only
|
||||
and self.server_args.encoder_transfer_backend
|
||||
and get_disagg().language_only
|
||||
and get_disagg().encoder_transfer_backend
|
||||
in ["zmq_to_scheduler", "mooncake"]
|
||||
):
|
||||
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
|
||||
|
||||
@@ -36,6 +36,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
||||
from sglang.srt.runtime_context import get_disagg
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
|
||||
from sglang.srt.utils.common import get_device_module, is_xpu
|
||||
@@ -479,7 +480,7 @@ class SchedulerPPMixin:
|
||||
)
|
||||
)
|
||||
|
||||
if self.server_args.disaggregation_decode_enable_offload_kvcache:
|
||||
if get_disagg().disaggregation_decode_enable_offload_kvcache:
|
||||
self.decode_offload_manager.check_offload_progress()
|
||||
|
||||
if rmbs[next_mb_id] is not None:
|
||||
@@ -549,7 +550,7 @@ class SchedulerPPMixin:
|
||||
+ len(self.disagg_decode_transfer_queue.queue)
|
||||
+ len(self.disagg_decode_prealloc_queue.queue)
|
||||
)
|
||||
if self.server_args.disaggregation_decode_enable_offload_kvcache:
|
||||
if get_disagg().disaggregation_decode_enable_offload_kvcache:
|
||||
queue_size += len(self.decode_offload_manager.ongoing_offload)
|
||||
|
||||
if server_is_idle and queue_size == 0:
|
||||
|
||||
@@ -74,6 +74,7 @@ from sglang.srt.managers.io_struct import (
|
||||
UpdateWeightsFromTensorReqOutput,
|
||||
)
|
||||
from sglang.srt.managers.load_snapshot import LoadSnapshot
|
||||
from sglang.srt.runtime_context import get_lora
|
||||
from sglang.srt.server_args import LoRARef, ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
get_bool_env_var,
|
||||
@@ -569,7 +570,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."
|
||||
)
|
||||
@@ -602,10 +603,10 @@ class TokenizerControlMixin:
|
||||
await self.lora_registry.register(new_adapter)
|
||||
self.lora_ref_cache[obj.lora_name] = new_adapter
|
||||
|
||||
if self.server_args.max_loaded_loras is not None:
|
||||
if get_lora().max_loaded_loras is not None:
|
||||
while (
|
||||
self.lora_registry.num_registered_loras
|
||||
> self.server_args.max_loaded_loras
|
||||
> get_lora().max_loaded_loras
|
||||
):
|
||||
lru_lora_name = await self.lora_registry.lru_lora_name(
|
||||
exclude_pinned=True
|
||||
@@ -619,7 +620,7 @@ class TokenizerControlMixin:
|
||||
logger.info(
|
||||
f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
|
||||
f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
|
||||
f"max allowed: {self.server_args.max_loaded_loras})"
|
||||
f"max allowed: {get_lora().max_loaded_loras})"
|
||||
)
|
||||
|
||||
unload_result = await self._unload_lora_adapter_locked(
|
||||
@@ -647,7 +648,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."
|
||||
)
|
||||
@@ -672,10 +673,10 @@ class TokenizerControlMixin:
|
||||
if result.success:
|
||||
await self.lora_registry.register(new_adapter)
|
||||
self.lora_ref_cache[obj.lora_name] = new_adapter
|
||||
if self.server_args.max_loaded_loras is not None:
|
||||
if get_lora().max_loaded_loras is not None:
|
||||
while (
|
||||
self.lora_registry.num_registered_loras
|
||||
> self.server_args.max_loaded_loras
|
||||
> get_lora().max_loaded_loras
|
||||
):
|
||||
lru_lora_name = await self.lora_registry.lru_lora_name(
|
||||
exclude_pinned=True
|
||||
@@ -689,7 +690,7 @@ class TokenizerControlMixin:
|
||||
logger.info(
|
||||
f"Unloading least recently used LoRA adapter '{lru_lora_name}' "
|
||||
f"(current number of adapters: {self.lora_registry.num_registered_loras}, "
|
||||
f"max allowed: {self.server_args.max_loaded_loras})"
|
||||
f"max allowed: {get_lora().max_loaded_loras})"
|
||||
)
|
||||
|
||||
unload_result = await self._unload_lora_adapter_locked(
|
||||
@@ -717,7 +718,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."
|
||||
)
|
||||
@@ -893,6 +894,8 @@ class TokenizerControlMixin:
|
||||
) -> None:
|
||||
"""Update weight version if provided."""
|
||||
if weight_version is not None:
|
||||
self.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"tokenizer.weight_version", weight_version=weight_version
|
||||
)
|
||||
|
||||
@@ -110,6 +110,14 @@ from sglang.srt.observability.request_metrics_exporter import (
|
||||
RequestMetricsExporterManager,
|
||||
)
|
||||
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_disagg,
|
||||
get_lora,
|
||||
get_model,
|
||||
get_observability,
|
||||
get_serving,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||
from sglang.srt.server_args import (
|
||||
PortArgs,
|
||||
@@ -463,10 +471,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# TODO: Refactor and organize the log export code.
|
||||
# Request logging
|
||||
self.request_logger = RequestLogger(
|
||||
log_requests=self.server_args.log_requests,
|
||||
log_requests_level=self.server_args.log_requests_level,
|
||||
log_requests_format=self.server_args.log_requests_format,
|
||||
log_requests_target=self.server_args.log_requests_target,
|
||||
log_requests=get_observability().log_requests,
|
||||
log_requests_level=get_observability().log_requests_level,
|
||||
log_requests_format=get_observability().log_requests_format,
|
||||
log_requests_target=get_observability().log_requests_target,
|
||||
)
|
||||
|
||||
# Dumping
|
||||
@@ -489,7 +497,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
def init_weight_update(self):
|
||||
# Initial weights status
|
||||
self.initial_weights_loaded = True
|
||||
if self.server_args.checkpoint_engine_wait_weights_before_ready:
|
||||
if get_model().checkpoint_engine_wait_weights_before_ready:
|
||||
self.initial_weights_loaded = False
|
||||
|
||||
# Weight updates
|
||||
@@ -509,7 +517,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# The registry dynamically updates as adapters are loaded / unloaded during runtime. It
|
||||
# serves as the source of truth for available adapters and maps user-friendly LoRA names
|
||||
# to internally used unique LoRA IDs.
|
||||
self.lora_registry = LoRARegistry(self.server_args.lora_paths)
|
||||
self.lora_registry = LoRARegistry(get_lora().lora_paths)
|
||||
# Lock to serialize LoRA update operations.
|
||||
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
|
||||
# LoRA updates and inference to overlap.
|
||||
@@ -518,15 +526,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# point to their latest LoRARef objects, so that they can be
|
||||
# dynamically loaded if needed for inference
|
||||
self.lora_ref_cache: Dict[str, LoRARef] = {}
|
||||
if self.server_args.lora_paths is not None:
|
||||
for lora_ref in self.server_args.lora_paths:
|
||||
if get_lora().lora_paths is not None:
|
||||
for lora_ref in get_lora().lora_paths:
|
||||
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
|
||||
|
||||
def init_disaggregation(self):
|
||||
# PD Disaggregation
|
||||
self.disaggregation_mode = DisaggregationMode(
|
||||
self.server_args.disaggregation_mode
|
||||
)
|
||||
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||
# Keep a reference so the bootstrap server is not garbage-collected.
|
||||
self.bootstrap_server = start_disagg_service(self.server_args)
|
||||
# Single-source counter for auto-assigning fake bootstrap_room.
|
||||
@@ -535,18 +541,16 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# Encoder Disaggregation
|
||||
self.encoder_bootstrap_server = None
|
||||
if self.server_args.language_only:
|
||||
from sglang.srt.disaggregation.encode_receiver import (
|
||||
EncoderBootstrapServer,
|
||||
)
|
||||
from sglang.srt.disaggregation.encode_receiver import EncoderBootstrapServer
|
||||
|
||||
# Shared mutable URL list: the bootstrap server appends / removes
|
||||
# entries as encoders register, the receiver reads from the same
|
||||
# list. Pre-populated with static --encoder-urls so the legacy
|
||||
# CLI flag still works (alongside dynamic registrations).
|
||||
self.encoder_urls: List[str] = list(self.server_args.encoder_urls)
|
||||
self.encoder_urls: List[str] = list(get_disagg().encoder_urls)
|
||||
self.encoder_bootstrap_server = EncoderBootstrapServer(
|
||||
host=self.server_args.host,
|
||||
port=self.server_args.encoder_bootstrap_port,
|
||||
host=get_serving().host,
|
||||
port=get_disagg().encoder_bootstrap_port,
|
||||
urls=self.encoder_urls,
|
||||
)
|
||||
self.mm_receiver = create_mm_receiver(
|
||||
@@ -560,20 +564,22 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
# Metrics
|
||||
if self.enable_metrics:
|
||||
engine_type = DisaggregationMode.to_engine_type(
|
||||
self.server_args.disaggregation_mode
|
||||
get_disagg().disaggregation_mode
|
||||
)
|
||||
|
||||
labels = {
|
||||
"model_name": self.server_args.served_model_name,
|
||||
"model_name": get_serving().served_model_name,
|
||||
"engine_type": engine_type,
|
||||
}
|
||||
if self.enable_priority_scheduling:
|
||||
labels["priority"] = ""
|
||||
if self.server_args.tokenizer_metrics_allowed_custom_labels:
|
||||
for label in self.server_args.tokenizer_metrics_allowed_custom_labels:
|
||||
if get_observability().tokenizer_metrics_allowed_custom_labels:
|
||||
for (
|
||||
label
|
||||
) in get_observability().tokenizer_metrics_allowed_custom_labels:
|
||||
labels[label] = ""
|
||||
if self.server_args.extra_metric_labels:
|
||||
labels.update(self.server_args.extra_metric_labels)
|
||||
if get_observability().extra_metric_labels:
|
||||
labels.update(get_observability().extra_metric_labels)
|
||||
tokenizer_collector_cls = resolve_collector_class(
|
||||
self.server_args,
|
||||
STAT_LOGGER_ROLE_TOKENIZER,
|
||||
@@ -582,18 +588,18 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
self.metrics_collector = tokenizer_collector_cls(
|
||||
server_args=self.server_args,
|
||||
labels=labels,
|
||||
bucket_time_to_first_token=self.server_args.bucket_time_to_first_token,
|
||||
bucket_e2e_request_latency=self.server_args.bucket_e2e_request_latency,
|
||||
bucket_inter_token_latency=self.server_args.bucket_inter_token_latency,
|
||||
bucket_time_to_first_token=get_observability().bucket_time_to_first_token,
|
||||
bucket_e2e_request_latency=get_observability().bucket_e2e_request_latency,
|
||||
bucket_inter_token_latency=get_observability().bucket_inter_token_latency,
|
||||
)
|
||||
|
||||
start_cpu_monitor_thread("tokenizer")
|
||||
|
||||
if self.server_args.gc_warning_threshold_secs > 0.0:
|
||||
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
|
||||
if get_observability().gc_warning_threshold_secs > 0.0:
|
||||
configure_gc_warning(get_observability().gc_warning_threshold_secs)
|
||||
self.soft_watchdog = Watchdog.create(
|
||||
debug_name="TokenizerManager",
|
||||
watchdog_timeout=self.server_args.soft_watchdog_timeout,
|
||||
watchdog_timeout=get_device().soft_watchdog_timeout,
|
||||
soft=True,
|
||||
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
|
||||
)
|
||||
@@ -1757,7 +1763,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
|
||||
# default the load format to the server_args
|
||||
if obj.load_format is None:
|
||||
obj.load_format = self.server_args.load_format
|
||||
obj.load_format = get_model().load_format
|
||||
logger.info("Start update_weights. Load format=%s", obj.load_format)
|
||||
|
||||
if obj.abort_all_requests:
|
||||
@@ -1783,7 +1789,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
|
||||
def _update_model_path_info(self, model_path: str, load_format: str):
|
||||
self.served_model_name = model_path
|
||||
self.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"tokenizer.update_weights", model_path=model_path, load_format=load_format
|
||||
)
|
||||
self.model_path = model_path
|
||||
@@ -1927,7 +1935,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
"id": rid,
|
||||
"finish_reason": recv_obj.finished_reasons[i],
|
||||
"prompt_tokens": recv_obj.prompt_tokens[i],
|
||||
"weight_version": self.server_args.weight_version,
|
||||
"weight_version": get_serving().weight_version,
|
||||
"num_retractions": recv_obj.retraction_counts[i],
|
||||
}
|
||||
|
||||
@@ -2801,7 +2809,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
|
||||
meta_info = {
|
||||
"id": recv_obj.rid,
|
||||
"finish_reason": finish_reason,
|
||||
"weight_version": self.server_args.weight_version,
|
||||
"weight_version": get_serving().weight_version,
|
||||
"e2e_latency": state.time_stats.get_e2e_latency(),
|
||||
}
|
||||
is_stream = getattr(state.obj, "stream", False)
|
||||
|
||||
@@ -597,7 +597,10 @@ class TokenizerManagerScoreMixin:
|
||||
f"Token ID {token_id} is out of vocabulary (vocab size: {vocab_size})"
|
||||
)
|
||||
|
||||
# Check if multi-item scoring is enabled
|
||||
# Check if multi-item scoring is enabled. enable_mis is a static startup
|
||||
# feature flag (never overridden post-publish), and score_request is also
|
||||
# exercised on a bare mixin without a published context, so read it off
|
||||
# server_args rather than the resolved-config bag.
|
||||
use_multi_item_scoring = self.server_args.enable_mis
|
||||
|
||||
input_ids = None
|
||||
|
||||
@@ -47,6 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||
from sglang.srt.runtime_context import get_exec, get_model, get_schedule, get_spec
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
|
||||
from sglang.srt.utils.hf_transformers_utils import (
|
||||
@@ -405,14 +406,14 @@ class TpModelWorker(BaseTpWorker):
|
||||
self.model_config = ModelConfig.from_server_args(
|
||||
self.server_args,
|
||||
model_path=(
|
||||
self.server_args.model_path
|
||||
get_model().model_path
|
||||
if not self.is_draft_worker
|
||||
else self.server_args.speculative_draft_model_path
|
||||
else get_spec().speculative_draft_model_path
|
||||
),
|
||||
model_revision=(
|
||||
self.server_args.revision
|
||||
get_model().revision
|
||||
if not self.is_draft_worker
|
||||
else self.server_args.speculative_draft_model_revision
|
||||
else get_spec().speculative_draft_model_revision
|
||||
),
|
||||
is_draft_model=self.is_draft_worker,
|
||||
context_length=self.context_length,
|
||||
@@ -423,7 +424,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
|
||||
self._model_runner = ModelRunner(
|
||||
model_config=self.model_config,
|
||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||
gpu_id=self.gpu_id,
|
||||
ps=self.ps,
|
||||
nccl_port=self.nccl_port,
|
||||
@@ -439,11 +440,11 @@ class TpModelWorker(BaseTpWorker):
|
||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||
|
||||
self.model_runner_list.append(self.model_runner)
|
||||
for i in range(1, self.server_args.speculative_num_steps):
|
||||
for i in range(1, get_spec().speculative_num_steps):
|
||||
self.model_runner_list.append(
|
||||
ModelRunner(
|
||||
model_config=self.model_config,
|
||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
||||
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||
gpu_id=self.gpu_id,
|
||||
ps=self.ps,
|
||||
nccl_port=self.nccl_port,
|
||||
@@ -459,7 +460,7 @@ class TpModelWorker(BaseTpWorker):
|
||||
def _init_dllm_algorithm(self):
|
||||
from sglang.srt.dllm.algorithm.base import DllmAlgorithm
|
||||
|
||||
if self.server_args.dllm_algorithm is not None:
|
||||
if get_exec().dllm.dllm_algorithm is not None:
|
||||
self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args)
|
||||
else:
|
||||
self.dllm_algorithm = None
|
||||
@@ -485,9 +486,9 @@ class TpModelWorker(BaseTpWorker):
|
||||
)
|
||||
return (
|
||||
self.model_runner.max_total_num_tokens,
|
||||
self.server_args.max_prefill_tokens,
|
||||
get_schedule().max_prefill_tokens,
|
||||
self.model_runner.max_running_requests,
|
||||
self.server_args.max_queued_requests,
|
||||
get_schedule().max_queued_requests,
|
||||
max_req_len,
|
||||
max_req_len - 5,
|
||||
self.random_seed,
|
||||
|
||||
@@ -26,7 +26,7 @@ from sglang.srt.mem_cache.common import (
|
||||
evict_from_tree_cache,
|
||||
)
|
||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_server_args
|
||||
from sglang.srt.utils import (
|
||||
is_cpu,
|
||||
is_cuda,
|
||||
@@ -65,7 +65,7 @@ def write_cache_indices(
|
||||
prefix_tensors: list[torch.Tensor],
|
||||
req_to_token_pool: ReqToTokenPool,
|
||||
):
|
||||
if support_triton(get_server_args().attention_backend):
|
||||
if support_triton(get_exec().kernel.attention_backend):
|
||||
prefix_pointers = torch.tensor(
|
||||
[t.data_ptr() for t in prefix_tensors],
|
||||
dtype=torch.uint64,
|
||||
@@ -106,7 +106,7 @@ def get_last_loc(
|
||||
req_pool_indices_tensor: torch.Tensor,
|
||||
prefix_lens_tensor: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
attn_backend = get_server_args().attention_backend
|
||||
attn_backend = get_exec().kernel.attention_backend
|
||||
uses_triton_dispatch = attn_backend not in ("ascend", "torch_native")
|
||||
|
||||
if _is_hip and uses_triton_dispatch:
|
||||
|
||||
@@ -16,7 +16,7 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
|
||||
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_server_args, get_serving
|
||||
from sglang.srt.utils.common import ceil_align
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -183,7 +183,7 @@ def _release_overallocated_kv_indices(
|
||||
|
||||
# strip_thinking_cache intentionally reports output tokens as overallocated
|
||||
# so they fall into the free path below (#22373).
|
||||
if spec_algo is None and not global_server_args.strip_thinking_cache:
|
||||
if spec_algo is None and not get_serving().strip_thinking_cache:
|
||||
assert (
|
||||
start_p == end_p
|
||||
), f"Unexpected overallocated KV cache, {req.kv_committed_len=}, {req.kv.kv_allocated_len=}"
|
||||
|
||||
@@ -21,7 +21,7 @@ from sglang.srt.environ import envs
|
||||
from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool
|
||||
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
|
||||
from sglang.srt.mem_cache.memory_pool import KVCache
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_server_args
|
||||
from sglang.srt.utils import ceil_div, is_hip
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -276,7 +276,7 @@ class DeepSeekV4IndexerPool(KVCache):
|
||||
end_layer,
|
||||
)
|
||||
self.index_head_dim = index_head_dim
|
||||
self.use_fp4_indexer = get_server_args().enable_deepseek_v4_fp4_indexer
|
||||
self.use_fp4_indexer = get_exec().kernel.enable_deepseek_v4_fp4_indexer
|
||||
|
||||
self._create_buffer()
|
||||
|
||||
|
||||
@@ -58,7 +58,15 @@ from sglang.srt.mem_cache.memory_pool import (
|
||||
)
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import get_model, get_parallel
|
||||
from sglang.srt.runtime_context import (
|
||||
get_disagg,
|
||||
get_exec,
|
||||
get_memory,
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.utils.common import (
|
||||
@@ -115,9 +123,7 @@ if TYPE_CHECKING:
|
||||
from sglang.srt.model_executor.model_runner_components.spec_aux_hidden_state import (
|
||||
SpecAuxHiddenStateConfig,
|
||||
)
|
||||
from sglang.srt.model_executor.pool_configurator import (
|
||||
MemoryPoolConfig,
|
||||
)
|
||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||
|
||||
|
||||
class KVCacheConfigResult(msgspec.Struct, frozen=True, kw_only=True):
|
||||
@@ -308,8 +314,8 @@ class KVCacheConfigurator:
|
||||
# from one byte buffer, then return. Gated to the target worker
|
||||
# (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4).
|
||||
if (
|
||||
self.server_args.enable_unified_memory
|
||||
and self.server_args.disaggregation_mode == "null"
|
||||
get_memory().enable_unified_memory
|
||||
and get_disagg().disaggregation_mode == "null"
|
||||
and req_to_token_pool is None
|
||||
):
|
||||
if self.mambaish_config is not None:
|
||||
@@ -358,13 +364,13 @@ class KVCacheConfigurator:
|
||||
# TARGET_VERIFY, so their pools skip the per-step intermediate
|
||||
# (SpeculativeState) buffers only the target pool consumes.
|
||||
req_to_token_pool = req_to_token_pool.clone_with_new_mamba(
|
||||
mamba_size=self.server_args.max_mamba_cache_size,
|
||||
mamba_size=get_schedule().max_mamba_cache_size,
|
||||
mamba_spec_state_size=sizes.max_running_requests,
|
||||
cache_params=self.mambaish_config.mamba2_cache_params,
|
||||
device=self.device,
|
||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||
draft_model_idx=self.draft_model_idx,
|
||||
speculative_eagle_topk=self.server_args.speculative_eagle_topk,
|
||||
speculative_eagle_topk=get_spec().speculative_eagle_topk,
|
||||
)
|
||||
|
||||
# Initialize token_to_kv_pool
|
||||
@@ -394,7 +400,7 @@ class KVCacheConfigurator:
|
||||
# unsupported pool families before allocation. Keep this guard here so
|
||||
# future pool-selection refactors fail at boot instead of on first use.
|
||||
if (
|
||||
self.server_args.prefill_only_disable_kv_cache
|
||||
get_schedule().prefill_only_disable_kv_cache
|
||||
and not self.is_draft_worker
|
||||
and not isinstance(token_to_kv_pool, NoOpMHATokenToKVPool)
|
||||
):
|
||||
@@ -432,8 +438,8 @@ class KVCacheConfigurator:
|
||||
assert self.page_size >= 1, f"page_size must be >= 1, got {self.page_size}"
|
||||
# Mirror the non-shared path's extra_max_context_len computation.
|
||||
extra_max_context_len = 4
|
||||
if self.server_args.speculative_num_draft_tokens is not None:
|
||||
extra_max_context_len += self.server_args.speculative_num_draft_tokens
|
||||
if get_spec().speculative_num_draft_tokens is not None:
|
||||
extra_max_context_len += get_spec().speculative_num_draft_tokens
|
||||
|
||||
mamba_layer_ids = [
|
||||
i
|
||||
@@ -462,14 +468,14 @@ class KVCacheConfigurator:
|
||||
model_context_len=self.model_config.context_len,
|
||||
extra_max_context_len=extra_max_context_len,
|
||||
max_total_num_tokens=max_total_num_tokens,
|
||||
max_mamba_cache_size=self.server_args.max_mamba_cache_size,
|
||||
max_mamba_cache_size=get_schedule().max_mamba_cache_size,
|
||||
max_num_reqs=max_num_reqs,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
||||
disable_overlap_schedule=self.server_args.disable_overlap_schedule,
|
||||
need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"),
|
||||
mamba_full_memory_ratio=self.server_args.mamba_full_memory_ratio,
|
||||
speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens,
|
||||
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
|
||||
mamba_full_memory_ratio=get_schedule().mamba_full_memory_ratio,
|
||||
# Overlap mode: the allocator's `free` drops a wait_stream(forward_stream)
|
||||
# barrier so eager compaction serializes after the in-flight forward's
|
||||
# v2p/KV reads. Near-no-op in normal mode.
|
||||
@@ -502,13 +508,13 @@ class KVCacheConfigurator:
|
||||
), "unified memory pool does not support MLA-SWA hybrid yet"
|
||||
# Mirror the non-shared path's extra_max_context_len computation.
|
||||
extra_max_context_len = 4
|
||||
if self.server_args.speculative_num_draft_tokens is not None:
|
||||
extra_max_context_len += self.server_args.speculative_num_draft_tokens
|
||||
if get_spec().speculative_num_draft_tokens is not None:
|
||||
extra_max_context_len += get_spec().speculative_num_draft_tokens
|
||||
req_to_token_pool = ReqToTokenPool(
|
||||
size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
)
|
||||
|
||||
head_num = self.model_config.get_num_kv_heads(get_parallel().attn_tp_size)
|
||||
@@ -558,8 +564,8 @@ class KVCacheConfigurator:
|
||||
full_attention_layer_ids=full_attention_layer_ids,
|
||||
full_max_total_num_tokens=full_max_total_num_tokens,
|
||||
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"),
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
|
||||
# Overlap mode: same wait_stream(forward_stream) rationale as
|
||||
# `_init_unified_mamba_pools`.
|
||||
forward_stream=self.forward_stream,
|
||||
@@ -579,7 +585,7 @@ class KVCacheConfigurator:
|
||||
is_dsv4_model: bool,
|
||||
current_platform,
|
||||
):
|
||||
if not self.server_args.prefill_only_disable_kv_cache or self.is_draft_worker:
|
||||
if not get_schedule().prefill_only_disable_kv_cache or self.is_draft_worker:
|
||||
return
|
||||
|
||||
unsupported_pool_family = None
|
||||
@@ -588,7 +594,7 @@ class KVCacheConfigurator:
|
||||
elif current_platform.is_out_of_tree() and not self.mambaish_config:
|
||||
unsupported_pool_family = "out-of-tree platform KV pool"
|
||||
elif (
|
||||
self.server_args.attention_backend == "ascend" and not self.mambaish_config
|
||||
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
|
||||
):
|
||||
unsupported_pool_family = "NPU/Ascend KV pool"
|
||||
elif self.use_mla_backend and is_dsa_model:
|
||||
@@ -614,9 +620,9 @@ class KVCacheConfigurator:
|
||||
def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool:
|
||||
extra_max_context_len = get_req_to_token_extra_context_len(self.server_args)
|
||||
|
||||
if self.server_args.disaggregation_mode == "decode":
|
||||
if get_disagg().disaggregation_mode == "decode":
|
||||
# Extra slots for pre-allocated requests
|
||||
pre_alloc_size = self.server_args.disaggregation_decode_extra_slots
|
||||
pre_alloc_size = get_disagg().disaggregation_decode_extra_slots
|
||||
if self.mambaish_config:
|
||||
req_to_token_pool = self._build_hybrid_mamba_decode_req_pool(
|
||||
max_num_reqs=max_num_reqs,
|
||||
@@ -648,15 +654,13 @@ class KVCacheConfigurator:
|
||||
extra_max_context_len: int,
|
||||
pre_alloc_size: int,
|
||||
) -> ReqToTokenPool:
|
||||
from sglang.srt.disaggregation.decode import (
|
||||
HybridMambaDecodeReqToTokenPool,
|
||||
)
|
||||
from sglang.srt.disaggregation.decode import HybridMambaDecodeReqToTokenPool
|
||||
|
||||
req_to_token_pool = HybridMambaDecodeReqToTokenPool(
|
||||
size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
cache_params=self.mambaish_config.mamba2_cache_params,
|
||||
mamba_layer_ids=(
|
||||
[
|
||||
@@ -666,11 +670,11 @@ class KVCacheConfigurator:
|
||||
]
|
||||
),
|
||||
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
|
||||
speculative_eagle_topk=self.server_args.speculative_eagle_topk,
|
||||
speculative_eagle_topk=get_spec().speculative_eagle_topk,
|
||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||
pre_alloc_size=pre_alloc_size,
|
||||
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
||||
mamba_size=self.server_args.max_mamba_cache_size,
|
||||
enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
|
||||
mamba_size=get_schedule().max_mamba_cache_size,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
)
|
||||
return req_to_token_pool
|
||||
@@ -688,7 +692,7 @@ class KVCacheConfigurator:
|
||||
size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
pre_alloc_size=pre_alloc_size,
|
||||
)
|
||||
return req_to_token_pool
|
||||
@@ -701,11 +705,11 @@ class KVCacheConfigurator:
|
||||
) -> ReqToTokenPool:
|
||||
req_to_token_pool = HybridReqToTokenPool(
|
||||
size=max_num_reqs,
|
||||
mamba_size=self.server_args.max_mamba_cache_size,
|
||||
mamba_size=get_schedule().max_mamba_cache_size,
|
||||
mamba_spec_state_size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
cache_params=self.mambaish_config.mamba2_cache_params,
|
||||
mamba_layer_ids=(
|
||||
[
|
||||
@@ -717,18 +721,18 @@ class KVCacheConfigurator:
|
||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||
enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(),
|
||||
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
|
||||
speculative_eagle_topk=self.server_args.speculative_eagle_topk,
|
||||
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
||||
speculative_eagle_topk=get_spec().speculative_eagle_topk,
|
||||
enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
enable_linear_replayssm=self.server_args.enable_linear_replayssm,
|
||||
linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len,
|
||||
mamba_envelope_layout=self.server_args.enable_page_major_kv_layout,
|
||||
enable_linear_replayssm=get_exec().mamba.enable_linear_replayssm,
|
||||
linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len,
|
||||
mamba_envelope_layout=get_memory().enable_page_major_kv_layout,
|
||||
# ReplaySSM spec-verify is GDN-only: activate the pool machinery
|
||||
# (rings + cursors + the intermediate_ssm gate) only for GDN-hybrid
|
||||
# models, so any other mamba-ish model (Mamba2/Nemotron, lightning,
|
||||
# ...) run with the flag set stays byte-identical to flag-off.
|
||||
enable_gdn_replayssm_spec=(
|
||||
self.server_args.enable_gdn_replayssm_spec
|
||||
get_exec().mamba.enable_gdn_replayssm_spec
|
||||
and self.hybrid_gdn_config is not None
|
||||
),
|
||||
)
|
||||
@@ -754,7 +758,7 @@ class KVCacheConfigurator:
|
||||
size=max_num_reqs,
|
||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
)
|
||||
return req_to_token_pool
|
||||
|
||||
@@ -770,7 +774,7 @@ class KVCacheConfigurator:
|
||||
# selected by swapping in the PageMajorMHATokenToKVPool subclass. The
|
||||
# default keeps upstream's per-layer layout. The Mamba state pool is routed
|
||||
# separately via `mamba_envelope_layout` on the req-to-token pool above.
|
||||
enable_page_major = self.server_args.enable_page_major_kv_layout
|
||||
enable_page_major = get_memory().enable_page_major_kv_layout
|
||||
mha_pool_class = (
|
||||
PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool
|
||||
)
|
||||
@@ -802,7 +806,7 @@ class KVCacheConfigurator:
|
||||
max_total_num_tokens=sizes.max_total_num_tokens,
|
||||
)
|
||||
elif (
|
||||
self.server_args.attention_backend == "ascend" and not self.mambaish_config
|
||||
get_exec().kernel.attention_backend == "ascend" and not self.mambaish_config
|
||||
):
|
||||
if self.is_hybrid_swa:
|
||||
token_to_kv_pool = self._build_ascend_swa_kv_pool(
|
||||
@@ -878,14 +882,12 @@ class KVCacheConfigurator:
|
||||
c128_state_dtype: Optional[torch.dtype],
|
||||
req_to_token_pool: ReqToTokenPool,
|
||||
) -> KVCache:
|
||||
swa_page_size = self.server_args.page_size
|
||||
swa_page_size = get_schedule().page_size
|
||||
if not _is_npu:
|
||||
assert swa_page_size == 256, "In paged swa mode, page_size must be 256."
|
||||
|
||||
if self.is_draft_worker:
|
||||
from sglang.srt.models.deepseek_v4_nextn import (
|
||||
COMPRESS_RATIO_NEXTN_LAYER,
|
||||
)
|
||||
from sglang.srt.models.deepseek_v4_nextn import COMPRESS_RATIO_NEXTN_LAYER
|
||||
|
||||
compression_ratios = [
|
||||
COMPRESS_RATIO_NEXTN_LAYER
|
||||
@@ -912,12 +914,12 @@ class KVCacheConfigurator:
|
||||
# sliding eviction in ``ScheduleBatch._evict_swa``.
|
||||
c4_state_pool_size = npu_state_pool_size(
|
||||
ratio=4,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
max_num_reqs=max_running_requests,
|
||||
)
|
||||
c128_state_pool_size = npu_state_pool_size(
|
||||
ratio=128,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
max_num_reqs=max_running_requests,
|
||||
)
|
||||
else:
|
||||
@@ -935,7 +937,7 @@ class KVCacheConfigurator:
|
||||
c128_size=c128_max_total_num_tokens,
|
||||
c4_state_pool_size=c4_state_pool_size,
|
||||
c128_state_pool_size=c128_state_pool_size,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
swa_page_size=swa_page_size,
|
||||
sliding_window=self.model_config.window_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
@@ -946,11 +948,11 @@ class KVCacheConfigurator:
|
||||
indexer_head_dim=self.model_config.index_head_dim,
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
compression_ratios=compression_ratios,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
enable_hisparse=self.server_args.enable_hisparse,
|
||||
enable_hisparse=get_memory().enable_hisparse,
|
||||
online_mtp_max_draft_tokens=(
|
||||
self.server_args.max_speculative_num_draft_tokens or 0
|
||||
),
|
||||
@@ -961,7 +963,7 @@ class KVCacheConfigurator:
|
||||
PoolCls = current_platform.get_dsa_kv_pool_cls()
|
||||
token_to_kv_pool = PoolCls(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||
@@ -972,7 +974,7 @@ class KVCacheConfigurator:
|
||||
kv_cache_dtype=self.kv_cache_dtype,
|
||||
server_args=self.server_args,
|
||||
),
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
|
||||
@@ -985,14 +987,14 @@ class KVCacheConfigurator:
|
||||
PoolCls = current_platform.get_mla_kv_pool_cls()
|
||||
token_to_kv_pool = PoolCls(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||
index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None),
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
@@ -1002,13 +1004,13 @@ class KVCacheConfigurator:
|
||||
PoolCls = current_platform.get_mha_kv_pool_cls()
|
||||
token_to_kv_pool = PoolCls(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
head_dim=self.model_config.head_dim,
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
@@ -1020,9 +1022,7 @@ class KVCacheConfigurator:
|
||||
full_max_total_num_tokens: Optional[int],
|
||||
swa_max_total_num_tokens: Optional[int],
|
||||
) -> KVCache:
|
||||
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
||||
NPUMHATokenToKVPool,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool
|
||||
|
||||
kwargs = {}
|
||||
if self.is_hybrid_swa_compress:
|
||||
@@ -1039,7 +1039,7 @@ class KVCacheConfigurator:
|
||||
token_to_kv_pool = SWAKVPool(
|
||||
size=full_max_total_num_tokens,
|
||||
size_swa=swa_max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
post_capture_active=self.post_capture_kv_active,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
@@ -1055,39 +1055,35 @@ class KVCacheConfigurator:
|
||||
def _build_ascend_mla_kv_pool(
|
||||
self, *, max_total_num_tokens: int, is_dsa_model: bool
|
||||
) -> KVCache:
|
||||
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
||||
NPUMLATokenToKVPool,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool
|
||||
|
||||
token_to_kv_pool = NPUMLATokenToKVPool(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||
index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None),
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
return token_to_kv_pool
|
||||
|
||||
def _build_ascend_mha_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
||||
NPUMHATokenToKVPool,
|
||||
)
|
||||
from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMHATokenToKVPool
|
||||
|
||||
token_to_kv_pool = NPUMHATokenToKVPool(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
head_dim=self.model_config.head_dim,
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
@@ -1101,7 +1097,7 @@ class KVCacheConfigurator:
|
||||
dsa_cp_layer_shard_size,
|
||||
) = get_glm_dsa_cp_layer_shard_info(self)
|
||||
pool_kwargs = {}
|
||||
if self.server_args.enable_hisparse:
|
||||
if get_memory().enable_hisparse:
|
||||
PoolCls = HiSparseDSATokenToKVPool
|
||||
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||
|
||||
@@ -1121,7 +1117,7 @@ class KVCacheConfigurator:
|
||||
PoolCls = DSATokenToKVPool
|
||||
token_to_kv_pool = PoolCls(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||
@@ -1132,7 +1128,7 @@ class KVCacheConfigurator:
|
||||
kv_cache_dtype=self.kv_cache_dtype,
|
||||
server_args=self.server_args,
|
||||
),
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
|
||||
@@ -1143,13 +1139,13 @@ class KVCacheConfigurator:
|
||||
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||
token_to_kv_pool = MLATokenToKVPoolFP4(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
@@ -1158,13 +1154,13 @@ class KVCacheConfigurator:
|
||||
def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||
token_to_kv_pool = MLATokenToKVPool(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
@@ -1221,7 +1217,7 @@ class KVCacheConfigurator:
|
||||
token_to_kv_pool = SWAKVPool(
|
||||
size=full_max_total_num_tokens,
|
||||
size_swa=size_swa,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
post_capture_active=self.post_capture_kv_active,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
@@ -1229,7 +1225,7 @@ class KVCacheConfigurator:
|
||||
swa_attention_layer_ids=swa_attention_layer_ids,
|
||||
full_attention_layer_ids=full_attention_layer_ids,
|
||||
device=self.device,
|
||||
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
|
||||
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
|
||||
token_to_kv_pool_class=swa_pool_class,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -1244,7 +1240,7 @@ class KVCacheConfigurator:
|
||||
)
|
||||
token_to_kv_pool = MiniMaxSparseKVPool(
|
||||
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),
|
||||
@@ -1254,7 +1250,7 @@ class KVCacheConfigurator:
|
||||
sparse_layer_ids=sparse_layer_ids,
|
||||
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
)
|
||||
@@ -1293,7 +1289,7 @@ class KVCacheConfigurator:
|
||||
else mha_pool_class
|
||||
)
|
||||
token_to_kv_pool = HybridLinearKVPool(
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
size=max_total_num_tokens,
|
||||
dtype=self.kv_cache_dtype,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
@@ -1302,8 +1298,8 @@ class KVCacheConfigurator:
|
||||
full_attention_layer_ids=full_attention_layer_ids,
|
||||
device=self.device,
|
||||
mamba_pool=req_to_token_pool.mamba_pool,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
|
||||
use_mla=self.use_mla_backend,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
full_kv_pool_class=full_pool_class,
|
||||
@@ -1316,18 +1312,18 @@ class KVCacheConfigurator:
|
||||
def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||
token_to_kv_pool = MHATokenToKVPoolFP4(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
head_dim=self.model_config.head_dim,
|
||||
v_head_dim=self.model_config.v_head_dim,
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
enable_alt_stream=not self.server_args.enable_pdmux,
|
||||
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
|
||||
enable_alt_stream=not get_disagg().enable_pdmux,
|
||||
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
|
||||
)
|
||||
return token_to_kv_pool
|
||||
|
||||
@@ -1339,7 +1335,7 @@ class KVCacheConfigurator:
|
||||
else:
|
||||
pool_cls = (
|
||||
NoOpMHATokenToKVPool
|
||||
if self.server_args.prefill_only_disable_kv_cache
|
||||
if get_schedule().prefill_only_disable_kv_cache
|
||||
else mha_pool_class
|
||||
)
|
||||
pool_kwargs = {}
|
||||
@@ -1349,18 +1345,18 @@ class KVCacheConfigurator:
|
||||
pool_kwargs["post_capture_active"] = self.post_capture_kv_active
|
||||
token_to_kv_pool = pool_cls(
|
||||
max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||
head_dim=self.model_config.head_dim,
|
||||
v_head_dim=self.model_config.v_head_dim,
|
||||
layer_num=self.layer_info.num_effective_layers,
|
||||
device=self.device,
|
||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
||||
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||
start_layer=self.layer_info.start_layer,
|
||||
end_layer=self.layer_info.end_layer,
|
||||
enable_alt_stream=not self.server_args.enable_pdmux,
|
||||
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
|
||||
enable_alt_stream=not get_disagg().enable_pdmux,
|
||||
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
|
||||
**pool_kwargs,
|
||||
)
|
||||
return token_to_kv_pool
|
||||
@@ -1375,20 +1371,20 @@ class KVCacheConfigurator:
|
||||
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator],
|
||||
) -> BaseTokenToKVPoolAllocator:
|
||||
# Initialize token_to_kv_pool_allocator
|
||||
need_sort = self.server_args.disaggregation_mode in ("decode", "prefill")
|
||||
need_sort = get_disagg().disaggregation_mode in ("decode", "prefill")
|
||||
if token_to_kv_pool_allocator is None:
|
||||
if current_platform.is_out_of_tree():
|
||||
AllocatorCls = current_platform.get_paged_allocator_cls()
|
||||
token_to_kv_pool_allocator = AllocatorCls(
|
||||
sizes.max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device=self.device,
|
||||
kvcache=token_to_kv_pool,
|
||||
need_sort=need_sort,
|
||||
)
|
||||
elif _is_npu and (
|
||||
self.server_args.attention_backend == "ascend"
|
||||
get_exec().kernel.attention_backend == "ascend"
|
||||
or is_dsv4_model
|
||||
or self.hybrid_gdn_config is not None
|
||||
):
|
||||
@@ -1406,7 +1402,7 @@ class KVCacheConfigurator:
|
||||
token_to_kv_pool_allocator = swa_allocator_cls(
|
||||
sizes.full_max_total_num_tokens,
|
||||
sizes.swa_max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device=self.device,
|
||||
kvcache=token_to_kv_pool,
|
||||
@@ -1419,7 +1415,7 @@ class KVCacheConfigurator:
|
||||
|
||||
token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
|
||||
sizes.max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device=self.device,
|
||||
kvcache=token_to_kv_pool,
|
||||
@@ -1429,7 +1425,7 @@ class KVCacheConfigurator:
|
||||
if self.is_hybrid_swa and sizes.full_max_total_num_tokens == 0:
|
||||
token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator(
|
||||
sizes.swa_max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device=self.device,
|
||||
kvcache=token_to_kv_pool,
|
||||
@@ -1439,22 +1435,20 @@ class KVCacheConfigurator:
|
||||
token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
||||
sizes.full_max_total_num_tokens,
|
||||
sizes.swa_max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device=self.device,
|
||||
kvcache=token_to_kv_pool,
|
||||
need_sort=need_sort,
|
||||
)
|
||||
else:
|
||||
if self.server_args.enable_hisparse:
|
||||
from sglang.srt.mem_cache.sparsity import (
|
||||
parse_hisparse_config,
|
||||
)
|
||||
if get_memory().enable_hisparse:
|
||||
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||
|
||||
hisparse_cfg = parse_hisparse_config(self.server_args)
|
||||
token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator(
|
||||
sizes.max_total_num_tokens,
|
||||
page_size=self.server_args.page_size,
|
||||
page_size=get_schedule().page_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device=self.device,
|
||||
kvcache=token_to_kv_pool,
|
||||
@@ -1462,8 +1456,7 @@ class KVCacheConfigurator:
|
||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
||||
)
|
||||
elif (
|
||||
self.server_args.page_size == 1
|
||||
and self.server_args.dcp_size == 1
|
||||
get_schedule().page_size == 1 and self.server_args.dcp_size == 1
|
||||
):
|
||||
token_to_kv_pool_allocator = TokenToKVPoolAllocator(
|
||||
sizes.max_total_num_tokens,
|
||||
@@ -1475,7 +1468,7 @@ class KVCacheConfigurator:
|
||||
else:
|
||||
token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
|
||||
sizes.max_total_num_tokens * self.server_args.dcp_size,
|
||||
page_size=self.server_args.page_size
|
||||
page_size=get_schedule().page_size
|
||||
* self.server_args.dcp_size,
|
||||
dtype=self.kv_cache_dtype,
|
||||
device=self.device,
|
||||
@@ -1483,7 +1476,7 @@ class KVCacheConfigurator:
|
||||
need_sort=need_sort,
|
||||
)
|
||||
|
||||
if self.server_args.enable_hisparse and is_dsv4_model:
|
||||
if get_memory().enable_hisparse and is_dsv4_model:
|
||||
assert self.is_hybrid_swa, "DeepSeek V4 HiSparse requires SWA mode."
|
||||
token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator(
|
||||
token_to_kv_pool_allocator
|
||||
@@ -1535,7 +1528,7 @@ class KVCacheConfigurator:
|
||||
cpu_group=get_world_group().cpu_group,
|
||||
)
|
||||
|
||||
slack_gb = pre_model_load_memory * (1 - self.server_args.mem_fraction_static)
|
||||
slack_gb = pre_model_load_memory * (1 - get_schedule().mem_fraction_static)
|
||||
if self.mambaish_config is not None and self.post_capture_kv_active:
|
||||
# Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack.
|
||||
slack_gb = max(
|
||||
@@ -1559,7 +1552,7 @@ class KVCacheConfigurator:
|
||||
)
|
||||
raise ValueError(
|
||||
f"Loaded weights leave no GPU memory for the KV cache under "
|
||||
f"--mem-fraction-static={self.server_args.mem_fraction_static}. "
|
||||
f"--mem-fraction-static={get_schedule().mem_fraction_static}. "
|
||||
f"Raise --mem-fraction-static above "
|
||||
f"{suggested_mem_fraction_static:.3f} "
|
||||
f"(minimum viable = 1 - available/pre = "
|
||||
@@ -1570,14 +1563,14 @@ class KVCacheConfigurator:
|
||||
return int(rest_memory * (1 << 30)) # return in bytes
|
||||
|
||||
def _calculate_mamba_ratio(self) -> int:
|
||||
if self.server_args.disable_radix_cache:
|
||||
if get_memory().disable_radix_cache:
|
||||
return 1
|
||||
|
||||
additional_ratio = 0
|
||||
if self.server_args.enable_mamba_extra_buffer():
|
||||
# ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise.
|
||||
# Lazy mode saves 1 slot (2 → 1) for overlap; non-overlap already uses 1.
|
||||
if not self.server_args.disable_overlap_schedule:
|
||||
if not get_schedule().disable_overlap_schedule:
|
||||
if self.server_args.enable_mamba_extra_buffer_lazy():
|
||||
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY
|
||||
else:
|
||||
@@ -1596,7 +1589,7 @@ class KVCacheConfigurator:
|
||||
Page alignment is handled by the configurator, not here.
|
||||
If constraints change the value, the configurator re-runs and re-aligns.
|
||||
"""
|
||||
user_limit = self.server_args.max_total_tokens
|
||||
user_limit = get_schedule().max_total_tokens
|
||||
|
||||
# Apply user-specified upper bound
|
||||
if user_limit is not None:
|
||||
@@ -1626,7 +1619,7 @@ class KVCacheConfigurator:
|
||||
estimated = int(token_capacity / self.model_config.context_len * 512)
|
||||
estimated = max(min(estimated, 4096), 2048)
|
||||
|
||||
max_num_reqs = self.server_args.max_running_requests
|
||||
max_num_reqs = get_schedule().max_running_requests
|
||||
if max_num_reqs is not None:
|
||||
requested_per_worker = max_num_reqs // self.ps.attn_dp_size
|
||||
max_num_reqs = min(requested_per_worker, token_capacity // 2)
|
||||
@@ -1637,13 +1630,13 @@ class KVCacheConfigurator:
|
||||
if self.mambaish_config is not None:
|
||||
ratio = self._calculate_mamba_ratio()
|
||||
max_num_reqs = min(
|
||||
max_num_reqs, self.server_args.max_mamba_cache_size // ratio
|
||||
max_num_reqs, get_schedule().max_mamba_cache_size // ratio
|
||||
)
|
||||
|
||||
if max_num_reqs <= 0:
|
||||
raise RuntimeError(
|
||||
f"Hybrid (mamba/linear-attention) state cache is too small to serve "
|
||||
f"any requests. max_mamba_cache_size={self.server_args.max_mamba_cache_size}, "
|
||||
f"any requests. max_mamba_cache_size={get_schedule().max_mamba_cache_size}, "
|
||||
f"mamba_ratio={ratio}, resulting max_num_reqs={max_num_reqs}. "
|
||||
f"Try: (1) reduce --max-running-requests, "
|
||||
f"(2) increase --mem-fraction-static, or "
|
||||
@@ -1673,7 +1666,7 @@ class KVCacheConfigurator:
|
||||
)
|
||||
configurator = create_memory_pool_configurator(self)
|
||||
config = configurator.finalize_with_max_running_requests(config)
|
||||
config.mem_fraction_static = self.server_args.mem_fraction_static
|
||||
config.mem_fraction_static = get_schedule().mem_fraction_static
|
||||
return config
|
||||
|
||||
def config_from_budget(
|
||||
@@ -1689,18 +1682,20 @@ class KVCacheConfigurator:
|
||||
|
||||
configurator = create_memory_pool_configurator(self)
|
||||
config = configurator.calculate_pool_sizes(
|
||||
budget_bytes, self.server_args.page_size
|
||||
budget_bytes, get_schedule().page_size
|
||||
)
|
||||
max_tokens = self._apply_token_constraints(config.max_total_num_tokens)
|
||||
if cap_tokens is not None:
|
||||
max_tokens = min(max_tokens, cap_tokens)
|
||||
if max_tokens != config.max_total_num_tokens:
|
||||
config = configurator.calculate_pool_sizes_from_max_tokens(
|
||||
max_tokens, self.server_args.page_size
|
||||
max_tokens, get_schedule().page_size
|
||||
)
|
||||
return config
|
||||
|
||||
def _handle_max_mamba_cache(self, total_rest_memory):
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
config = self.mambaish_config
|
||||
server_args = self.server_args
|
||||
assert config is not None
|
||||
@@ -1710,11 +1705,11 @@ class KVCacheConfigurator:
|
||||
assert server_args.speculative_num_draft_tokens is not None
|
||||
assert server_args.max_running_requests is not None
|
||||
|
||||
if server_args.max_mamba_cache_size is not None:
|
||||
if get_schedule().max_mamba_cache_size is not None:
|
||||
# Use explicitly set max_mamba_cache_size
|
||||
server_args.override(
|
||||
get_context().override(
|
||||
"mamba_pool.per_dp_shard",
|
||||
max_mamba_cache_size=server_args.max_mamba_cache_size
|
||||
max_mamba_cache_size=get_schedule().max_mamba_cache_size
|
||||
// self.ps.attn_dp_size,
|
||||
)
|
||||
# Reserve intermediate memory based on capped max_num_reqs
|
||||
@@ -1722,7 +1717,7 @@ class KVCacheConfigurator:
|
||||
ratio = self._calculate_mamba_ratio()
|
||||
capped_reqs = min(
|
||||
server_args.max_running_requests // self.ps.attn_dp_size,
|
||||
server_args.max_mamba_cache_size // ratio,
|
||||
get_schedule().max_mamba_cache_size // ratio,
|
||||
)
|
||||
intermediate_size = (
|
||||
config.mamba2_cache_params.mamba_cache_per_req
|
||||
@@ -1735,7 +1730,7 @@ class KVCacheConfigurator:
|
||||
and server_args.max_running_requests is not None
|
||||
):
|
||||
# Use explicitly set max_running_requests when radix cache is disabled
|
||||
server_args.override(
|
||||
get_context().override(
|
||||
"mamba_pool.from_max_running_requests",
|
||||
max_mamba_cache_size=server_args.max_running_requests
|
||||
// self.ps.attn_dp_size,
|
||||
@@ -1744,7 +1739,7 @@ class KVCacheConfigurator:
|
||||
if has_spec_dec:
|
||||
intermediate_size = (
|
||||
config.mamba2_cache_params.mamba_cache_per_req
|
||||
* server_args.max_mamba_cache_size
|
||||
* get_schedule().max_mamba_cache_size
|
||||
* server_args.speculative_num_draft_tokens
|
||||
)
|
||||
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
||||
@@ -1769,7 +1764,7 @@ class KVCacheConfigurator:
|
||||
ratio = self._calculate_mamba_ratio()
|
||||
D = server_args.speculative_num_draft_tokens
|
||||
# Joint solve: main_state + intermediate = mamba_budget
|
||||
server_args.override(
|
||||
get_context().override(
|
||||
"mamba_pool.memory_budget_spec",
|
||||
max_mamba_cache_size=int(
|
||||
mamba_budget_bytes // (per_req * (1 + D / ratio))
|
||||
@@ -1779,12 +1774,12 @@ class KVCacheConfigurator:
|
||||
# so the return value only has main_state subtracted from total
|
||||
capped_reqs = min(
|
||||
server_args.max_running_requests // self.ps.attn_dp_size,
|
||||
server_args.max_mamba_cache_size // ratio,
|
||||
get_schedule().max_mamba_cache_size // ratio,
|
||||
)
|
||||
intermediate_size = per_req * capped_reqs * D
|
||||
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
||||
else:
|
||||
server_args.override(
|
||||
get_context().override(
|
||||
"mamba_pool.memory_budget",
|
||||
max_mamba_cache_size=int(mamba_budget_bytes // per_req),
|
||||
)
|
||||
@@ -1793,10 +1788,10 @@ class KVCacheConfigurator:
|
||||
# A non-positive value means GPU memory is insufficient for the requested
|
||||
# configuration. Fail fast with actionable advice instead of silently
|
||||
# producing garbled output at runtime.
|
||||
if server_args.max_mamba_cache_size <= 0:
|
||||
if get_schedule().max_mamba_cache_size <= 0:
|
||||
raise RuntimeError(
|
||||
f"Not enough GPU memory for hybrid (mamba/linear-attention) state cache. "
|
||||
f"Computed max_mamba_cache_size={server_args.max_mamba_cache_size} "
|
||||
f"Computed max_mamba_cache_size={get_schedule().max_mamba_cache_size} "
|
||||
f"(total_rest_memory={total_rest_memory:.2f} GB, "
|
||||
f"mamba_cache_per_req={config.mamba2_cache_params.mamba_cache_per_req / (1 << 20):.2f} MB). "
|
||||
f"Try: (1) reduce --max-running-requests, "
|
||||
@@ -1806,7 +1801,7 @@ class KVCacheConfigurator:
|
||||
)
|
||||
|
||||
mamba_state_memory = (
|
||||
server_args.max_mamba_cache_size
|
||||
get_schedule().max_mamba_cache_size
|
||||
* config.mamba2_cache_params.mamba_cache_per_req
|
||||
/ (1 << 30)
|
||||
)
|
||||
|
||||
@@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
||||
MatchResult,
|
||||
)
|
||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_memory, get_server_args
|
||||
|
||||
try:
|
||||
from lmcache.integration.sglang.multi_process_adapter import LMCacheMPConnector
|
||||
@@ -108,7 +108,7 @@ class LMCRadixCache(RadixCache):
|
||||
):
|
||||
super().__init__(params)
|
||||
|
||||
cli_lmc_cfg = get_server_args().lmcache_config_file or ""
|
||||
cli_lmc_cfg = get_memory().lmcache_config_file or ""
|
||||
|
||||
kvcache = self.token_to_kv_pool_allocator.get_kvcache()
|
||||
connector_kwargs = dict(
|
||||
|
||||
@@ -51,13 +51,8 @@ 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_parallel, get_server_args
|
||||
from sglang.srt.utils import (
|
||||
is_cuda,
|
||||
is_hip,
|
||||
is_npu,
|
||||
support_triton,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel
|
||||
from sglang.srt.utils import is_cuda, is_hip, is_npu, support_triton
|
||||
from sglang.srt.utils.common import ceil_align, is_pin_memory_available
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -941,7 +936,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# --enable-mis: every request must carry delimiter indices (the score
|
||||
# endpoint always produces MIS-structured requests; consumers index
|
||||
# without None-checking).
|
||||
if get_server_args().enable_mis and any(
|
||||
if get_exec().features.enable_mis and any(
|
||||
r.multi_item_delimiter_indices is not None for r in batch.reqs
|
||||
):
|
||||
assert all(
|
||||
@@ -1110,7 +1105,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
||||
# batch_size * [3 * seq_len]
|
||||
batch_size = self.seq_lens_cpu.shape[0]
|
||||
mrope_positions_list = [[]] * batch_size
|
||||
rl_on_policy_target = get_server_args().rl_on_policy_target
|
||||
rl_on_policy_target = get_exec().deterministic.rl_on_policy_target
|
||||
for batch_idx in range(batch_size):
|
||||
mm_input = batch.multimodal_inputs[batch_idx]
|
||||
if self.forward_mode.is_decode():
|
||||
|
||||
@@ -26,11 +26,7 @@ import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from sglang.srt.configs.load_config import LoadConfig
|
||||
from sglang.srt.configs.model_config import (
|
||||
AttentionArch,
|
||||
ModelConfig,
|
||||
ModelImpl,
|
||||
)
|
||||
from sglang.srt.configs.model_config import AttentionArch, ModelConfig, ModelImpl
|
||||
from sglang.srt.configs.update_config import adjust_config_with_unaligned_cpu_tp
|
||||
from sglang.srt.debug_utils.dumper import dumper
|
||||
from sglang.srt.distributed import bootstrap
|
||||
@@ -74,9 +70,7 @@ from sglang.srt.kv_canary.runner.canary_manager import context_tuple
|
||||
from sglang.srt.kv_canary.token_oracle.install import install_token_oracle_from_env
|
||||
from sglang.srt.layers import deep_gemm_wrapper, model_parallel
|
||||
from sglang.srt.layers.attention.dsa.utils import is_dsa_enable_prefill_cp
|
||||
from sglang.srt.layers.cp.utils import (
|
||||
get_cp_strategy,
|
||||
)
|
||||
from sglang.srt.layers.cp.utils import get_cp_strategy
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||
from sglang.srt.layers.sampler import create_sampler
|
||||
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
|
||||
@@ -86,17 +80,10 @@ from sglang.srt.lora.lora_registry import LoRARef
|
||||
from sglang.srt.managers.schedule_batch import sanity_check_mm_pad_shift_value
|
||||
from sglang.srt.mem_cache import kv_cache_dtype
|
||||
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
|
||||
from sglang.srt.mem_cache.kv_cache_configurator import (
|
||||
KVCacheConfigurator,
|
||||
)
|
||||
from sglang.srt.mem_cache.kv_cache_configurator import KVCacheConfigurator
|
||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||
from sglang.srt.model_executor.cuda_graph_config import (
|
||||
cuda_graph_fully_disabled,
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import (
|
||||
ForwardBatch,
|
||||
PPProxyTensors,
|
||||
)
|
||||
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
from sglang.srt.model_executor.forward_context import (
|
||||
ForwardContext,
|
||||
forward_context,
|
||||
@@ -155,14 +142,15 @@ from sglang.srt.model_executor.model_runner_components.weight_updater import (
|
||||
WeightUpdater,
|
||||
)
|
||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
||||
from sglang.srt.model_executor.runner import (
|
||||
EagerRunner,
|
||||
get_batch_sizes_to_capture,
|
||||
)
|
||||
from sglang.srt.model_executor.runner import EagerRunner, get_batch_sizes_to_capture
|
||||
from sglang.srt.platforms import current_platform
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_exec,
|
||||
get_global_dwdp_manager,
|
||||
get_server_args,
|
||||
get_lora,
|
||||
get_model,
|
||||
get_schedule,
|
||||
set_global_dwdp_manager,
|
||||
)
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
@@ -319,7 +307,7 @@ class ModelRunner:
|
||||
self.init_threads_binding()
|
||||
|
||||
# Set float32 matmul precision
|
||||
if get_server_args().enable_tf32_matmul:
|
||||
if get_exec().features.enable_tf32_matmul:
|
||||
torch.set_float32_matmul_precision("high")
|
||||
|
||||
# Set device early so that TransferEngine init (e.g. Ascend NPU)
|
||||
@@ -396,12 +384,12 @@ class ModelRunner:
|
||||
|
||||
def _initialize_elastic_ep_joiner(self) -> None:
|
||||
if not (
|
||||
self.server_args.elastic_ep_backend is not None
|
||||
get_exec().moe.elastic_ep_backend is not None
|
||||
and self.server_args.is_ep_joiner
|
||||
):
|
||||
return
|
||||
|
||||
is_scale_join = self.server_args.ep_join_mode == "scale"
|
||||
is_scale_join = get_exec().moe.ep_join_mode == "scale"
|
||||
if is_scale_join:
|
||||
join_effective_ep_size = (
|
||||
self.server_args.ep_join_rank_offset + self.ps.tp_size
|
||||
@@ -484,7 +472,7 @@ class ModelRunner:
|
||||
device=self.device,
|
||||
gpu_id=self.gpu_id,
|
||||
model_config=self.model_config,
|
||||
custom_weight_loaders=self.server_args.custom_weight_loader,
|
||||
custom_weight_loaders=get_model().custom_weight_loader,
|
||||
get_model=lambda: self.model,
|
||||
update_model_fields=self.update_model_fields,
|
||||
recapture_cuda_graph=self.init_decode_cuda_graph,
|
||||
@@ -561,7 +549,7 @@ class ModelRunner:
|
||||
def init_mindspore_runner(self):
|
||||
# Init the mindspore runner
|
||||
# for now, there is only some communication initialization work
|
||||
if self.server_args.model_impl.lower() == ModelImpl.MINDSPORE and _is_npu:
|
||||
if get_model().model_impl.lower() == ModelImpl.MINDSPORE and _is_npu:
|
||||
from sglang.srt.model_executor.mindspore_runner import init_ms_distributed
|
||||
|
||||
init_ms_distributed(
|
||||
@@ -618,7 +606,7 @@ class ModelRunner:
|
||||
|
||||
def init_memory_saver_adapter(self):
|
||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||
enable=self.server_args.enable_memory_saver
|
||||
enable=get_exec().features.enable_memory_saver
|
||||
)
|
||||
|
||||
def maybe_init_remote_instance_transfer_engine(self):
|
||||
@@ -654,7 +642,7 @@ class ModelRunner:
|
||||
)
|
||||
|
||||
def maybe_init_lplb_solvers(self):
|
||||
if self.server_args.ep_dispatch_algorithm == "lp" and not self.is_draft_worker:
|
||||
if get_exec().moe.ep_dispatch_algorithm == "lp" and not self.is_draft_worker:
|
||||
init_lplb_solvers(model_config=self.model_config)
|
||||
|
||||
def maybe_init_eplb_manager(self):
|
||||
@@ -668,12 +656,12 @@ class ModelRunner:
|
||||
get_expert_backup_client=lambda: self.expert_backup_client,
|
||||
get_weight_updater=lambda: self.weight_updater,
|
||||
)
|
||||
if self.server_args.enable_eplb and (not self.is_draft_worker)
|
||||
if get_exec().moe.enable_eplb and (not self.is_draft_worker)
|
||||
else None
|
||||
)
|
||||
|
||||
def maybe_init_elastic_ep(self):
|
||||
if self.server_args.elastic_ep_backend:
|
||||
if get_exec().moe.elastic_ep_backend:
|
||||
ElasticEPStateManager.init(self.server_args)
|
||||
|
||||
def init_token_oracle(self):
|
||||
@@ -692,8 +680,8 @@ class ModelRunner:
|
||||
get_model=lambda: self.model,
|
||||
)
|
||||
if (
|
||||
self.server_args.enable_elastic_expert_backup
|
||||
and self.server_args.elastic_ep_backend is not None
|
||||
get_exec().moe.enable_elastic_expert_backup
|
||||
and get_exec().moe.elastic_ep_backend is not None
|
||||
)
|
||||
else None
|
||||
)
|
||||
@@ -702,17 +690,17 @@ class ModelRunner:
|
||||
# In layered loading, torchao may have been applied
|
||||
torchao_applied = getattr(self.model, "torchao_applied", False)
|
||||
if not torchao_applied:
|
||||
apply_torchao_config_to_model(self.model, get_server_args().torchao_config)
|
||||
apply_torchao_config_to_model(self.model, get_exec().graph.torchao_config)
|
||||
supports_torch_tp = getattr(self.model, "supports_torch_tp", False)
|
||||
if self.ps.tp_size > 1 and supports_torch_tp:
|
||||
self.apply_torch_tp()
|
||||
|
||||
def maybe_init_lora_manager(self):
|
||||
if self.server_args.enable_lora:
|
||||
if get_lora().enable_lora:
|
||||
self.init_lora_manager()
|
||||
|
||||
def maybe_enable_batch_invariant_mode(self):
|
||||
if self.server_args.enable_deterministic_inference:
|
||||
if get_exec().deterministic.enable_deterministic_inference:
|
||||
from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode
|
||||
|
||||
enable_batch_invariant_mode()
|
||||
@@ -973,7 +961,7 @@ class ModelRunner:
|
||||
get_offloader().post_init()
|
||||
|
||||
# Register model for layerwise NVTX profiling if enabled
|
||||
if self.server_args.enable_layerwise_nvtx_marker:
|
||||
if get_exec().comm.enable_layerwise_nvtx_marker:
|
||||
pyt_hooks = PytHooks()
|
||||
pyt_hooks.register_hooks(self.model, module_prefix="model")
|
||||
|
||||
@@ -1030,7 +1018,7 @@ class ModelRunner:
|
||||
)
|
||||
|
||||
dist_barrier_after_load(
|
||||
elastic_ep_backend=self.server_args.elastic_ep_backend,
|
||||
elastic_ep_backend=get_exec().moe.elastic_ep_backend,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
is_ep_scale_joiner=self.server_args.is_ep_scale_joiner,
|
||||
)
|
||||
@@ -1050,16 +1038,16 @@ class ModelRunner:
|
||||
self.lora_manager = LoRAManager(
|
||||
base_model=self.model,
|
||||
base_hf_config=self.model_config.hf_config,
|
||||
max_loras_per_batch=self.server_args.max_loras_per_batch,
|
||||
max_loras_per_batch=get_lora().max_loras_per_batch,
|
||||
load_config=self.load_config,
|
||||
dtype=self.dtype,
|
||||
server_args=self.server_args,
|
||||
lora_backend=self.server_args.lora_backend,
|
||||
lora_backend=get_lora().lora_backend,
|
||||
tp_size=self.ps.tp_size,
|
||||
tp_rank=self.ps.tp_rank,
|
||||
max_lora_rank=self.server_args.max_lora_rank,
|
||||
target_modules=self.server_args.lora_target_modules,
|
||||
lora_paths=self.server_args.lora_paths,
|
||||
max_lora_rank=get_lora().max_lora_rank,
|
||||
target_modules=get_lora().lora_target_modules,
|
||||
lora_paths=get_lora().lora_paths,
|
||||
)
|
||||
if not cuda_graph_fully_disabled():
|
||||
init_lora_cuda_graph_moe_buffers(
|
||||
@@ -1331,7 +1319,7 @@ class ModelRunner:
|
||||
)
|
||||
output.expert_distribution_metrics = recorder_outputs.get("metrics")
|
||||
|
||||
no_copy_to_cpu = not self.server_args.disable_overlap_schedule
|
||||
no_copy_to_cpu = not get_schedule().disable_overlap_schedule
|
||||
if (
|
||||
not self.is_draft_worker
|
||||
and (experts_capturer := get_global_experts_capturer()) is not None
|
||||
@@ -1361,7 +1349,7 @@ class ModelRunner:
|
||||
self.msprobe_debugger.stop()
|
||||
self.msprobe_debugger.step()
|
||||
|
||||
if self.server_args.elastic_ep_backend is not None:
|
||||
if get_exec().moe.elastic_ep_backend is not None:
|
||||
self.maybe_join_ep_ranks()
|
||||
|
||||
return output
|
||||
@@ -1765,7 +1753,7 @@ class ModelRunner:
|
||||
recovered = maybe_recover_ep_ranks(
|
||||
tp_group=self.tp_group,
|
||||
eplb_manager=self.eplb_manager,
|
||||
random_seed=self.server_args.random_seed,
|
||||
random_seed=get_device().random_seed,
|
||||
)
|
||||
if recovered:
|
||||
self.forward_pass_id = 0
|
||||
@@ -1774,7 +1762,7 @@ class ModelRunner:
|
||||
local_timeout = (
|
||||
state.pending_since is not None
|
||||
and time.monotonic() - state.pending_since
|
||||
> self.server_args.elastic_ep_scale_timeout
|
||||
> get_exec().moe.elastic_ep_scale_timeout
|
||||
)
|
||||
timeout = state.active_ranks.new_tensor(int(local_timeout))
|
||||
dist.all_reduce(timeout, op=dist.ReduceOp.MAX, group=dist.group.WORLD)
|
||||
@@ -1842,7 +1830,9 @@ class ModelRunner:
|
||||
load_config: LoadConfig,
|
||||
) -> None:
|
||||
self.model = new_model
|
||||
self.server_args.override(
|
||||
from sglang.srt.runtime_context import get_context
|
||||
|
||||
get_context().override(
|
||||
"model_runner.update_model_fields",
|
||||
model_path=model_path,
|
||||
load_format=load_format,
|
||||
|
||||
@@ -24,6 +24,12 @@ def maybe_disable_chunked_prefix_cache(
|
||||
# model's (often non-MLA) config must not flip the shared setting.
|
||||
if is_draft_worker:
|
||||
return
|
||||
|
||||
# This is a load-time gate that runs in ModelRunner.__init__ BEFORE the
|
||||
# runner publishes its config (and direct/benchmark construction never
|
||||
# publishes earlier), so read/write the supplied server_args. The runner's
|
||||
# subsequent publish snapshots this into the schedule bag for get_schedule()
|
||||
# readers.
|
||||
if (
|
||||
not use_mla_backend
|
||||
or server_args.attention_backend
|
||||
|
||||
+3
-2
@@ -11,6 +11,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
RemoteInstanceWeightLoaderBackend,
|
||||
register_memory_region,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_model
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||
|
||||
@@ -58,7 +59,7 @@ class RemoteInstanceWeightTransporter:
|
||||
# ModelExpress owns TransferEngine memory registration and metadata
|
||||
# publishing for backend=modelexpress. Re-registering here would
|
||||
# overlap the same weight buffers.
|
||||
and self.server_args.remote_instance_weight_loader_backend
|
||||
and get_model().remote_instance_weight_loader_backend
|
||||
!= RemoteInstanceWeightLoaderBackend.MODELEXPRESS
|
||||
and self.engine is not None
|
||||
and self.weight_info is None
|
||||
@@ -84,7 +85,7 @@ class RemoteInstanceWeightTransporter:
|
||||
else:
|
||||
bootstrap_host = "127.0.0.1"
|
||||
|
||||
bootstrap_port = self.server_args.engine_info_bootstrap_port
|
||||
bootstrap_port = get_model().engine_info_bootstrap_port
|
||||
bootstrap_na = NetworkAddress(bootstrap_host, bootstrap_port)
|
||||
url = f"{bootstrap_na.to_url()}/register_transfer_engine_info"
|
||||
|
||||
|
||||
@@ -44,7 +44,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
get_remote_instance_transfer_engine_info_per_rank,
|
||||
register_memory_region,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_server_args
|
||||
from sglang.srt.utils import get_available_gpu_memory
|
||||
|
||||
# Try to import accelerate (optional dependency)
|
||||
@@ -71,9 +71,7 @@ from sglang.srt.connector import (
|
||||
get_connector_type,
|
||||
)
|
||||
from sglang.srt.connector.utils import parse_model_name
|
||||
from sglang.srt.distributed import (
|
||||
model_parallel_is_initialized,
|
||||
)
|
||||
from sglang.srt.distributed import model_parallel_is_initialized
|
||||
from sglang.srt.layers.modelopt_utils import QUANT_CFG_CHOICES
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
||||
@@ -865,9 +863,8 @@ class LayeredModelLoader(DefaultModelLoader):
|
||||
device_config: DeviceConfig,
|
||||
) -> nn.Module:
|
||||
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
torchao_config = get_server_args().torchao_config
|
||||
torchao_config = get_exec().graph.torchao_config
|
||||
target_device = torch.device(device_config.device)
|
||||
quant_config = _get_quantization_config(model_config, self.load_config)
|
||||
|
||||
|
||||
@@ -41,9 +41,7 @@ from sglang.srt.layers.communicator import (
|
||||
LayerScatterModes,
|
||||
enable_moe_dense_fully_dp,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
@@ -78,6 +76,7 @@ from sglang.srt.models.utils import (
|
||||
enable_fused_set_kv_buffer,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
@@ -209,7 +208,7 @@ class BailingMoESparseMoeBlock(nn.Module):
|
||||
self.router_dtype = torch.bfloat16
|
||||
|
||||
# TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now
|
||||
assert get_server_args().ep_num_redundant_experts == 0
|
||||
assert get_exec().moe.ep_num_redundant_experts == 0
|
||||
# check group topk
|
||||
self.num_expert_group = getattr(config, "n_group", 0)
|
||||
self.topk_group = getattr(config, "topk_group", 0)
|
||||
@@ -223,9 +222,7 @@ class BailingMoESparseMoeBlock(nn.Module):
|
||||
self.num_expert_group = self.topk_group = None
|
||||
self.use_grouped_topk = False
|
||||
|
||||
self.num_experts = (
|
||||
config.num_experts + get_server_args().ep_num_redundant_experts
|
||||
)
|
||||
self.num_experts = config.num_experts + get_exec().moe.ep_num_redundant_experts
|
||||
|
||||
self.gate = BailingMoEGate(
|
||||
config=config,
|
||||
|
||||
@@ -12,17 +12,12 @@ from transformers import PretrainedConfig
|
||||
from sglang.kernels.ops.attention.fla.layernorm_gated import RMSNorm as RMSNormGated
|
||||
from sglang.kernels.ops.attention.fla.layernorm_gated import layernorm_fn
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import is_fp8_fnuz
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce
|
||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||
from sglang.srt.layers import deep_gemm_wrapper
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
@@ -59,6 +54,7 @@ from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
|
||||
from sglang.srt.models.utils import WeightsMapper
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
@@ -529,7 +525,7 @@ class BailingMoELinearAttention(nn.Module):
|
||||
base=self.rope_theta,
|
||||
rope_scaling=config.rope_scaling,
|
||||
is_neox_style=True,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
@@ -690,7 +686,7 @@ class BailingMoEAttention(nn.Module):
|
||||
max_position=self.max_position_embeddings,
|
||||
base=self.rope_theta,
|
||||
rope_scaling=config.rope_scaling,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
self.attn = RadixAttention(
|
||||
self.num_heads,
|
||||
|
||||
@@ -16,7 +16,7 @@ from sglang.srt.layers.radix_attention import AttentionType, RadixAttention
|
||||
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_model, get_parallel
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
BertConfig = None
|
||||
@@ -365,9 +365,7 @@ class BertModel(nn.Module):
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("encoder", prefix),
|
||||
)
|
||||
pooling_type = (
|
||||
PoolingType.CLS if get_server_args().is_embedding else PoolingType.LAST
|
||||
)
|
||||
pooling_type = PoolingType.CLS if get_model().is_embedding else PoolingType.LAST
|
||||
self.pooler = (
|
||||
BertPooler(config)
|
||||
if self.use_bert_pooler
|
||||
|
||||
@@ -11,7 +11,7 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods
|
||||
AttnForwardMethod,
|
||||
)
|
||||
from sglang.srt.models.deepseek_common.utils import _is_hip
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.runtime_context import get_exec
|
||||
from sglang.srt.utils import use_intel_amx_backend
|
||||
|
||||
MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"]
|
||||
@@ -118,7 +118,7 @@ def handle_attention_flashinfer(attn, forward_batch):
|
||||
|
||||
def handle_attention_fa3(attn, forward_batch):
|
||||
# when deterministic inference is enabled, use MLA
|
||||
if get_server_args().enable_deterministic_inference:
|
||||
if get_exec().deterministic.enable_deterministic_inference:
|
||||
return _dispatch_mla_subtype(attn, forward_batch)
|
||||
else:
|
||||
return _handle_attention_backend(attn, forward_batch, "fa3")
|
||||
@@ -187,7 +187,7 @@ def handle_attention_triton(attn, forward_batch):
|
||||
return AttnForwardMethod.MLA
|
||||
|
||||
# when deterministic inference is enabled, use MLA
|
||||
if get_server_args().enable_deterministic_inference:
|
||||
if get_exec().deterministic.enable_deterministic_inference:
|
||||
return _dispatch_mla_subtype(attn, forward_batch)
|
||||
|
||||
if (
|
||||
|
||||
@@ -30,7 +30,11 @@ from sglang.srt.models.deepseek_common.utils import (
|
||||
_use_aiter_bpreshuffle_gfx95,
|
||||
_use_aiter_gfx95,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_schedule,
|
||||
)
|
||||
from sglang.srt.utils import BumpAllocator, get_bool_env_var, next_power_of_2
|
||||
|
||||
_use_fp8_prefill_attn = (
|
||||
@@ -142,9 +146,7 @@ def _forward_dsa_indexer_for_mha(
|
||||
|
||||
class DeepseekMHAForwardMixin:
|
||||
def init_mha_forward(self: DeepseekV2AttentionMLA):
|
||||
self.disable_chunked_prefix_cache = (
|
||||
get_server_args().disable_chunked_prefix_cache
|
||||
)
|
||||
self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache
|
||||
|
||||
# TODO: Design a finer way to determine the threshold
|
||||
self.chunked_prefix_cache_threshold = (
|
||||
@@ -305,8 +307,8 @@ class DeepseekMHAForwardMixin:
|
||||
self.use_dsa
|
||||
and self.kv_cache_dtype == "fp8_e4m3"
|
||||
and (
|
||||
not get_server_args().dsa_decode_backend == "trtllm"
|
||||
or not get_server_args().dsa_prefill_backend == "trtllm"
|
||||
not get_exec().kernel.dsa_decode_backend == "trtllm"
|
||||
or not get_exec().kernel.dsa_prefill_backend == "trtllm"
|
||||
)
|
||||
):
|
||||
# FP8 path: dequantize DSA-specific FP8 format to BF16
|
||||
|
||||
@@ -65,10 +65,8 @@ from sglang.srt.models.deepseek_common.utils import (
|
||||
_use_aiter_bpreshuffle_gfx95,
|
||||
_use_aiter_gfx95,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.state_capturer.indexer_topk import (
|
||||
maybe_capture_indexer_topk,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.state_capturer.indexer_topk import maybe_capture_indexer_topk
|
||||
from sglang.srt.utils import BumpAllocator
|
||||
from sglang.srt.utils.custom_op import register_custom_op
|
||||
|
||||
@@ -153,7 +151,7 @@ def _should_defer_dsa_cp_kv_gather(
|
||||
class DeepseekMLAForwardMixin:
|
||||
def init_mla_forward(self: DeepseekV2AttentionMLA):
|
||||
self.flashinfer_mla_disable_ragged = (
|
||||
get_server_args().flashinfer_mla_disable_ragged
|
||||
get_exec().kernel.flashinfer_mla_disable_ragged
|
||||
)
|
||||
|
||||
def should_run_indexer(
|
||||
@@ -990,8 +988,8 @@ class DeepseekMLAForwardMixin:
|
||||
"""
|
||||
if self.current_attention_backend in ("dsa", "nsa"):
|
||||
return (
|
||||
get_server_args().dsa_decode_backend == "trtllm"
|
||||
or get_server_args().dsa_prefill_backend == "trtllm"
|
||||
get_exec().kernel.dsa_decode_backend == "trtllm"
|
||||
or get_exec().kernel.dsa_prefill_backend == "trtllm"
|
||||
) and get_attn_backend().kv_cache_dtype == torch.float8_e4m3fn
|
||||
|
||||
return (
|
||||
|
||||
@@ -59,7 +59,12 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.models.deepseek_common.utils import enable_nextn_moe_bf16_cast_to_fp8
|
||||
from sglang.srt.models.deepseek_v2 import DeepseekV2DecoderLayer, DeepseekV3ForCausalLM
|
||||
from sglang.srt.models.utils import WeightsMapper
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_cuda, is_npu
|
||||
|
||||
|
||||
@@ -148,7 +153,7 @@ class DeepseekModelNextN(nn.Module):
|
||||
|
||||
self.rot_weight = None
|
||||
if _is_npu:
|
||||
rot_weight_path = get_server_args().model_path + "/rot.safetensors"
|
||||
rot_weight_path = get_model().model_path + "/rot.safetensors"
|
||||
if os.path.isfile(rot_weight_path):
|
||||
self.rot_weight = load_file(rot_weight_path)
|
||||
self.rot_weight = self.rot_weight["rot.weight"].npu()
|
||||
@@ -161,8 +166,7 @@ class DeepseekModelNextN(nn.Module):
|
||||
|
||||
layer_name = "decoder"
|
||||
if _is_npu and (
|
||||
get_server_args().speculative_draft_model_path
|
||||
== get_server_args().model_path
|
||||
get_spec().speculative_draft_model_path == get_model().model_path
|
||||
):
|
||||
layer_name = "layers." + str(config.num_hidden_layers)
|
||||
|
||||
@@ -201,7 +205,7 @@ class DeepseekModelNextN(nn.Module):
|
||||
if (
|
||||
_is_npu
|
||||
and self.quant_config is None
|
||||
and get_server_args().quantization is not None
|
||||
and get_model().quantization is not None
|
||||
):
|
||||
# ascend mtp unquant
|
||||
exit_stack.enter_context(envs.SGLANG_DEEPEP_BF16_DISPATCH.override(True))
|
||||
|
||||
@@ -29,10 +29,7 @@ import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.jit_kernel.dsv4 import (
|
||||
silu_and_mul_clamp,
|
||||
silu_and_mul_contig_post_quant,
|
||||
)
|
||||
from sglang.jit_kernel.dsv4 import silu_and_mul_clamp, silu_and_mul_contig_post_quant
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
create_per_token_group_quant_fp8_output_scale,
|
||||
)
|
||||
@@ -78,9 +75,7 @@ from sglang.srt.layers.communicator_dsa_cp import (
|
||||
maybe_prefetch_next_full_attention_kv,
|
||||
)
|
||||
from sglang.srt.layers.cp.utils import is_cp_v2_active
|
||||
from sglang.srt.layers.dcp.planner import (
|
||||
prepare_decode_context_parallel_metadata,
|
||||
)
|
||||
from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
ColumnParallelLinear,
|
||||
@@ -115,9 +110,7 @@ from sglang.srt.layers.moe.utils import (
|
||||
)
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8Config
|
||||
from sglang.srt.layers.quantization.fp8_utils import (
|
||||
materialize_bpreshuffle_fp8_scale,
|
||||
)
|
||||
from sglang.srt.layers.quantization.fp8_utils import materialize_bpreshuffle_fp8_scale
|
||||
from sglang.srt.layers.quantization.mxfp4_flashinfer_trtllm_moe import (
|
||||
maybe_fuse_routed_scale_and_shared_add,
|
||||
)
|
||||
@@ -183,11 +176,14 @@ from sglang.srt.models.deepseek_common.utils import (
|
||||
is_wint4afp8_or_wint4a16_config,
|
||||
)
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_exec,
|
||||
get_flags,
|
||||
get_forward,
|
||||
get_model,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_spec,
|
||||
)
|
||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||
from sglang.srt.utils import (
|
||||
@@ -381,9 +377,7 @@ class DeepseekV2MLP(nn.Module):
|
||||
return down_output
|
||||
|
||||
if self.use_fused_clamp_act_mul and self.swiglu_limit is not None:
|
||||
from aiter.ops.triton.fusions.fused_clamp_act_mul import (
|
||||
fused_clamp_act_mul,
|
||||
)
|
||||
from aiter.ops.triton.fusions.fused_clamp_act_mul import fused_clamp_act_mul
|
||||
|
||||
if not self._fused_clamp_fp8_checked:
|
||||
from sglang.srt.layers.quantization.fp8 import Fp8LinearMethod
|
||||
@@ -494,7 +488,7 @@ class MoEGate(nn.Module):
|
||||
True, # is_vnni
|
||||
)
|
||||
|
||||
if get_server_args().enable_deterministic_inference:
|
||||
if get_exec().deterministic.enable_deterministic_inference:
|
||||
return F.linear(hidden_states, self.weight, None)
|
||||
|
||||
if (
|
||||
@@ -560,7 +554,7 @@ class DeepseekV2MoE(nn.Module):
|
||||
n_shared_experts = (
|
||||
0 if config.n_shared_experts is None else int(config.n_shared_experts)
|
||||
)
|
||||
_fusion_disabled = get_server_args().disable_shared_experts_fusion
|
||||
_fusion_disabled = get_exec().moe.disable_shared_experts_fusion
|
||||
|
||||
# num_fused_shared_experts drives weight remapping in deepseek_weight_loader:
|
||||
# mlp.shared_experts → mlp.experts.256 when > 0.
|
||||
@@ -630,8 +624,7 @@ class DeepseekV2MoE(nn.Module):
|
||||
fused_shared_experts_scaling_factor = 1.0 / float(self.moe_ep_size)
|
||||
|
||||
self.experts = get_moe_impl_class(quant_config)(
|
||||
num_experts=num_experts_for_moe
|
||||
+ get_server_args().ep_num_redundant_experts,
|
||||
num_experts=num_experts_for_moe + get_exec().moe.ep_num_redundant_experts,
|
||||
num_fused_shared_experts=self.num_fused_shared_experts,
|
||||
top_k=top_k_for_moe,
|
||||
hidden_size=config.hidden_size,
|
||||
@@ -804,7 +797,7 @@ class DeepseekV2MoE(nn.Module):
|
||||
# TODO: we will support tp < ep in the future
|
||||
self.ep_size = get_parallel().moe_ep_size
|
||||
self.num_experts = (
|
||||
config.n_routed_experts + get_server_args().ep_num_redundant_experts
|
||||
config.n_routed_experts + get_exec().moe.ep_num_redundant_experts
|
||||
)
|
||||
self.renormalize = config.norm_topk_prob
|
||||
self.topk_group = config.topk_group
|
||||
@@ -1718,7 +1711,7 @@ class DeepseekV2AttentionMLA(
|
||||
base=rope_theta,
|
||||
rope_scaling=rope_scaling,
|
||||
is_neox_style=is_neox_style,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
|
||||
if rope_scaling and rope_scaling.get("apply_yarn_scaling", True):
|
||||
@@ -2069,7 +2062,7 @@ class DeepseekV2DecoderLayer(nn.Module):
|
||||
rope_scaling = config.rope_scaling
|
||||
max_position_embeddings = config.max_position_embeddings
|
||||
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
|
||||
get_server_args().speculative_algorithm
|
||||
get_spec().speculative_algorithm
|
||||
)
|
||||
self.dsa_enable_prefill_cp = dsa_enable_prefill_cp
|
||||
self.mla_enable_prefill_cp = mla_enable_prefill_cp
|
||||
@@ -2765,7 +2758,7 @@ class DeepseekV2ForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
||||
self.num_fused_shared_experts = 0
|
||||
server_args = get_server_args()
|
||||
|
||||
if get_server_args().disable_shared_experts_fusion:
|
||||
if get_exec().moe.disable_shared_experts_fusion:
|
||||
return
|
||||
|
||||
disable_reason = None
|
||||
|
||||
@@ -29,18 +29,11 @@ from sglang.jit_kernel.dsv4 import (
|
||||
fused_rope_inplace,
|
||||
sglang_per_token_group_quant_fp8_dsv4_wo_a,
|
||||
)
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import (
|
||||
v4_rope_inplace_npu,
|
||||
)
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||
sglang_per_token_group_quant_fp8,
|
||||
)
|
||||
from sglang.kernels.ops.attention.deepseek_v4_rope import v4_rope_inplace_npu
|
||||
from sglang.kernels.ops.quantization.fp8_kernel import sglang_per_token_group_quant_fp8
|
||||
from sglang.srt.compilation.compilation_config import register_split_op
|
||||
from sglang.srt.configs.deepseek_v4 import DeepSeekV4Config
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
get_tp_group,
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group, get_tp_group
|
||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||
use_symmetric_memory,
|
||||
)
|
||||
@@ -136,7 +129,13 @@ from sglang.srt.models.deepseek_v2 import (
|
||||
_is_npu,
|
||||
_is_xpu,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_forward, get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import (
|
||||
get_device,
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
)
|
||||
|
||||
if not _is_hip:
|
||||
from sglang.srt.layers.utils.cp_utils import (
|
||||
@@ -311,9 +310,7 @@ def _freqs_cis_to_cos_sin(
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend import (
|
||||
DeepseekV4AttnBackend,
|
||||
)
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend import DeepseekV4AttnBackend
|
||||
from sglang.srt.layers.attention.deepseek_v4_backend_hip_radix import (
|
||||
DeepseekV4HipRadixBackend,
|
||||
)
|
||||
@@ -575,7 +572,7 @@ class MQALayer(MqaAttentionBase):
|
||||
base=self.rope_base,
|
||||
rope_scaling=self.rope_scaling,
|
||||
is_neox_style=False,
|
||||
device=get_server_args().device,
|
||||
device=get_device().device,
|
||||
)
|
||||
|
||||
if _is_hip:
|
||||
@@ -2458,11 +2455,11 @@ class DeepseekV4ForCausalLM(nn.Module):
|
||||
|
||||
def determine_num_fused_shared_experts(self):
|
||||
self.num_fused_shared_experts = 0
|
||||
if get_server_args().disable_shared_experts_fusion:
|
||||
if get_exec().moe.disable_shared_experts_fusion:
|
||||
return
|
||||
|
||||
disable_reason = None
|
||||
if get_server_args().enforce_shared_experts_fusion:
|
||||
if get_exec().moe.enforce_shared_experts_fusion:
|
||||
if self.config.n_shared_experts != 1:
|
||||
raise ValueError(
|
||||
"DeepSeek V4 shared-experts fusion expects exactly one shared "
|
||||
|
||||
@@ -24,17 +24,12 @@ import torch
|
||||
from torch import nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
tensor_model_parallel_all_reduce,
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group, tensor_model_parallel_all_reduce
|
||||
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
|
||||
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
|
||||
from sglang.srt.eplb.expert_location_dispatch import ExpertLocationDispatchInfo
|
||||
from sglang.srt.layers.activation import SiluAndMul
|
||||
from sglang.srt.layers.dp_attention import (
|
||||
is_dp_attention_enabled,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
|
||||
from sglang.srt.layers.layernorm import RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
MergedColumnParallelLinear,
|
||||
@@ -62,7 +57,12 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args, get_stream
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
get_stream,
|
||||
)
|
||||
from sglang.srt.utils import LazyValue, add_prefix, is_cuda, make_layers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -165,7 +165,7 @@ class ExaoneMoESparseMoEBlock(nn.Module):
|
||||
)
|
||||
|
||||
self.experts = get_moe_impl_class(quant_config)(
|
||||
num_experts=config.num_experts + get_server_args().ep_num_redundant_experts,
|
||||
num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts,
|
||||
top_k=config.num_experts_per_tok,
|
||||
hidden_size=config.hidden_size,
|
||||
intermediate_size=config.moe_intermediate_size,
|
||||
@@ -206,7 +206,7 @@ class ExaoneMoESparseMoEBlock(nn.Module):
|
||||
if get_moe_a2a_backend().is_deepep():
|
||||
self.ep_size = get_parallel().moe_ep_size
|
||||
self.num_experts = (
|
||||
config.num_experts + get_server_args().ep_num_redundant_experts
|
||||
config.num_experts + get_exec().moe.ep_num_redundant_experts
|
||||
)
|
||||
self.top_k = config.num_experts_per_tok
|
||||
|
||||
|
||||
@@ -18,11 +18,7 @@ from typing import Iterable, List, Optional, Set, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import (
|
||||
Gemma4TextConfig,
|
||||
PretrainedConfig,
|
||||
PreTrainedModel,
|
||||
)
|
||||
from transformers import Gemma4TextConfig, PretrainedConfig, PreTrainedModel
|
||||
|
||||
from sglang.kernels.ops.layernorm.gemma4_fused_ops import (
|
||||
gemma4_fused_routing,
|
||||
@@ -31,9 +27,7 @@ from sglang.kernels.ops.layernorm.gemma4_fused_ops import (
|
||||
gemma_rmsnorm_residual_scalar,
|
||||
gemma_routing_post_topk,
|
||||
)
|
||||
from sglang.srt.distributed import (
|
||||
get_pp_group,
|
||||
)
|
||||
from sglang.srt.distributed import get_pp_group
|
||||
from sglang.srt.layers.layernorm import Gemma4RMSNorm, RMSNorm
|
||||
from sglang.srt.layers.linear import (
|
||||
QKVParallelLinear,
|
||||
@@ -55,10 +49,8 @@ from sglang.srt.model_loader.weight_utils import (
|
||||
maybe_remap_kv_scale_name,
|
||||
)
|
||||
from sglang.srt.models.gemma3_causal import Gemma3MLP, Gemma3TextScaledWordEmbedding
|
||||
from sglang.srt.models.utils import (
|
||||
create_fused_set_kv_buffer_arg,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.models.utils import create_fused_set_kv_buffer_arg
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
||||
from sglang.srt.utils import add_prefix, make_layers
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -254,7 +246,7 @@ class Gemma4MoE(nn.Module):
|
||||
experts_type = get_moe_impl_class(quant_config)
|
||||
|
||||
self.experts = experts_type(
|
||||
num_experts=config.num_experts + get_server_args().ep_num_redundant_experts,
|
||||
num_experts=config.num_experts + get_exec().moe.ep_num_redundant_experts,
|
||||
hidden_size=config.hidden_size,
|
||||
intermediate_size=config.moe_intermediate_size,
|
||||
layer_id=layer_id,
|
||||
|
||||
@@ -29,7 +29,7 @@ from sglang.srt.layers.clippable_linear import (
|
||||
)
|
||||
from sglang.srt.layers.layernorm import Gemma4RMSNorm
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.runtime_context import get_parallel
|
||||
from sglang.srt.runtime_context import get_mm, get_parallel
|
||||
from sglang.srt.utils import add_prefix, get_device_capability, is_cuda, is_hip
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -181,9 +181,8 @@ class Gemma4VisionAttention(nn.Module):
|
||||
@staticmethod
|
||||
def _select_backend() -> str:
|
||||
"""Mirror VisionAttention._determine_attention_backend for consistency."""
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
|
||||
override = get_server_args().mm_attention_backend
|
||||
override = get_mm().mm_attention_backend
|
||||
if override is not None:
|
||||
return override
|
||||
if is_cuda():
|
||||
|
||||
@@ -84,6 +84,7 @@ from sglang.srt.models.deepseek_nextn import DeepseekV3ForCausalLMNextN
|
||||
from sglang.srt.models.deepseek_v2 import DeepseekV2ForCausalLM
|
||||
from sglang.srt.models.utils import WeightsMapper, apply_qk_norm
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
@@ -406,7 +407,7 @@ class Glm4MoeSparseMoeBlock(nn.Module):
|
||||
self.n_shared_experts = config.n_shared_experts
|
||||
self.num_fused_shared_experts = (
|
||||
0
|
||||
if get_server_args().disable_shared_experts_fusion
|
||||
if get_exec().moe.disable_shared_experts_fusion
|
||||
else config.n_shared_experts
|
||||
)
|
||||
|
||||
@@ -526,7 +527,7 @@ class Glm4MoeSparseMoeBlock(nn.Module):
|
||||
# TODO: we will support tp < ep in the future
|
||||
self.ep_size = get_parallel().moe_ep_size
|
||||
self.num_experts = (
|
||||
config.n_routed_experts + get_server_args().ep_num_redundant_experts
|
||||
config.n_routed_experts + get_exec().moe.ep_num_redundant_experts
|
||||
)
|
||||
self.renormalize = config.norm_topk_prob
|
||||
self.topk_group = config.topk_group
|
||||
@@ -1178,7 +1179,7 @@ class Glm4MoeForCausalLM(nn.Module):
|
||||
self.capture_aux_hidden_states = False
|
||||
|
||||
def determine_num_fused_shared_experts(self):
|
||||
if get_server_args().disable_shared_experts_fusion:
|
||||
if get_exec().moe.disable_shared_experts_fusion:
|
||||
return
|
||||
|
||||
disable_reason = None
|
||||
|
||||
@@ -75,6 +75,7 @@ from sglang.srt.models.deepseek_common.deepseek_weight_loader import (
|
||||
from sglang.srt.models.deepseek_common.utils import _is_cuda, _use_aiter
|
||||
from sglang.srt.models.deepseek_v2 import DeepseekV2AttentionMLA
|
||||
from sglang.srt.runtime_context import (
|
||||
get_exec,
|
||||
get_forward,
|
||||
get_parallel,
|
||||
get_server_args,
|
||||
@@ -189,7 +190,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
|
||||
self.n_shared_experts = config.n_shared_experts
|
||||
self.num_fused_shared_experts = (
|
||||
0
|
||||
if get_server_args().disable_shared_experts_fusion
|
||||
if get_exec().moe.disable_shared_experts_fusion
|
||||
else config.n_shared_experts
|
||||
)
|
||||
self.config = config
|
||||
@@ -216,7 +217,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
|
||||
self.experts = get_moe_impl_class(quant_config)(
|
||||
num_experts=config.n_routed_experts
|
||||
+ self.num_fused_shared_experts
|
||||
+ get_server_args().ep_num_redundant_experts,
|
||||
+ get_exec().moe.ep_num_redundant_experts,
|
||||
num_fused_shared_experts=self.num_fused_shared_experts,
|
||||
top_k=config.num_experts_per_tok + self.num_fused_shared_experts,
|
||||
hidden_size=config.hidden_size,
|
||||
@@ -284,7 +285,7 @@ class Glm4MoeLiteSparseMoeBlock(nn.Module):
|
||||
# TODO: we will support tp < ep in the future
|
||||
self.ep_size = get_parallel().moe_ep_size
|
||||
self.num_experts = (
|
||||
config.n_routed_experts + get_server_args().ep_num_redundant_experts
|
||||
config.n_routed_experts + get_exec().moe.ep_num_redundant_experts
|
||||
)
|
||||
self.renormalize = config.norm_topk_prob
|
||||
self.topk_group = config.topk_group
|
||||
@@ -928,7 +929,7 @@ class Glm4MoeLiteForCausalLM(nn.Module, DeepseekV2WeightLoaderMixin):
|
||||
self, architecture: str = "Glm4MoeLiteForCausalLM"
|
||||
):
|
||||
self.num_fused_shared_experts = 0
|
||||
if get_server_args().disable_shared_experts_fusion:
|
||||
if get_exec().moe.disable_shared_experts_fusion:
|
||||
return
|
||||
|
||||
disable_reason = None
|
||||
|
||||
@@ -35,7 +35,7 @@ from sglang.srt.models.glm4_moe_lite import (
|
||||
Glm4MoeLiteDecoderLayer,
|
||||
Glm4MoeLiteForCausalLM,
|
||||
)
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
|
||||
from sglang.srt.utils import BumpAllocator, add_prefix, is_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -139,7 +139,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
|
||||
nn.Module.__init__(self)
|
||||
self.config = config
|
||||
self.tp_size = get_parallel().tp_size
|
||||
if is_npu() and get_server_args().speculative_draft_model_quantization is None:
|
||||
if is_npu() and get_spec().speculative_draft_model_quantization is None:
|
||||
quant_config = None
|
||||
self.quant_config = quant_config
|
||||
|
||||
@@ -156,7 +156,7 @@ class Glm4MoeLiteForCausalLMNextN(Glm4MoeLiteForCausalLM):
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
self.num_fused_shared_experts = (
|
||||
0 if get_server_args().disable_shared_experts_fusion else 1
|
||||
0 if get_exec().moe.disable_shared_experts_fusion else 1
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
@@ -32,7 +32,7 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
||||
)
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||
from sglang.srt.models.glm4_moe import Glm4MoeDecoderLayer, Glm4MoeForCausalLM
|
||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args, get_spec
|
||||
from sglang.srt.utils import add_prefix, is_npu
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -125,7 +125,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
|
||||
nn.Module.__init__(self)
|
||||
self.config = config
|
||||
self.tp_size = get_parallel().tp_size
|
||||
if is_npu() and get_server_args().speculative_draft_model_quantization is None:
|
||||
if is_npu() and get_spec().speculative_draft_model_quantization is None:
|
||||
quant_config = None
|
||||
self.quant_config = quant_config
|
||||
|
||||
@@ -142,7 +142,7 @@ class Glm4MoeForCausalLMNextN(Glm4MoeForCausalLM):
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
self.num_fused_shared_experts = (
|
||||
0 if get_server_args().disable_shared_experts_fusion else 1
|
||||
0 if get_exec().moe.disable_shared_experts_fusion else 1
|
||||
)
|
||||
|
||||
@torch.no_grad()
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user