config: read resolved config via namespace accessors (#33013)
This commit is contained in:
@@ -261,19 +261,14 @@ def mamba_extra_buffer_of(cfg: Any) -> bool:
|
|||||||
|
|
||||||
def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None:
|
def declare_load_time_override(source: str, declared: Dict[str, Any]) -> None:
|
||||||
"""Declare a load-time resolved field (model-file config overrides,
|
"""Declare a load-time resolved field (model-file config overrides,
|
||||||
weight-resolved dtypes) on the published ``server_args``: resolution has
|
weight-resolved dtypes): validated against the resolvable whitelist, then
|
||||||
already materialized, so the declaration writes through, joining the
|
written to the config bags via ``get_context().override``; ``server_args``
|
||||||
declaration stash for provenance and republish consistency."""
|
stays the pristine startup record."""
|
||||||
from sglang.srt.runtime_context import get_context
|
from sglang.srt.runtime_context import get_context
|
||||||
|
|
||||||
server_args = get_context().server_args
|
context = get_context()
|
||||||
validate_declarations(server_args, [(source, dict(declared))])
|
validate_declarations(context.server_args, [(source, dict(declared))])
|
||||||
override = getattr(server_args, "override", None)
|
context.override(source, **declared)
|
||||||
if override is not None:
|
|
||||||
override(source, **declared)
|
|
||||||
else:
|
|
||||||
# Config-shaped fixtures without the mutation entry point.
|
|
||||||
_apply_fields(server_args, declared)
|
|
||||||
|
|
||||||
|
|
||||||
def collect_model_override_declarations(
|
def collect_model_override_declarations(
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
compute_position,
|
compute_position,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
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_parallel, get_server_args
|
||||||
from sglang.srt.speculative.spec_info import SpecInput
|
from sglang.srt.speculative.spec_info import SpecInput
|
||||||
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
|
from sglang.srt.utils import BumpAllocator, empty_context, get_bool_env_var, is_hip
|
||||||
|
|
||||||
@@ -184,7 +184,7 @@ def _update_device_and_sum_field_from_cpu_field(
|
|||||||
cpu_value
|
cpu_value
|
||||||
if isinstance(cpu_value, torch.Tensor)
|
if isinstance(cpu_value, torch.Tensor)
|
||||||
else torch.tensor(cpu_value, dtype=old_device_value.dtype)
|
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)
|
setattr(batch, device_field, new_device_value)
|
||||||
|
|
||||||
if sum_field is not None:
|
if sum_field is not None:
|
||||||
@@ -336,7 +336,7 @@ def compute_split_indices_for_cuda_graph_replay(
|
|||||||
class TboCudaGraphRunnerPlugin:
|
class TboCudaGraphRunnerPlugin:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
self._tbo_children_num_token_non_padded = torch.zeros(
|
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):
|
def capture_one_batch_size(self, batch: ForwardBatch, num_tokens: int):
|
||||||
@@ -835,7 +835,7 @@ class TboForwardBatchPreparer:
|
|||||||
value_a = min(tbo_split_token_index, num_token_non_padded)
|
value_a = min(tbo_split_token_index, num_token_non_padded)
|
||||||
value_b = max(0, num_token_non_padded - tbo_split_token_index)
|
value_b = max(0, num_token_non_padded - tbo_split_token_index)
|
||||||
return torch.tensor([value_a, value_b], dtype=torch.int32).to(
|
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
|
@classmethod
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from transformers import CONFIG_MAPPING
|
|||||||
from transformers.configuration_utils import PretrainedConfig
|
from transformers.configuration_utils import PretrainedConfig
|
||||||
|
|
||||||
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
|
from sglang.srt.configs.mamba_utils import BaseLinearStateParams
|
||||||
|
from sglang.srt.runtime_context import get_exec
|
||||||
|
|
||||||
|
|
||||||
class InklingModelConfig(PretrainedConfig):
|
class InklingModelConfig(PretrainedConfig):
|
||||||
@@ -224,9 +225,8 @@ class InklingModelConfig(PretrainedConfig):
|
|||||||
self.swa_num_key_value_heads, self.swa_head_dim
|
self.swa_num_key_value_heads, self.swa_head_dim
|
||||||
)
|
)
|
||||||
stream_dim = self.hidden_size
|
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]
|
# Scattered sconv: the attn/mlp output sconvs run on the [T, H/P]
|
||||||
# hidden shard, so their conv-state caches shard with them.
|
# hidden shard, so their conv-state caches shard with them.
|
||||||
assert (
|
assert (
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ import torch.distributed as dist
|
|||||||
import zmq
|
import zmq
|
||||||
|
|
||||||
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
from sglang.srt.managers.io_struct import sock_recv, sock_send, wrap_as_pickle
|
||||||
|
from sglang.srt.runtime_context import get_serving
|
||||||
|
|
||||||
# -------------------------------------- config base ------------------------------------------
|
# -------------------------------------- config base ------------------------------------------
|
||||||
|
|
||||||
@@ -1798,7 +1799,7 @@ class _SGLangPlugin(_FrameworkPlugin):
|
|||||||
if args is None:
|
if args is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
return args.tokenizer_path
|
return get_serving().tokenizer_path
|
||||||
except Exception:
|
except Exception:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|||||||
@@ -41,6 +41,7 @@ class KVArgs:
|
|||||||
kv_data_lens: List[int]
|
kv_data_lens: List[int]
|
||||||
kv_item_lens: List[int]
|
kv_item_lens: List[int]
|
||||||
kv_layer_ids: List[int]
|
kv_layer_ids: List[int]
|
||||||
|
kv_cache_dtype_str: str
|
||||||
aux_data_ptrs: List[int]
|
aux_data_ptrs: List[int]
|
||||||
aux_data_lens: List[int]
|
aux_data_lens: List[int]
|
||||||
aux_item_lens: List[int]
|
aux_item_lens: List[int]
|
||||||
|
|||||||
@@ -36,7 +36,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
get_attention_dp_rank,
|
get_attention_dp_rank,
|
||||||
get_attention_dp_size,
|
get_attention_dp_size,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_model, get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_serving
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils.network import (
|
from sglang.srt.utils.network import (
|
||||||
NetworkAddress,
|
NetworkAddress,
|
||||||
@@ -148,6 +148,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
is_mla_backend: Optional[bool] = False,
|
is_mla_backend: Optional[bool] = False,
|
||||||
):
|
):
|
||||||
self.kv_args = args
|
self.kv_args = args
|
||||||
|
self.kv_cache_dtype_str = args.kv_cache_dtype_str
|
||||||
self.kv_item_lens_sum = sum(args.kv_item_lens)
|
self.kv_item_lens_sum = sum(args.kv_item_lens)
|
||||||
self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp)
|
self.state_item_lens_sum = sum(x for comp in args.state_item_lens for x in comp)
|
||||||
self.is_mla_backend = is_mla_backend
|
self.is_mla_backend = is_mla_backend
|
||||||
@@ -533,11 +534,11 @@ class CommonKVManager(BaseKVManager):
|
|||||||
|
|
||||||
if (
|
if (
|
||||||
info.kv_cache_dtype is not None
|
info.kv_cache_dtype is not None
|
||||||
and info.kv_cache_dtype != get_model().kv_cache_dtype
|
and info.kv_cache_dtype != self.kv_cache_dtype_str
|
||||||
):
|
):
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"KV cache dtype mismatch: prefill server has kv_cache_dtype={info.kv_cache_dtype}, "
|
f"KV cache dtype mismatch: prefill server has kv_cache_dtype={info.kv_cache_dtype}, "
|
||||||
f"but decode server has kv_cache_dtype={get_model().kv_cache_dtype}. "
|
f"but decode server has kv_cache_dtype={self.kv_cache_dtype_str}. "
|
||||||
f"Both servers must use the same --kv-cache-dtype value."
|
f"Both servers must use the same --kv-cache-dtype value."
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -701,7 +702,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
"rank_ip": self.local_ip,
|
"rank_ip": self.local_ip,
|
||||||
"rank_port": self.rank_port,
|
"rank_port": self.rank_port,
|
||||||
"page_size": self.kv_args.page_size,
|
"page_size": self.kv_args.page_size,
|
||||||
"kv_cache_dtype": get_model().kv_cache_dtype,
|
"kv_cache_dtype": self.kv_cache_dtype_str,
|
||||||
"load_balance_method": self.server_args.load_balance_method,
|
"load_balance_method": self.server_args.load_balance_method,
|
||||||
"enable_dsa_cache_layer_split": getattr(
|
"enable_dsa_cache_layer_split": getattr(
|
||||||
self.server_args, "enable_dsa_cache_layer_split", False
|
self.server_args, "enable_dsa_cache_layer_split", False
|
||||||
@@ -709,7 +710,7 @@ class CommonKVManager(BaseKVManager):
|
|||||||
# Self-register the HTTP API port so the decode can derive the PD
|
# Self-register the HTTP API port so the decode can derive the PD
|
||||||
# retract rebootstrap /generate URL from bootstrap info instead of a
|
# retract rebootstrap /generate URL from bootstrap info instead of a
|
||||||
# router-injected pd_rebootstrap_prefill_url.
|
# 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
|
max_retries, initial_delay, max_delay = 5, 1.0, 30.0
|
||||||
|
|||||||
@@ -87,7 +87,7 @@ from sglang.srt.observability.req_time_stats import (
|
|||||||
set_schedule_time_batch,
|
set_schedule_time_batch,
|
||||||
set_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, is_npu
|
from sglang.srt.utils import get_num_new_pages, is_npu
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||||
@@ -423,6 +423,9 @@ class DecodePreallocQueue(DecodeHiCachePreallocMixin):
|
|||||||
|
|
||||||
kv_args.pp_rank = self.pp_rank
|
kv_args.pp_rank = self.pp_rank
|
||||||
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
||||||
|
kv_args.kv_cache_dtype_str = (
|
||||||
|
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
|
||||||
|
)
|
||||||
transfer_kv_pool = (
|
transfer_kv_pool = (
|
||||||
self.scheduler.hisparse_coordinator.mem_pool_host
|
self.scheduler.hisparse_coordinator.mem_pool_host
|
||||||
if self.scheduler.enable_hisparse
|
if self.scheduler.enable_hisparse
|
||||||
@@ -2244,7 +2247,7 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
# Decode-radix path: new requests already matched in
|
# Decode-radix path: new requests already matched in
|
||||||
# `pop_preallocated`. Retracted requests reset `last_node`,
|
# `pop_preallocated`. Retracted requests reset `last_node`,
|
||||||
# so re-match only when that state is missing.
|
# 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
|
tree_cache = self.tree_cache if req.last_node is None else None
|
||||||
else:
|
else:
|
||||||
tree_cache = self.tree_cache
|
tree_cache = self.tree_cache
|
||||||
@@ -2284,7 +2287,7 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
if self.enable_decode_hicache:
|
if self.enable_decode_hicache:
|
||||||
self.tree_cache.check_hicache_events()
|
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()
|
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
|
# try to resume retracted requests if there are enough space for another `num_reserved_decode_tokens` decode steps
|
||||||
@@ -2296,9 +2299,7 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
|
|
||||||
if not hasattr(self, "polling_count"):
|
if not hasattr(self, "polling_count"):
|
||||||
self.polling_count = 0
|
self.polling_count = 0
|
||||||
self.polling_interval = (
|
self.polling_interval = get_disagg().disaggregation_decode_polling_interval
|
||||||
self.server_args.disaggregation_decode_polling_interval
|
|
||||||
)
|
|
||||||
|
|
||||||
self.polling_count = (self.polling_count + 1) % self.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.io_struct import async_sock_send, wrap_as_pickle
|
||||||
from sglang.srt.managers.schedule_batch import Modality
|
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.server_args import PortArgs, ServerArgs
|
||||||
from sglang.srt.utils import random_uuid
|
from sglang.srt.utils import random_uuid
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_zmq_socket
|
from sglang.srt.utils.network import NetworkAddress, get_zmq_socket
|
||||||
@@ -117,13 +118,13 @@ class SGLangEncoderServer(SGLangEncoderServicer):
|
|||||||
context.set_details(error_msg)
|
context.set_details(error_msg)
|
||||||
return sglang_encoder_pb2.EncodeResponse()
|
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(
|
return sglang_encoder_pb2.EncodeResponse(
|
||||||
embedding_size=nbytes,
|
embedding_size=nbytes,
|
||||||
embedding_len=embedding_len,
|
embedding_len=embedding_len,
|
||||||
embedding_dim=embedding_dim,
|
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)
|
embedding_ports = list(request.embedding_port)
|
||||||
logger.info(f"embedding_port = {embedding_ports}")
|
logger.info(f"embedding_port = {embedding_ports}")
|
||||||
if not embedding_ports:
|
if not embedding_ports:
|
||||||
@@ -141,7 +142,7 @@ class SGLangEncoderServer(SGLangEncoderServicer):
|
|||||||
await asyncio.gather(*tasks)
|
await asyncio.gather(*tasks)
|
||||||
self.encoder.embedding_to_send.pop(request.req_id, None)
|
self.encoder.embedding_to_send.pop(request.req_id, None)
|
||||||
return sglang_encoder_pb2.EncodeResponse()
|
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 = (
|
embedding_port = (
|
||||||
request.embedding_port[0] if request.embedding_port else 0
|
request.embedding_port[0] if request.embedding_port else 0
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -63,7 +63,7 @@ from sglang.srt.observability.trace import (
|
|||||||
process_tracing_init,
|
process_tracing_init,
|
||||||
trace_set_thread_info,
|
trace_set_thread_info,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import publish
|
from sglang.srt.runtime_context import get_disagg, get_exec, get_mm, publish
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
PortArgs,
|
PortArgs,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
@@ -352,7 +352,7 @@ class MMEncoder:
|
|||||||
[], dtype=self._embedding_dtype
|
[], dtype=self._embedding_dtype
|
||||||
).element_size()
|
).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 (
|
from sglang.srt.mem_cache.storage.mooncake_store.embedding_cache_controller import (
|
||||||
EmbeddingCacheController,
|
EmbeddingCacheController,
|
||||||
)
|
)
|
||||||
@@ -370,15 +370,15 @@ class MMEncoder:
|
|||||||
self.mm_global_cache = None
|
self.mm_global_cache = None
|
||||||
|
|
||||||
# Pre-compute embedding metadata (needed by all ranks for mooncake)
|
# 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()
|
self._embedding_dims = self._infer_embedding_dims()
|
||||||
|
|
||||||
if self.rank == 0:
|
if self.rank == 0:
|
||||||
logger.info(
|
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.local_ip = get_local_ip_auto()
|
||||||
|
|
||||||
self.engine = get_mooncake_transfer_engine()
|
self.engine = get_mooncake_transfer_engine()
|
||||||
@@ -391,8 +391,8 @@ class MMEncoder:
|
|||||||
hostname=self.local_ip,
|
hostname=self.local_ip,
|
||||||
gpu_id=self.gpu_id,
|
gpu_id=self.gpu_id,
|
||||||
ib_device=(
|
ib_device=(
|
||||||
self.server_args.disaggregation_ib_device
|
get_disagg().disaggregation_ib_device
|
||||||
or self.server_args.mooncake_ib_device
|
or get_exec().moe.mooncake_ib_device
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -401,7 +401,7 @@ class MMEncoder:
|
|||||||
self.encode_dispatch_lock = asyncio.Lock()
|
self.encode_dispatch_lock = asyncio.Lock()
|
||||||
|
|
||||||
# Async mooncake state: track background VIT forward completion
|
# 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_ready_events: Dict[str, asyncio.Event] = {}
|
||||||
self._forward_results: Dict[str, dict] = {}
|
self._forward_results: Dict[str, dict] = {}
|
||||||
# when multiple decoder TP ranks call
|
# when multiple decoder TP ranks call
|
||||||
@@ -415,12 +415,12 @@ class MMEncoder:
|
|||||||
|
|
||||||
# Bind unified encode entry point based on backend and cache config
|
# Bind unified encode entry point based on backend and cache config
|
||||||
if self.mm_global_cache is not None:
|
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
|
self._encode_fn = self.encode_with_global_cache_mooncake
|
||||||
else:
|
else:
|
||||||
self._encode_fn = self.encode_with_global_cache
|
self._encode_fn = self.encode_with_global_cache
|
||||||
else:
|
else:
|
||||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||||
self._encode_fn = self.encode_with_mooncake
|
self._encode_fn = self.encode_with_mooncake
|
||||||
else:
|
else:
|
||||||
self._encode_fn = self.encode
|
self._encode_fn = self.encode
|
||||||
@@ -1710,7 +1710,7 @@ class MMEncoder:
|
|||||||
mm_item.set(k, _convert(v))
|
mm_item.set(k, _convert(v))
|
||||||
|
|
||||||
cache_hit = False
|
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:
|
if use_mm_cache:
|
||||||
mm_item.set_pad_value()
|
mm_item.set_pad_value()
|
||||||
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
|
mm_hash = MultiModalStaticCache.combine_hashes([mm_item.hash])
|
||||||
@@ -1806,7 +1806,7 @@ class MMEncoder:
|
|||||||
embedding_port=None,
|
embedding_port=None,
|
||||||
url=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
|
# Wait for async VIT forward completion if needed
|
||||||
req_id = mm_data.req_id
|
req_id = mm_data.req_id
|
||||||
if req_id in self._forward_ready_events:
|
if req_id in self._forward_ready_events:
|
||||||
@@ -1878,7 +1878,7 @@ class MMEncoder:
|
|||||||
logger.info(f"{endpoint = }")
|
logger.info(f"{endpoint = }")
|
||||||
|
|
||||||
# Serialize data
|
# Serialize data
|
||||||
if self.server_args.encoder_transfer_backend == "mooncake":
|
if get_disagg().encoder_transfer_backend == "mooncake":
|
||||||
# Mooncake already pushed the embedding via RDMA;
|
# Mooncake already pushed the embedding via RDMA;
|
||||||
new_mm_data = mm_data.copy_without_embedding()
|
new_mm_data = mm_data.copy_without_embedding()
|
||||||
serialized_data = pickle.dumps(new_mm_data)
|
serialized_data = pickle.dumps(new_mm_data)
|
||||||
@@ -1910,11 +1910,11 @@ class MMEncoder:
|
|||||||
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
|
await asyncio.get_event_loop().run_in_executor(self.executor, send_with_socket)
|
||||||
if (
|
if (
|
||||||
encoder_metrics_collector is not None
|
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(
|
encoder_metrics_collector.observe_transfer(
|
||||||
time.perf_counter() - _zmq_xfer_start,
|
time.perf_counter() - _zmq_xfer_start,
|
||||||
backend=self.server_args.encoder_transfer_backend,
|
backend=get_disagg().encoder_transfer_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def encode(
|
async def encode(
|
||||||
|
|||||||
@@ -60,6 +60,7 @@ from sglang.srt.observability.trace import (
|
|||||||
TraceReqContext,
|
TraceReqContext,
|
||||||
trace_set_thread_info,
|
trace_set_thread_info,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_schedule
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils.network import NetworkAddress
|
from sglang.srt.utils.network import NetworkAddress
|
||||||
|
|
||||||
@@ -327,7 +328,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
lambda ptr, size: self.engine.batch_register([ptr], [size]),
|
lambda ptr, size: self.engine.batch_register([ptr], [size]),
|
||||||
self.kv_args,
|
self.kv_args,
|
||||||
count,
|
count,
|
||||||
self.server_args.chunked_prefill_size,
|
get_schedule().chunked_prefill_size,
|
||||||
)
|
)
|
||||||
self.kv_buffer_tensors = None
|
self.kv_buffer_tensors = None
|
||||||
|
|
||||||
@@ -498,7 +499,7 @@ class MooncakeKVManager(CommonKVManager):
|
|||||||
room,
|
room,
|
||||||
self.transfer_infos,
|
self.transfer_infos,
|
||||||
self.kv_buffer_tensors,
|
self.kv_buffer_tensors,
|
||||||
self.server_args.chunked_prefill_size,
|
get_schedule().chunked_prefill_size,
|
||||||
self._staging_ctx.prefetch_requested,
|
self._staging_ctx.prefetch_requested,
|
||||||
self._staging_ctx.prefetch_sockets,
|
self._staging_ctx.prefetch_sockets,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -44,6 +44,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
resolve_dcp_dst_entry_indices,
|
resolve_dcp_dst_entry_indices,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.runtime_context import get_schedule
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -538,7 +539,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
|
lambda ptr, size: self._register_staging_memory(ptr, size, gpu_id),
|
||||||
self.kv_args,
|
self.kv_args,
|
||||||
count,
|
count,
|
||||||
self.server_args.chunked_prefill_size,
|
get_schedule().chunked_prefill_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _init_staging_allocator(self):
|
def _init_staging_allocator(self):
|
||||||
@@ -670,7 +671,7 @@ class NixlKVManager(CommonKVManager):
|
|||||||
room,
|
room,
|
||||||
self.transfer_infos,
|
self.transfer_infos,
|
||||||
self.kv_buffer_tensors,
|
self.kv_buffer_tensors,
|
||||||
self.server_args.chunked_prefill_size,
|
get_schedule().chunked_prefill_size,
|
||||||
self._staging_ctx.prefetch_requested,
|
self._staging_ctx.prefetch_requested,
|
||||||
self._staging_ctx.prefetch_sockets,
|
self._staging_ctx.prefetch_sockets,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -65,6 +65,7 @@ from sglang.srt.mem_cache.common import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
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.observability.req_time_stats import set_schedule_time_batch
|
||||||
|
from sglang.srt.runtime_context import get_disagg
|
||||||
from sglang.srt.utils import is_npu
|
from sglang.srt.utils import is_npu
|
||||||
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
from sglang.srt.utils.nvtx_utils import scheduler_nvtx_method
|
||||||
|
|
||||||
@@ -183,6 +184,9 @@ class PrefillBootstrapQueue:
|
|||||||
kv_args.engine_rank = self.tp_rank
|
kv_args.engine_rank = self.tp_rank
|
||||||
kv_args.pp_rank = self.pp_rank
|
kv_args.pp_rank = self.pp_rank
|
||||||
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
kv_args.system_dp_rank = self.scheduler.ps.dp_rank
|
||||||
|
kv_args.kv_cache_dtype_str = (
|
||||||
|
self.scheduler.tp_worker.model_runner.kv_cache_dtype_str
|
||||||
|
)
|
||||||
layer_shard_enabled = getattr(
|
layer_shard_enabled = getattr(
|
||||||
self.token_to_kv_pool, "layer_shard_enabled", False
|
self.token_to_kv_pool, "layer_shard_enabled", False
|
||||||
)
|
)
|
||||||
@@ -1249,7 +1253,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
|
|
||||||
def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None:
|
def optimistic_release_and_requeue(self: Scheduler, req: Req) -> None:
|
||||||
"""Release KV cache and requeue an optimistic prefill request."""
|
"""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)
|
maybe_cache_unfinished_req(req, self.tree_cache)
|
||||||
release_kv_cache(req, self.tree_cache)
|
release_kv_cache(req, self.tree_cache)
|
||||||
req.reset_for_retract()
|
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 (
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||||
is_in_tc_piecewise_cuda_graph,
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -25,7 +25,7 @@ class PyMscclppCommunicator:
|
|||||||
|
|
||||||
def _is_symm_mem_enabled(self) -> bool:
|
def _is_symm_mem_enabled(self) -> bool:
|
||||||
try:
|
try:
|
||||||
return get_server_args().enable_symm_mem
|
return get_exec().comm.enable_symm_mem
|
||||||
except ValueError:
|
except ValueError:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from torch.cuda.memory import (
|
|||||||
|
|
||||||
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
from sglang.srt.distributed.parallel_state import GroupCoordinator
|
||||||
from sglang.srt.environ import envs
|
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
|
from sglang.srt.utils.common import torch_release
|
||||||
|
|
||||||
after_2_8_0 = torch_release >= (2, 8)
|
after_2_8_0 = torch_release >= (2, 8)
|
||||||
@@ -159,7 +159,7 @@ _register_func = None
|
|||||||
|
|
||||||
def is_symmetric_memory_enabled():
|
def is_symmetric_memory_enabled():
|
||||||
try:
|
try:
|
||||||
return get_server_args().enable_symm_mem
|
return get_exec().comm.enable_symm_mem
|
||||||
except ValueError:
|
except ValueError:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ from sglang.srt.distributed.device_communicators.all_reduce_utils import (
|
|||||||
TORCH_SYMM_MEM_ALL_REDUCE_MAX_SIZES,
|
TORCH_SYMM_MEM_ALL_REDUCE_MAX_SIZES,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
|
from sglang.srt.runtime_context import get_exec
|
||||||
from sglang.srt.utils import is_cuda, is_hip
|
from sglang.srt.utils import is_cuda, is_hip
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -98,10 +99,9 @@ class TorchSymmMemCommunicator:
|
|||||||
# ([16384, 6144] bf16 = 192 MiB), including room for tail regions.
|
# ([16384, 6144] bf16 = 192 MiB), including room for tail regions.
|
||||||
if envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get():
|
if envs.SGLANG_OPT_USE_INKLING_CUSTOM_AR.get():
|
||||||
self.max_size = max(self.max_size, 256 * 1024 * 1024)
|
self.max_size = max(self.max_size, 256 * 1024 * 1024)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
if (
|
if (
|
||||||
get_server_args().enable_scattered_sconv
|
get_exec().comm.enable_scattered_sconv
|
||||||
or envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get()
|
or envs.SGLANG_OPT_USE_INKLING_FUSED_AR_SCONV.get()
|
||||||
):
|
):
|
||||||
# Fused extend kernels are out-of-place, so OUT must hold the
|
# 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.mem_cache.common import release_kv_cache
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -22,7 +23,7 @@ class SchedulerDllmMixin:
|
|||||||
def init_diffusion_llm(self: Scheduler):
|
def init_diffusion_llm(self: Scheduler):
|
||||||
self.dllm_config = (
|
self.dllm_config = (
|
||||||
DllmConfig.from_server_args(self.server_args)
|
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
|
else None
|
||||||
)
|
)
|
||||||
self.dllm_manager = DllmManager(dllm_config=self.dllm_config)
|
self.dllm_manager = DllmManager(dllm_config=self.dllm_config)
|
||||||
@@ -200,7 +201,7 @@ class SchedulerDllmMixin:
|
|||||||
self.chunked_prefill_size,
|
self.chunked_prefill_size,
|
||||||
running_bs if self.is_mixed_chunk else 0,
|
running_bs if self.is_mixed_chunk else 0,
|
||||||
self.priority_scheduling_preemption_threshold,
|
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,
|
dllm_config=self.dllm_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.srt.distributed.parallel_state import (
|
|||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
|
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.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.server_args import ServerArgs
|
||||||
from sglang.srt.utils.network import get_local_ip_auto
|
from sglang.srt.utils.network import get_local_ip_auto
|
||||||
|
|
||||||
@@ -111,7 +112,7 @@ class ExpertBackupClient:
|
|||||||
global_expert_location_metadata = get_global_expert_location_metadata()
|
global_expert_location_metadata = get_global_expert_location_metadata()
|
||||||
num_experts = (
|
num_experts = (
|
||||||
self.model_config.hf_config.n_routed_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
|
num_local_experts = num_experts // self.moe_ep_size
|
||||||
for i in range(self.engine_num):
|
for i in range(self.engine_num):
|
||||||
|
|||||||
@@ -377,9 +377,9 @@ class RuntimeHandle:
|
|||||||
model_config = self.tokenizer_manager.model_config
|
model_config = self.tokenizer_manager.model_config
|
||||||
result = {
|
result = {
|
||||||
"model_path": self.tokenizer_manager.model_path,
|
"model_path": self.tokenizer_manager.model_path,
|
||||||
"tokenizer_path": self.server_args.tokenizer_path,
|
"tokenizer_path": self.tokenizer_manager.server_args.tokenizer_path,
|
||||||
"is_generation": self.tokenizer_manager.is_generation,
|
"is_generation": self.tokenizer_manager.is_generation,
|
||||||
"weight_version": self.server_args.weight_version,
|
"weight_version": self.tokenizer_manager.server_args.weight_version,
|
||||||
"model_type": getattr(model_config.hf_config, "model_type", None),
|
"model_type": getattr(model_config.hf_config, "model_type", None),
|
||||||
"architectures": getattr(model_config.hf_config, "architectures", None),
|
"architectures": getattr(model_config.hf_config, "architectures", None),
|
||||||
}
|
}
|
||||||
@@ -432,7 +432,7 @@ class RuntimeHandle:
|
|||||||
"max_model_len": self.tokenizer_manager.model_config.context_len,
|
"max_model_len": self.tokenizer_manager.model_config.context_len,
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
if self.server_args.enable_lora and hasattr(
|
if self.tokenizer_manager.server_args.enable_lora and hasattr(
|
||||||
self.tokenizer_manager, "lora_registry"
|
self.tokenizer_manager, "lora_registry"
|
||||||
):
|
):
|
||||||
lora_registry = self.tokenizer_manager.lora_registry
|
lora_registry = self.tokenizer_manager.lora_registry
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ from sglang.srt.eplb.expert_location import (
|
|||||||
get_global_expert_location_metadata,
|
get_global_expert_location_metadata,
|
||||||
)
|
)
|
||||||
from sglang.srt.eplb.expert_location_updater import ExpertLocationUpdater
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
@@ -343,8 +343,8 @@ def update_expert_location_with_recovery(
|
|||||||
else:
|
else:
|
||||||
# Load the missing weights from disk
|
# Load the missing weights from disk
|
||||||
update_weights_from_disk_callable(
|
update_weights_from_disk_callable(
|
||||||
get_server_args().model_path,
|
get_model().model_path,
|
||||||
get_server_args().load_format,
|
get_model().load_format,
|
||||||
weight_name_filter=weight_name_filter,
|
weight_name_filter=weight_name_filter,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ from typing import Literal, Optional
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.eplb.expert_location import get_global_expert_location_metadata
|
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
|
@dataclass
|
||||||
@@ -40,8 +40,7 @@ class ExpertLocationDispatchInfo:
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def init_new(cls, layer_id: int):
|
def init_new(cls, layer_id: int):
|
||||||
server_args = get_server_args()
|
ep_dispatch_algorithm = get_exec().moe.ep_dispatch_algorithm
|
||||||
ep_dispatch_algorithm = server_args.ep_dispatch_algorithm
|
|
||||||
expert_location_metadata = get_global_expert_location_metadata()
|
expert_location_metadata = get_global_expert_location_metadata()
|
||||||
assert expert_location_metadata is not None
|
assert expert_location_metadata is not None
|
||||||
|
|
||||||
@@ -50,7 +49,7 @@ class ExpertLocationDispatchInfo:
|
|||||||
|
|
||||||
return cls(
|
return cls(
|
||||||
ep_dispatch_algorithm=ep_dispatch_algorithm,
|
ep_dispatch_algorithm=ep_dispatch_algorithm,
|
||||||
rank_invariant=server_args.moe_a2a_backend == "none",
|
rank_invariant=get_exec().moe.moe_a2a_backend == "none",
|
||||||
partial_logical_to_rank_dispatch_physical_map=(
|
partial_logical_to_rank_dispatch_physical_map=(
|
||||||
expert_location_metadata.logical_to_rank_dispatch_physical_map[
|
expert_location_metadata.logical_to_rank_dispatch_physical_map[
|
||||||
layer_id, :
|
layer_id, :
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from sglang.srt.eplb.expert_location import (
|
|||||||
ExpertLocationMetadata,
|
ExpertLocationMetadata,
|
||||||
get_global_expert_location_metadata,
|
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
|
from sglang.srt.utils import get_bool_env_var
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -109,7 +109,7 @@ def _update_expert_weights_with_canary(
|
|||||||
canary_tensor = (
|
canary_tensor = (
|
||||||
_get_canary_value(old_expert_location_metadata, layer_id)
|
_get_canary_value(old_expert_location_metadata, layer_id)
|
||||||
.clone()
|
.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)
|
routed_experts_weights_of_layer[layer_id].append(canary_tensor)
|
||||||
|
|
||||||
|
|||||||
@@ -19,6 +19,7 @@ from sglang.srt.model_executor.model_runner import ModelRunner
|
|||||||
from sglang.srt.model_executor.model_runner_components.layer_setup import (
|
from sglang.srt.model_executor.model_runner_components.layer_setup import (
|
||||||
ModelLayerInfo,
|
ModelLayerInfo,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_exec, get_memory, get_schedule
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -144,13 +145,13 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
(``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for
|
(``MlxAuxiliaryStateComponent``) raises ``NotImplementedError`` for
|
||||||
the mode.
|
the mode.
|
||||||
"""
|
"""
|
||||||
if self.server_args.disable_radix_cache:
|
if get_memory().disable_radix_cache:
|
||||||
return 1
|
return 1
|
||||||
return MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO
|
return MLX_AUX_STATE_SIZE_MAX_RUNNING_REQUESTS_RATIO
|
||||||
|
|
||||||
def _explicit_aux_state_size_per_worker(self) -> int | None:
|
def _explicit_aux_state_size_per_worker(self) -> int | None:
|
||||||
"""Return the explicit auxiliary-state cap for this attention-DP owner."""
|
"""Return the explicit auxiliary-state cap for this attention-DP owner."""
|
||||||
aux_state_size = self.server_args.max_mamba_cache_size
|
aux_state_size = get_schedule().max_mamba_cache_size
|
||||||
if aux_state_size is None:
|
if aux_state_size is None:
|
||||||
return None
|
return None
|
||||||
return aux_state_size // self.ps.attn_dp_size
|
return aux_state_size // self.ps.attn_dp_size
|
||||||
@@ -173,7 +174,7 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
Requires ``self.max_total_num_tokens`` to already be set.
|
Requires ``self.max_total_num_tokens`` to already be set.
|
||||||
"""
|
"""
|
||||||
capacity_cap = self.max_total_num_tokens // 2
|
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:
|
if requested is None:
|
||||||
requested_per_worker = None
|
requested_per_worker = None
|
||||||
resolved = min(capacity_cap, 4096)
|
resolved = min(capacity_cap, 4096)
|
||||||
@@ -189,7 +190,7 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
ratio = self._aux_state_slots_per_request()
|
ratio = self._aux_state_slots_per_request()
|
||||||
resolved = min(resolved, aux_state_size // ratio)
|
resolved = min(resolved, aux_state_size // ratio)
|
||||||
if resolved <= 0:
|
if resolved <= 0:
|
||||||
global_aux_state_size = self.server_args.max_mamba_cache_size
|
global_aux_state_size = get_schedule().max_mamba_cache_size
|
||||||
min_global_aux_state_size = ratio * self.ps.attn_dp_size
|
min_global_aux_state_size = ratio * self.ps.attn_dp_size
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"MLX auxiliary-state cache is too small to serve any "
|
f"MLX auxiliary-state cache is too small to serve any "
|
||||||
@@ -221,7 +222,7 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
from sglang.srt.utils.torch_memory_saver_adapter import TorchMemorySaverAdapter
|
||||||
|
|
||||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
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)
|
# Load model (sets metadata only)
|
||||||
@@ -267,7 +268,7 @@ class MlxModelRunnerStub(ModelRunner):
|
|||||||
# With the radix cache disabled no tree component exists to
|
# With the radix cache disabled no tree component exists to
|
||||||
# release auxiliary slots, so the pool owns their release
|
# release auxiliary slots, so the pool owns their release
|
||||||
# (see MlxAuxiliaryStateReqToTokenPool docstring).
|
# (see MlxAuxiliaryStateReqToTokenPool docstring).
|
||||||
owns_auxiliary_state_release=self.server_args.disable_radix_cache,
|
owns_auxiliary_state_release=get_memory().disable_radix_cache,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.req_to_token_pool = ReqToTokenPool(
|
self.req_to_token_pool = ReqToTokenPool(
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardBatch,
|
ForwardBatch,
|
||||||
PPProxyTensors,
|
PPProxyTensors,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_memory, get_model, get_schedule
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -53,19 +54,19 @@ class MlxTpModelWorker(TpModelWorker):
|
|||||||
|
|
||||||
logger.info("Initializing MlxModelRunner for end-to-end MLX inference")
|
logger.info("Initializing MlxModelRunner for end-to-end MLX inference")
|
||||||
init_kwargs = dict(
|
init_kwargs = dict(
|
||||||
model_path=self.server_args.model_path,
|
model_path=get_model().model_path,
|
||||||
trust_remote_code=self.server_args.trust_remote_code,
|
trust_remote_code=get_model().trust_remote_code,
|
||||||
disable_radix_cache=self.server_args.disable_radix_cache,
|
disable_radix_cache=get_memory().disable_radix_cache,
|
||||||
mem_fraction_static=self.server_args.mem_fraction_static,
|
mem_fraction_static=get_schedule().mem_fraction_static,
|
||||||
quantization=self.server_args.quantization,
|
quantization=get_model().quantization,
|
||||||
)
|
)
|
||||||
if self.server_args.max_total_tokens is not None:
|
if get_schedule().max_total_tokens is not None:
|
||||||
init_kwargs["pool_size"] = self.server_args.max_total_tokens
|
init_kwargs["pool_size"] = get_schedule().max_total_tokens
|
||||||
self._mlx_runner = MlxModelRunner(**init_kwargs)
|
self._mlx_runner = MlxModelRunner(**init_kwargs)
|
||||||
|
|
||||||
self._model_runner = MlxModelRunnerStub(
|
self._model_runner = MlxModelRunnerStub(
|
||||||
model_config=self.model_config,
|
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,
|
gpu_id=self.gpu_id,
|
||||||
ps=self.ps,
|
ps=self.ps,
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
|
|||||||
@@ -23,7 +23,7 @@ from sglang.srt.layers.utils.cp_utils import (
|
|||||||
cp_allgather_and_save_kv_cache,
|
cp_allgather_and_save_kv_cache,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -515,7 +515,7 @@ class MusaFlashAttentionBackend(FlashAttentionBackend):
|
|||||||
and not forward_batch.forward_mode.is_draft_extend_v2()
|
and not forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
):
|
):
|
||||||
if forward_batch.attn_attend_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
|
||||||
assert forward_batch.prefix_chunk_idx is not None
|
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_cu_seq_lens is not None
|
||||||
assert forward_batch.prefix_chunk_max_seq_lens is not None
|
assert forward_batch.prefix_chunk_max_seq_lens is not None
|
||||||
|
|||||||
@@ -26,7 +26,7 @@ from sglang.srt.layers.utils.cp_utils import cp_all_gather_rerange_kv_cache
|
|||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_flags
|
from sglang.srt.runtime_context import get_flags, get_spec
|
||||||
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
|
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
|
||||||
from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast
|
from sglang.srt.utils import get_bool_env_var, get_current_device_stream_fast
|
||||||
|
|
||||||
@@ -336,9 +336,7 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
self.use_fa = get_bool_env_var("ASCEND_USE_FA", "False")
|
self.use_fa = get_bool_env_var("ASCEND_USE_FA", "False")
|
||||||
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
|
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
|
||||||
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
self.enable_torch_compile = get_flags().capture.enable_torch_compile
|
||||||
self.speculative_num_draft_tokens = (
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
self.ascend_attn_mask_builder = AscendAttnMaskBuilder(
|
self.ascend_attn_mask_builder = AscendAttnMaskBuilder(
|
||||||
model_runner, self.device, self.use_fia, self.use_mla
|
model_runner, self.device, self.use_fia, self.use_mla
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from sglang.srt.layers.attention.dsv4.compressor import CompressorBackendMixin
|
|||||||
from sglang.srt.layers.attention.dsv4.indexer import C4IndexerBackendMixin
|
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_batch_info import DSV4OutCacheLoc, ForwardMode
|
||||||
from sglang.srt.model_executor.forward_context import get_attn_backend
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -1493,9 +1493,8 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
or forward_batch.forward_mode.is_draft_extend_v2()
|
||||||
):
|
):
|
||||||
B = forward_batch.batch_size
|
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(
|
actual_q = torch.arange(
|
||||||
n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device
|
n_draft, B * n_draft + 1, n_draft, dtype=torch.int32, device=device
|
||||||
)
|
)
|
||||||
@@ -1540,9 +1539,8 @@ class DeepseekV4AscendAttnBackend(
|
|||||||
forward_batch.forward_mode.is_target_verify()
|
forward_batch.forward_mode.is_target_verify()
|
||||||
or forward_batch.forward_mode.is_draft_extend_v2()
|
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:
|
else:
|
||||||
max_seqlen_q = 1
|
max_seqlen_q = 1
|
||||||
return self._kernel_metadata_from_parts(
|
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.layers.attention.vision import VisionAttention
|
||||||
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
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):
|
class ViTNpuGraphRunner(ViTCudaGraphRunner):
|
||||||
@@ -70,7 +70,7 @@ class ViTNpuGraphRunner(ViTCudaGraphRunner):
|
|||||||
graph = torch_npu.npu.NPUGraph()
|
graph = torch_npu.npu.NPUGraph()
|
||||||
vit = self.vit
|
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):
|
with torch_npu.npu.graph(graph, pool=ViTNpuGraphRunner._graph_memory_pool):
|
||||||
y = None
|
y = None
|
||||||
deepstack_outs: List[torch.Tensor] = []
|
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.hardware_backend.npu.utils import npu_format_cast
|
||||||
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
|
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPBuffer
|
||||||
from sglang.srt.layers.moe.utils import DeepEPMode
|
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:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
|
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()
|
envs.SGLANG_DEEPEP_NUM_MAX_DISPATCH_TOKENS_PER_RANK.get()
|
||||||
),
|
),
|
||||||
num_experts=layer.num_experts,
|
num_experts=layer.num_experts,
|
||||||
fuse_mode=get_server_args().fuseep_mode,
|
fuse_mode=get_exec().moe.fuseep_mode,
|
||||||
)
|
)
|
||||||
return hidden_states
|
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"``.
|
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 --
|
# -- The fused MoE optimization mode "1": dispatch_gmm_combine_decode --
|
||||||
if weight_prefix == "w13":
|
if weight_prefix == "w13":
|
||||||
cpu_w13 = layer.w13_weight.data.transpose(1, 2).cpu()
|
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(
|
layer.w2_weight_scale = torch.nn.Parameter(
|
||||||
w2_scale.to(torch.float32), requires_grad=False
|
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 --
|
# -- The fused MoE optimization mode "2": dispatch_ffn_combine --
|
||||||
if weight_prefix == "w13":
|
if weight_prefix == "w13":
|
||||||
w13_weight = _release_weight_cache(layer.w13_weight)
|
w13_weight = _release_weight_cache(layer.w13_weight)
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
Phase,
|
Phase,
|
||||||
check_cuda_graph_backend,
|
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 (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -130,7 +130,7 @@ logger = logging.getLogger(__name__)
|
|||||||
class SiluAndMul(MultiPlatformOp):
|
class SiluAndMul(MultiPlatformOp):
|
||||||
def __init__(self, *args, **kwargs):
|
def __init__(self, *args, **kwargs):
|
||||||
super().__init__(*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
|
self._forward_method = self.forward_native
|
||||||
elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get():
|
elif _use_aiter and envs.SGLANG_OPT_USE_AITER_SILU_MUL.get():
|
||||||
self._forward_method = self.forward_aiter
|
self._forward_method = self.forward_aiter
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||||
|
|
||||||
"""
|
"""
|
||||||
end to end attention solution with aiter kernels
|
end to end attention solution with aiter kernels
|
||||||
@@ -148,8 +148,8 @@ class AiterAttnBackend(AttentionBackend):
|
|||||||
|
|
||||||
self.device = model_runner.device
|
self.device = model_runner.device
|
||||||
self.is_multimodal = model_runner.model_config.is_multimodal
|
self.is_multimodal = model_runner.model_config.is_multimodal
|
||||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.topk = topk
|
self.topk = topk
|
||||||
self.num_head = (
|
self.num_head = (
|
||||||
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
|
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ from sglang.srt.layers.attention.dsv4.sparse_prefill_utils import (
|
|||||||
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
|
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||||
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
||||||
from sglang.srt.speculative.ragged_verify import (
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
RaggedVerifyMode,
|
RaggedVerifyMode,
|
||||||
@@ -537,9 +537,7 @@ class DeepseekV4AttnBackend(
|
|||||||
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
||||||
self.mtp_enabled = self.topk > 0
|
self.mtp_enabled = self.topk > 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens: int = (
|
self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
if self.speculative_num_draft_tokens is not None:
|
if self.speculative_num_draft_tokens is not None:
|
||||||
# Persistent target-verify metadata buffers. Allocated here (not
|
# Persistent target-verify metadata buffers. Allocated here (not
|
||||||
# lazily) so they are ordinary tensors: the first touch of a lazy
|
# lazily) so they are ordinary tensors: the first touch of a lazy
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ from sglang.srt.layers.attention.dsv4.metadata import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import DeepSeekV4TokenToKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||||
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
from sglang.srt.speculative.eagle_utils import per_step_draft_out_cache_loc
|
||||||
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
||||||
from sglang.srt.utils import ceil_align
|
from sglang.srt.utils import ceil_align
|
||||||
@@ -455,9 +455,7 @@ class DeepseekV4HipRadixBackend(
|
|||||||
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
assert self.topk in [0, 1], "MTP Topk > 1 not supported for DeepSeek V4"
|
||||||
self.mtp_enabled = self.topk > 0
|
self.mtp_enabled = self.topk > 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens: int = (
|
self.speculative_num_draft_tokens: int = get_spec().speculative_num_draft_tokens
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
self.forward_metadata: Union[
|
self.forward_metadata: Union[
|
||||||
DSV4Metadata,
|
DSV4Metadata,
|
||||||
|
|||||||
@@ -37,7 +37,14 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
|||||||
get_tc_piecewise_forward_context,
|
get_tc_piecewise_forward_context,
|
||||||
is_in_tc_piecewise_cuda_graph,
|
is_in_tc_piecewise_cuda_graph,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_device,
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
get_server_args,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.state_capturer.indexer_topk import (
|
from sglang.srt.state_capturer.indexer_topk import (
|
||||||
maybe_capture_indexer_topk,
|
maybe_capture_indexer_topk,
|
||||||
)
|
)
|
||||||
@@ -163,7 +170,7 @@ def _uses_dsa_attention_backend(forward_batch: ForwardBatch) -> bool:
|
|||||||
):
|
):
|
||||||
backend_name = (
|
backend_name = (
|
||||||
decode_backend
|
decode_backend
|
||||||
if server_args.speculative_attention_mode == "decode"
|
if get_spec().speculative_attention_mode == "decode"
|
||||||
else prefill_backend
|
else prefill_backend
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -460,7 +467,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
base=rope_theta, # type: ignore
|
base=rope_theta, # type: ignore
|
||||||
rope_scaling=rope_scaling,
|
rope_scaling=rope_scaling,
|
||||||
is_neox_style=is_neox_style,
|
is_neox_style=is_neox_style,
|
||||||
device=get_server_args().device,
|
device=get_device().device,
|
||||||
)
|
)
|
||||||
self.block_size = block_size
|
self.block_size = block_size
|
||||||
self.scale_fmt = scale_fmt
|
self.scale_fmt = scale_fmt
|
||||||
@@ -471,7 +478,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
self.num_local_tokens = getattr(config, "index_local_tokens", 0)
|
self.num_local_tokens = getattr(config, "index_local_tokens", 0)
|
||||||
|
|
||||||
self.paged_mqa_logits_backend = DSAPagedMQALogitsBackend.resolve(
|
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
|
@contextlib.contextmanager
|
||||||
@@ -1066,7 +1073,7 @@ class Indexer(MultiPlatformOp):
|
|||||||
total_mem = torch.cuda.get_device_properties(device_index).total_memory
|
total_mem = torch.cuda.get_device_properties(device_index).total_memory
|
||||||
|
|
||||||
total_mem_budget = int(total_mem * self._MQA_LOGITS_TOTAL_MEM_FRACTION)
|
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:
|
if mem_fraction_static is None:
|
||||||
static_budget = total_mem_budget
|
static_budget = total_mem_budget
|
||||||
else:
|
else:
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from typing import (
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
|
from sglang.srt.configs.model_config import get_dsa_index_topk, is_deepseek_dsa
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
|
from sglang.kernels.ops.attention.dsa.dequant_k_cache import (
|
||||||
@@ -469,9 +469,7 @@ class DeepseekSparseAttnBackend(
|
|||||||
# Speculative decoding
|
# Speculative decoding
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = (
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
self.use_fused_topk = should_use_dsa_fused_topk(
|
self.use_fused_topk = should_use_dsa_fused_topk(
|
||||||
model_runner.server_args, seed_dsa_topk_from_draft_extend
|
model_runner.server_args, seed_dsa_topk_from_draft_extend
|
||||||
|
|||||||
@@ -40,7 +40,7 @@ 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 (
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||||
is_in_tc_piecewise_cuda_graph,
|
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.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 import add_prefix, is_cuda, is_hip, is_xpu
|
||||||
from sglang.srt.utils.common import is_sm120_supported
|
from sglang.srt.utils.common import is_sm120_supported
|
||||||
@@ -922,9 +922,8 @@ class C4Indexer(nn.Module):
|
|||||||
self.rotary_emb = rotary_emb
|
self.rotary_emb = rotary_emb
|
||||||
self.freqs_cis = freqs_cis
|
self.freqs_cis = freqs_cis
|
||||||
self.weight_scale: float = self.softmax_scale * self.n_heads**-0.5
|
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
|
self.alt_streams = alt_streams
|
||||||
|
|
||||||
def compute_q(
|
def compute_q(
|
||||||
|
|||||||
@@ -30,7 +30,7 @@ from sglang.srt.layers.utils.cp_utils import (
|
|||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
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, get_spec
|
||||||
from sglang.srt.speculative.ragged_verify import build_ragged_target_verify_geometry
|
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_info import SpecInput, SpeculativeAlgorithm
|
||||||
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
from sglang.srt.speculative.spec_utils import resolve_num_tokens_per_req
|
||||||
@@ -204,9 +204,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
|
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = (
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
if (
|
if (
|
||||||
self.speculative_num_draft_tokens is not None
|
self.speculative_num_draft_tokens is not None
|
||||||
and model_runner.is_draft_worker
|
and model_runner.is_draft_worker
|
||||||
@@ -1513,7 +1511,7 @@ class FlashAttentionBackend(AttentionBackend):
|
|||||||
):
|
):
|
||||||
# Do multi-head attention with chunked prefix cache
|
# Do multi-head attention with chunked prefix cache
|
||||||
if forward_batch.attn_attend_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
|
# 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_idx is not None
|
||||||
assert forward_batch.prefix_chunk_cu_seq_lens is not None
|
assert forward_batch.prefix_chunk_cu_seq_lens is not None
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
from __future__ import annotations
|
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.
|
Support attention backend for flashinfer MLA.
|
||||||
@@ -33,7 +33,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 (
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||||
is_in_tc_piecewise_cuda_graph,
|
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_info import SpecInput
|
||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
draft_kv_indices_buffer_width,
|
draft_kv_indices_buffer_width,
|
||||||
@@ -224,9 +224,9 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
|||||||
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
self.token_to_kv_pool = model_runner.token_to_kv_pool
|
||||||
self.enable_chunk_kv = (
|
self.enable_chunk_kv = (
|
||||||
not skip_prefill
|
not skip_prefill
|
||||||
and get_server_args().disaggregation_mode != "decode"
|
and get_disagg().disaggregation_mode != "decode"
|
||||||
and not get_server_args().disable_chunked_prefix_cache
|
and not get_schedule().disable_chunked_prefix_cache
|
||||||
and not get_server_args().flashinfer_mla_disable_ragged
|
and not get_exec().kernel.flashinfer_mla_disable_ragged
|
||||||
)
|
)
|
||||||
self.page_size = model_runner.page_size
|
self.page_size = model_runner.page_size
|
||||||
|
|
||||||
@@ -402,7 +402,7 @@ class FlashInferMLAAttnBackend(AttentionBackend):
|
|||||||
prefix_lens = forward_batch.extend_prefix_lens
|
prefix_lens = forward_batch.extend_prefix_lens
|
||||||
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
|
||||||
use_ragged = (
|
use_ragged = (
|
||||||
not get_server_args().flashinfer_mla_disable_ragged
|
not get_exec().kernel.flashinfer_mla_disable_ragged
|
||||||
and extend_no_prefix
|
and extend_no_prefix
|
||||||
# Piecewise cuda graph should use paged prefill to be compatible with prefix cache
|
# Piecewise cuda graph should use paged prefill to be compatible with prefix cache
|
||||||
and not is_in_tc_piecewise_cuda_graph()
|
and not is_in_tc_piecewise_cuda_graph()
|
||||||
|
|||||||
@@ -20,7 +20,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import scaled_fp8_quant
|
|||||||
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
from sglang.srt.layers.attention.flashinfer_mla_backend import FlashInferMLAAttnBackend
|
||||||
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
|
from sglang.srt.layers.attention.verify_mask import VerifyMask, maybe_create_verify_mask
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -94,7 +94,7 @@ class FlashMLABackend(FlashInferMLAAttnBackend):
|
|||||||
torch.float8_e5m2,
|
torch.float8_e5m2,
|
||||||
}
|
}
|
||||||
|
|
||||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
|
|
||||||
self.cuda_graph_kv_indices = None
|
self.cuda_graph_kv_indices = None
|
||||||
self.cuda_graph_mla_metadata = None
|
self.cuda_graph_mla_metadata = None
|
||||||
|
|||||||
@@ -24,7 +24,7 @@ from sglang.srt.layers.radix_attention import RadixAttention
|
|||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool
|
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.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
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.eagle_info import EagleDraftInput, EagleVerifyInput
|
||||||
from sglang.srt.speculative.spec_info import SpecInput
|
from sglang.srt.speculative.spec_info import SpecInput
|
||||||
|
|
||||||
@@ -392,7 +392,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
"""Per-row (length bs) bool flush mask = the radix track's seq_lens_cpu %
|
"""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
|
mamba_track_interval == 0, so force-flush and snapshot fire on the same
|
||||||
steps (no off-by-one)."""
|
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:
|
if seq_lens_cpu is None:
|
||||||
# Should not happen for the supported config; stay safe and never flush.
|
# Should not happen for the supported config; stay safe and never flush.
|
||||||
return torch.zeros((bs,), dtype=torch.bool)
|
return torch.zeros((bs,), dtype=torch.bool)
|
||||||
@@ -823,7 +823,7 @@ class Mamba2AttnBackend(MambaAttnBackendBase):
|
|||||||
# Page-major stores state strided; only the stride-aware Triton causal-conv
|
# 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.
|
# reads it (CUDA causal_conv1d garbles it). A model may also force Triton.
|
||||||
use_triton_causal_conv = (
|
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)
|
layer_cache = self.req_to_token_pool.mamba2_layer_cache(layer_id)
|
||||||
mixer_out, intermediate_states = mixer.forward(
|
mixer_out, intermediate_states = mixer.forward(
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
|||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.runtime_context import get_spec
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -59,7 +60,7 @@ class IntelAMXAttnBackend(AttentionBackend):
|
|||||||
self.num_kv_splits = 8
|
self.num_kv_splits = 8
|
||||||
|
|
||||||
# speculative decoding params
|
# speculative decoding params
|
||||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
|
|
||||||
def _build_extend_metadata(self, forward_batch: ForwardBatch):
|
def _build_extend_metadata(self, forward_batch: ForwardBatch):
|
||||||
"""Resolve (seq_lens, extend_seq_lens, extend_start_loc, tree_mask) for
|
"""Resolve (seq_lens, extend_seq_lens, extend_start_loc, tree_mask) for
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ from sglang.srt.models.inkling_common.kernels.sconv import (
|
|||||||
fused_extend_sconv_metadata,
|
fused_extend_sconv_metadata,
|
||||||
precompute_helion_extend_metadata,
|
precompute_helion_extend_metadata,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_exec, get_server_args, get_spec
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -117,7 +117,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
|
|||||||
growing a buffer after a graph captured it moves the address that graph
|
growing a buffer after a graph captured it moves the address that graph
|
||||||
reads, and prefill captures before the decode runner reports its bounds."""
|
reads, and prefill captures before the decode runner reports its bounds."""
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
cuda_graph_config = server_args.cuda_graph_config
|
cuda_graph_config = get_exec().graph.cuda_graph_config
|
||||||
decode_bs: list[int] = []
|
decode_bs: list[int] = []
|
||||||
prefill_tokens: list[int] = []
|
prefill_tokens: list[int] = []
|
||||||
decode_max_bs = 0
|
decode_max_bs = 0
|
||||||
@@ -125,7 +125,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
|
|||||||
decode_bs = list(cuda_graph_config.decode.bs or [])
|
decode_bs = list(cuda_graph_config.decode.bs or [])
|
||||||
prefill_tokens = list(cuda_graph_config.prefill.bs or [])
|
prefill_tokens = list(cuda_graph_config.prefill.bs or [])
|
||||||
decode_max_bs = cuda_graph_config.decode.max_bs or 0
|
decode_max_bs = cuda_graph_config.decode.max_bs or 0
|
||||||
draft_token_num = server_args.speculative_num_draft_tokens or 1
|
draft_token_num = get_spec().speculative_num_draft_tokens or 1
|
||||||
# req_to_token_pool.size is the runner's max_bs for both graph phases.
|
# req_to_token_pool.size is the runner's max_bs for both graph phases.
|
||||||
max_bs = max([self.req_to_token_pool.size, decode_max_bs, *decode_bs])
|
max_bs = max([self.req_to_token_pool.size, decode_max_bs, *decode_bs])
|
||||||
max_tokens = max([max_bs, *prefill_tokens, max_bs * draft_token_num])
|
max_tokens = max([max_bs, *prefill_tokens, max_bs * draft_token_num])
|
||||||
|
|||||||
@@ -37,7 +37,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
cuda_graph_fully_disabled,
|
cuda_graph_fully_disabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
draft_kv_indices_buffer_width,
|
draft_kv_indices_buffer_width,
|
||||||
draft_kv_indices_used_len,
|
draft_kv_indices_used_len,
|
||||||
@@ -168,9 +168,9 @@ class TritonAttnBackend(AttentionBackend):
|
|||||||
self._translate_kv_loc = getattr(
|
self._translate_kv_loc = getattr(
|
||||||
self.token_to_kv_pool_allocator, "translate_kv_loc_dense", None
|
self.token_to_kv_pool_allocator, "translate_kv_loc_dense", None
|
||||||
) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None)
|
) or getattr(self.token_to_kv_pool_allocator, "translate_kv_loc", None)
|
||||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self.speculative_num_steps = model_runner.server_args.speculative_num_steps
|
self.speculative_num_steps = get_spec().speculative_num_steps
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = get_spec().speculative_eagle_topk or 0
|
||||||
# Split-KV verify is bit-equivalent only for a pure-causal chain (topk==1)
|
# Split-KV verify is bit-equivalent only for a pure-causal chain (topk==1)
|
||||||
# and is gfx95-only; else fall back to extend_attention_fwd.
|
# and is gfx95-only; else fall back to extend_attention_fwd.
|
||||||
self.use_verify_splitkv = (
|
self.use_verify_splitkv = (
|
||||||
|
|||||||
@@ -34,7 +34,7 @@ from sglang.srt.layers.radix_attention import AttentionType
|
|||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_buffer
|
from sglang.srt.runtime_context import get_buffer, get_spec
|
||||||
from sglang.srt.speculative.ragged_verify import (
|
from sglang.srt.speculative.ragged_verify import (
|
||||||
build_ragged_target_verify_geometry,
|
build_ragged_target_verify_geometry,
|
||||||
resolve_ragged_verify_layout,
|
resolve_ragged_verify_layout,
|
||||||
@@ -162,9 +162,7 @@ class TRTLLMHAAttnBackend(FlashInferAttnBackend):
|
|||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
self.target_verify_metadata = {}
|
self.target_verify_metadata = {}
|
||||||
|
|
||||||
self.speculative_num_draft_tokens = (
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
# True iff the model declares ENCODER_ONLY (bidirectional) layers, which
|
# True iff the model declares ENCODER_ONLY (bidirectional) layers, which
|
||||||
# need the expanded TARGET_VERIFY metadata (TRTLLMMHAMetadata.encoder_*).
|
# need the expanded TARGET_VERIFY metadata (TRTLLMMHAMetadata.encoder_*).
|
||||||
self.expand_encoder_only_verify = any(
|
self.expand_encoder_only_verify = any(
|
||||||
|
|||||||
@@ -41,7 +41,12 @@ from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMo
|
|||||||
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph import (
|
||||||
is_in_tc_piecewise_cuda_graph,
|
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,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2
|
from sglang.srt.utils import is_flashinfer_available, is_float4_e2m1fn_x2
|
||||||
|
|
||||||
if is_flashinfer_available():
|
if is_flashinfer_available():
|
||||||
@@ -238,11 +243,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
|
|||||||
self.forward_prefill_metadata: Optional[TRTLLMMLAPrefillMetadata] = None
|
self.forward_prefill_metadata: Optional[TRTLLMMLAPrefillMetadata] = None
|
||||||
self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None
|
self.forward_decode_metadata: Union[TRTLLMMLADecodeMetadata, None] = None
|
||||||
|
|
||||||
self.disable_chunked_prefix_cache = (
|
self.disable_chunked_prefix_cache = get_schedule().disable_chunked_prefix_cache
|
||||||
get_server_args().disable_chunked_prefix_cache
|
|
||||||
)
|
|
||||||
|
|
||||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
self._verify_mask = None
|
self._verify_mask = None
|
||||||
# Tree-mask scratch is fetched from the target backend only.
|
# Tree-mask scratch is fetched from the target backend only.
|
||||||
self.is_draft_runner = model_runner.is_draft_worker
|
self.is_draft_runner = model_runner.is_draft_worker
|
||||||
|
|||||||
@@ -17,7 +17,7 @@ from sglang.kernels.ops.layernorm.norm import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.models.utils import apply_qk_norm
|
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 (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -88,7 +88,6 @@ from sglang.srt.layers.linear import (
|
|||||||
from sglang.srt.layers.quantization import QuantizationConfig
|
from sglang.srt.layers.quantization import QuantizationConfig
|
||||||
from sglang.srt.layers.rotary_embedding import apply_rotary_pos_emb
|
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.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
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
@@ -1047,7 +1046,7 @@ class VisionAttention(nn.Module):
|
|||||||
# Select attention backend via a unified method
|
# Select attention backend via a unified method
|
||||||
_passed_backend = qkv_backend
|
_passed_backend = qkv_backend
|
||||||
qkv_backend = self._determine_attention_backend(_passed_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"Multimodal attention backend not set. Use {qkv_backend}.")
|
||||||
print_info_once(f"Using {qkv_backend} as multimodal attention backend.")
|
print_info_once(f"Using {qkv_backend} as multimodal attention backend.")
|
||||||
|
|
||||||
@@ -1126,7 +1125,7 @@ class VisionAttention(nn.Module):
|
|||||||
weight_dtype=torch.float32,
|
weight_dtype=torch.float32,
|
||||||
cast_x_before_out_mul=True,
|
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 {}
|
else {}
|
||||||
)
|
)
|
||||||
q_norm = RMSNorm(
|
q_norm = RMSNorm(
|
||||||
@@ -1154,7 +1153,7 @@ class VisionAttention(nn.Module):
|
|||||||
- CUDA (other): "triton_attn"
|
- CUDA (other): "triton_attn"
|
||||||
- Non-CUDA: "sdpa"
|
- Non-CUDA: "sdpa"
|
||||||
"""
|
"""
|
||||||
override_backend = get_server_args().mm_attention_backend
|
override_backend = get_mm().mm_attention_backend
|
||||||
if override_backend is not None:
|
if override_backend is not None:
|
||||||
backend = override_backend
|
backend = override_backend
|
||||||
elif passed_backend is not None:
|
elif passed_backend is not None:
|
||||||
@@ -1259,7 +1258,7 @@ class VisionAttention(nn.Module):
|
|||||||
x = x.unsqueeze(0)
|
x = x.unsqueeze(0)
|
||||||
assert x.dim() == 3, x.shape
|
assert x.dim() == 3, x.shape
|
||||||
if (
|
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
|
and position_embeddings is not None
|
||||||
):
|
):
|
||||||
assert isinstance(position_embeddings, tuple), (
|
assert isinstance(position_embeddings, tuple), (
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from sglang.kernels.ops.kvcache.kv_indices import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_parallel, get_spec
|
||||||
from sglang.srt.utils import get_bool_env_var, get_device_core_count
|
from sglang.srt.utils import get_bool_env_var, get_device_core_count
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -92,7 +92,7 @@ class WaveAttnBackend(AttentionBackend):
|
|||||||
(max_bs + 1,), dtype=torch.int64, device=model_runner.device
|
(max_bs + 1,), dtype=torch.int64, device=model_runner.device
|
||||||
)
|
)
|
||||||
|
|
||||||
self.num_draft_tokens = model_runner.server_args.speculative_num_draft_tokens
|
self.num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
|
|
||||||
self.num_head = (
|
self.num_head = (
|
||||||
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
|
model_runner.model_config.num_attention_heads // get_parallel().attn_tp_size
|
||||||
|
|||||||
@@ -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.memory_pool import KVWriteLoc
|
||||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
|
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, get_spec
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
@@ -86,9 +86,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
self.topk = model_runner.server_args.speculative_eagle_topk or 0
|
||||||
self.speculative_num_steps = speculative_num_steps
|
self.speculative_num_steps = speculative_num_steps
|
||||||
self.speculative_num_draft_tokens = (
|
self.speculative_num_draft_tokens = get_spec().speculative_num_draft_tokens
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
|
||||||
)
|
|
||||||
self.speculative_step_id = speculative_step_id
|
self.speculative_step_id = speculative_step_id
|
||||||
|
|
||||||
# Local attention settings
|
# Local attention settings
|
||||||
@@ -638,7 +636,7 @@ class XPUAttentionBackend(AttentionBackend):
|
|||||||
):
|
):
|
||||||
# Do multi-head attention with chunked prefix cache
|
# Do multi-head attention with chunked prefix cache
|
||||||
if forward_batch.attn_attend_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
|
# 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_idx is not None
|
||||||
assert forward_batch.prefix_chunk_cu_seq_lens is not None
|
assert forward_batch.prefix_chunk_cu_seq_lens is not None
|
||||||
|
|||||||
@@ -73,7 +73,13 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
check_cuda_graph_backend,
|
check_cuda_graph_backend,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
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.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -169,7 +175,7 @@ def apply_flashinfer_allreduce_fusion(batch_size: int):
|
|||||||
(_is_sm90_supported or _is_sm100_supported)
|
(_is_sm90_supported or _is_sm100_supported)
|
||||||
and _is_flashinfer_available
|
and _is_flashinfer_available
|
||||||
and not is_dp_attention_enabled()
|
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()
|
and not is_flashinfer_allreduce_unavailable()
|
||||||
# Symbolic size checks stay last: under Dynamo tracing they guard on
|
# Symbolic size checks stay last: under Dynamo tracing they guard on
|
||||||
# the dynamic token dim, so statically-off configs must short-circuit
|
# the dynamic token dim, so statically-off configs must short-circuit
|
||||||
@@ -190,7 +196,7 @@ def apply_aiter_all_reduce_fusion(input_tensor: torch.Tensor):
|
|||||||
and total_bytes <= 8 * 1024 * 8192
|
and total_bytes <= 8 * 1024 * 8192
|
||||||
and get_parallel().tp_size != 6
|
and get_parallel().tp_size != 6
|
||||||
and not is_dp_attention_enabled()
|
and not is_dp_attention_enabled()
|
||||||
and get_server_args().enable_aiter_allreduce_fusion
|
and get_exec().comm.enable_aiter_allreduce_fusion
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -278,7 +284,7 @@ class AttnTpContext:
|
|||||||
and get_moe_a2a_backend().is_none()
|
and get_moe_a2a_backend().is_none()
|
||||||
and not enable_moe_dense_fully_dp()
|
and not enable_moe_dense_fully_dp()
|
||||||
and not check_cuda_graph_backend(Phase.PREFILL, Backend.TC_PIECEWISE)
|
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 get_server_args().enable_attn_tp_input_scattered:
|
||||||
if not self.allow_input_scattered:
|
if not self.allow_input_scattered:
|
||||||
@@ -411,7 +417,7 @@ class LayerScatterModes:
|
|||||||
not context.is_layer_sparse
|
not context.is_layer_sparse
|
||||||
and context.is_next_layer_sparse
|
and context.is_next_layer_sparse
|
||||||
and enable_moe_dense_fully_dp()
|
and enable_moe_dense_fully_dp()
|
||||||
and get_server_args().enable_two_batch_overlap
|
and get_exec().overlap.enable_two_batch_overlap
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -475,7 +481,7 @@ class LayerCommunicator:
|
|||||||
)
|
)
|
||||||
self._post_init_communicate()
|
self._post_init_communicate()
|
||||||
self._speculative_algo = SpeculativeAlgorithm.from_string(
|
self._speculative_algo = SpeculativeAlgorithm.from_string(
|
||||||
get_server_args().speculative_algorithm
|
get_spec().speculative_algorithm
|
||||||
)
|
)
|
||||||
|
|
||||||
def _post_init_communicate(self):
|
def _post_init_communicate(self):
|
||||||
@@ -846,7 +852,7 @@ class LayerCommunicator:
|
|||||||
and get_parallel().tp_size != 6
|
and get_parallel().tp_size != 6
|
||||||
and not is_dp_attention_enabled()
|
and not is_dp_attention_enabled()
|
||||||
and get_moe_a2a_backend().is_none()
|
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)
|
and (not self.is_last_layer)
|
||||||
@@ -1151,7 +1157,7 @@ class CommunicateWithAllReduceAndLayerNormFn:
|
|||||||
if not handled:
|
if not handled:
|
||||||
quantize_communications = (
|
quantize_communications = (
|
||||||
not forward_batch.forward_mode.is_decode_or_idle()
|
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:
|
if quantize_communications:
|
||||||
hidden_states = attention_tensor_model_parallel_quant_all_reduce(
|
hidden_states = attention_tensor_model_parallel_quant_all_reduce(
|
||||||
|
|||||||
@@ -53,7 +53,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import KVWriteLoc
|
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.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
|
@dataclass
|
||||||
@@ -208,10 +208,8 @@ class ZigzagCPStrategy(ContextParallelStrategy):
|
|||||||
actual_seq_q_prev_list.append(block_sizes[cp_rank])
|
actual_seq_q_prev_list.append(block_sizes[cp_rank])
|
||||||
actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1])
|
actual_seq_q_next_list.append(block_sizes[cp_segment_num - cp_rank - 1])
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
device = torch.device(get_server_args().device)
|
device = torch.device(get_device().device)
|
||||||
except Exception:
|
except Exception:
|
||||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||||
cu_prev = [0] + list(accumulate(actual_seq_q_prev_list))
|
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.layout import update_local_kv_lens_for_dcp
|
||||||
from sglang.srt.layers.dcp.metadata import DecodeContextParallelMetadata
|
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(
|
def prepare_decode_context_parallel_metadata(
|
||||||
@@ -53,12 +53,12 @@ def prepare_decode_context_parallel_metadata(
|
|||||||
extend_prefix_starts = torch.zeros(
|
extend_prefix_starts = torch.zeros(
|
||||||
len(seq_lens),
|
len(seq_lens),
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=get_server_args().device,
|
device=get_device().device,
|
||||||
)
|
)
|
||||||
extend_cu_prefix_lens = torch.zeros(
|
extend_cu_prefix_lens = torch.zeros(
|
||||||
len(seq_lens) + 1,
|
len(seq_lens) + 1,
|
||||||
dtype=torch.int32,
|
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[1:] = torch.cumsum(extend_prefix_lens, dim=0)
|
||||||
extend_cu_prefix_lens = extend_cu_prefix_lens[:-1]
|
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(
|
dcp_prefix_kv_indices = torch.empty(
|
||||||
sum(extend_prefix_lens_cpu),
|
sum(extend_prefix_lens_cpu),
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=get_server_args().device,
|
device=get_device().device,
|
||||||
)
|
)
|
||||||
create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)](
|
create_chunked_prefix_cache_kv_indices_fn[(len(seq_lens),)](
|
||||||
req_to_token,
|
req_to_token,
|
||||||
@@ -81,20 +81,20 @@ def prepare_decode_context_parallel_metadata(
|
|||||||
dcp_kv_indptr = torch.zeros(
|
dcp_kv_indptr = torch.zeros(
|
||||||
len(seq_lens) + 1,
|
len(seq_lens) + 1,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=get_server_args().device,
|
device=get_device().device,
|
||||||
)
|
)
|
||||||
dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0)
|
dcp_kv_indptr[1:] = seq_lens.cumsum(dim=0)
|
||||||
dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)]
|
dcp_kv_indptr = dcp_kv_indptr[: (len(seq_lens) + 1)]
|
||||||
dcp_kv_indices = torch.zeros(
|
dcp_kv_indices = torch.zeros(
|
||||||
seq_lens_sum,
|
seq_lens_sum,
|
||||||
dtype=torch.int32,
|
dtype=torch.int32,
|
||||||
device=get_server_args().device,
|
device=get_device().device,
|
||||||
)
|
)
|
||||||
|
|
||||||
extend_cu_lens = torch.zeros(
|
extend_cu_lens = torch.zeros(
|
||||||
len(seq_lens) + 1,
|
len(seq_lens) + 1,
|
||||||
dtype=torch.int32,
|
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[1:] = torch.cumsum(extend_seq_lens, dim=0)
|
||||||
extend_cu_lens = extend_cu_lens[:-1]
|
extend_cu_lens = extend_cu_lens[:-1]
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
Phase,
|
Phase,
|
||||||
check_cuda_graph_backend,
|
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 (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -223,7 +223,7 @@ def _forward_with_allreduce_fusion(
|
|||||||
return fused_result
|
return fused_result
|
||||||
|
|
||||||
# For AITER route, preserve correctness when fused path is unavailable.
|
# 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)
|
x = tensor_model_parallel_all_reduce(x)
|
||||||
return norm_module.forward(x, residual, None)
|
return norm_module.forward(x, residual, None)
|
||||||
|
|
||||||
@@ -425,7 +425,7 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if (
|
if (
|
||||||
residual is not None
|
residual is not None
|
||||||
or self.cast_x_before_out_mul
|
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 self.forward_native(x, residual, post_residual_addition)
|
||||||
out = rms_norm_batch_invariant(
|
out = rms_norm_batch_invariant(
|
||||||
@@ -532,7 +532,7 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if (
|
if (
|
||||||
residual is not None
|
residual is not None
|
||||||
or self.cast_x_before_out_mul
|
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)
|
or (self._fused_pad_kernel is not None and self.x_pad_to_multiple > 0)
|
||||||
):
|
):
|
||||||
return self.forward_native(x, residual, post_residual_addition)
|
return self.forward_native(x, residual, post_residual_addition)
|
||||||
@@ -593,7 +593,7 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if (
|
if (
|
||||||
residual is not None
|
residual is not None
|
||||||
or self.cast_x_before_out_mul
|
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 self.forward_native(x, residual, post_residual_addition)
|
||||||
return rms_norm_batch_invariant(
|
return rms_norm_batch_invariant(
|
||||||
@@ -720,7 +720,10 @@ class RMSNorm(MultiPlatformOp):
|
|||||||
if self.variance_size_override is not None:
|
if self.variance_size_override is not None:
|
||||||
return self.forward_native(x, residual, post_residual_addition)
|
return self.forward_native(x, residual, post_residual_addition)
|
||||||
if is_batch_invariant_mode_enabled():
|
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 self.forward_native(x, residual, post_residual_addition)
|
||||||
return rms_norm_batch_invariant(
|
return rms_norm_batch_invariant(
|
||||||
x,
|
x,
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ from sglang.srt.layers.parameter import (
|
|||||||
_ColumnvLLMParameter,
|
_ColumnvLLMParameter,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.utils import pad_or_narrow_weight
|
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
|
from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -1597,7 +1597,7 @@ class RowParallelLinear(LinearBase):
|
|||||||
quantize_communications = (
|
quantize_communications = (
|
||||||
(
|
(
|
||||||
not forward_batch.forward_mode.is_decode_or_idle()
|
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
|
if forward_batch is not None
|
||||||
else False
|
else False
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardBatch,
|
ForwardBatch,
|
||||||
ForwardMode,
|
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 (
|
from sglang.srt.utils.common import (
|
||||||
is_cpu,
|
is_cpu,
|
||||||
is_npu,
|
is_npu,
|
||||||
@@ -350,7 +350,7 @@ class LogitsProcessor(nn.Module):
|
|||||||
self.vocab_size = config.vocab_size
|
self.vocab_size = config.vocab_size
|
||||||
self.logit_scale = logit_scale
|
self.logit_scale = logit_scale
|
||||||
self.use_attn_tp_group = get_server_args().enable_dp_lm_head
|
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:
|
if self.use_attn_tp_group:
|
||||||
self.attn_tp_size = get_parallel().attn_tp_size
|
self.attn_tp_size = get_parallel().attn_tp_size
|
||||||
self.do_tensor_parallel_all_gather = (
|
self.do_tensor_parallel_all_gather = (
|
||||||
@@ -374,8 +374,8 @@ class LogitsProcessor(nn.Module):
|
|||||||
self.final_logit_softcapping = None
|
self.final_logit_softcapping = None
|
||||||
|
|
||||||
self.return_full_logits = return_full_logits
|
self.return_full_logits = return_full_logits
|
||||||
self.enable_mis = get_server_args().enable_mis
|
self.enable_mis = get_exec().features.enable_mis
|
||||||
self.rl_on_policy_target = get_server_args().rl_on_policy_target
|
self.rl_on_policy_target = get_exec().deterministic.rl_on_policy_target
|
||||||
|
|
||||||
self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer(
|
self._logits_gatherer = triton_symm_mem_ag.MultimemAllGatherer(
|
||||||
max_tokens=triton_symm_mem_ag.recommended_max_tokens(
|
max_tokens=triton_symm_mem_ag.recommended_max_tokens(
|
||||||
|
|||||||
@@ -71,6 +71,7 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight
|
from sglang.srt.model_loader.weight_utils import narrow_padded_param_and_loaded_weight
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
get_global_dwdp_manager,
|
get_global_dwdp_manager,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
@@ -260,7 +261,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
|
|
||||||
self._num_global_routed = num_experts - num_shared_slots
|
self._num_global_routed = num_experts - num_shared_slots
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
if server_args.ep_join_mode == "scale":
|
if get_exec().moe.ep_join_mode == "scale":
|
||||||
storage_ep_size = server_args.elastic_ep_initial_size
|
storage_ep_size = server_args.elastic_ep_initial_size
|
||||||
assert storage_ep_size is not None
|
assert storage_ep_size is not None
|
||||||
self._expert_storage_rank = (
|
self._expert_storage_rank = (
|
||||||
@@ -359,7 +360,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
print_info_once(
|
print_info_once(
|
||||||
"FlashInfer TRTLLM MoE deferred finalize is "
|
"FlashInfer TRTLLM MoE deferred finalize is "
|
||||||
f"{'enabled' if self.supports_deferred_finalize else 'disabled'} "
|
f"{'enabled' if self.supports_deferred_finalize else 'disabled'} "
|
||||||
f"(moe_runner_backend={server_args.moe_runner_backend}, "
|
f"(moe_runner_backend={get_exec().moe.moe_runner_backend}, "
|
||||||
f"quant_method={type(self.quant_method).__name__})."
|
f"quant_method={type(self.quant_method).__name__})."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from sglang.srt.layers.moe.topk import (
|
|||||||
remap_topk_for_per_rank_shared_slots,
|
remap_topk_for_per_rank_shared_slots,
|
||||||
)
|
)
|
||||||
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.runtime_context import get_exec
|
||||||
from sglang.srt.utils import is_hip, is_npu
|
from sglang.srt.utils import is_hip, is_npu
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -44,10 +45,9 @@ class HashTopK(nn.Module):
|
|||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
self.enable_waterfill = (
|
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
|
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.dp_attention import is_allocation_symmetric
|
||||||
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
|
||||||
from sglang.srt.layers.moe.utils import get_moe_padding_size
|
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 (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -531,7 +531,7 @@ def _fused_moe_kernel_sequence(
|
|||||||
out_hidden_states = torch.empty_like(hidden_states)
|
out_hidden_states = torch.empty_like(hidden_states)
|
||||||
|
|
||||||
use_fused_moe_sum_all_reduce = (
|
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 (not no_combine)
|
||||||
and (topk > 2)
|
and (topk > 2)
|
||||||
and (not use_int8_w8a16)
|
and (not use_int8_w8a16)
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from typing import Any, Dict, List, Optional, Tuple
|
|||||||
import torch
|
import torch
|
||||||
import triton
|
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
|
from sglang.srt.utils import get_device_name, is_hip
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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
|
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.
|
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(
|
logger.warning(
|
||||||
"Deterministic inference is enabled, using default MoE kernel config."
|
"Deterministic inference is enabled, using default MoE kernel config."
|
||||||
)
|
)
|
||||||
@@ -187,7 +187,7 @@ def get_default_config(
|
|||||||
is_marlin: bool,
|
is_marlin: bool,
|
||||||
block_shape: Optional[List[int]] = None,
|
block_shape: Optional[List[int]] = None,
|
||||||
) -> Dict[str, int]:
|
) -> Dict[str, int]:
|
||||||
if get_server_args().enable_deterministic_inference:
|
if get_exec().deterministic.enable_deterministic_inference:
|
||||||
config = {
|
config = {
|
||||||
"BLOCK_SIZE_M": 64,
|
"BLOCK_SIZE_M": 64,
|
||||||
"BLOCK_SIZE_N": 64,
|
"BLOCK_SIZE_N": 64,
|
||||||
|
|||||||
@@ -27,7 +27,7 @@ from sglang.srt.layers.moe.topk import (
|
|||||||
TopKOutputChecker,
|
TopKOutputChecker,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.utils import get_moe_runner_backend
|
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.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils import get_int_env_var
|
from sglang.srt.utils import get_int_env_var
|
||||||
|
|
||||||
@@ -123,7 +123,7 @@ class FlashinferDispatcher(BaseDispatcher):
|
|||||||
# max_running_requests is not yet resolved at model-construction time,
|
# 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
|
# 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).
|
# (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)
|
default_max_tokens = max(cps if cps and cps > 0 else 4096, 4096)
|
||||||
self.max_num_tokens = get_int_env_var(
|
self.max_num_tokens = get_int_env_var(
|
||||||
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK",
|
"SGLANG_FLASHINFER_NUM_MAX_DISPATCH_TOKENS_PER_RANK",
|
||||||
@@ -132,7 +132,7 @@ class FlashinferDispatcher(BaseDispatcher):
|
|||||||
|
|
||||||
# Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized.
|
# Calculate workspace size. For eagle mode, use the larger workspace size since nextn layer will be unquantized.
|
||||||
speculative_algo = SpeculativeAlgorithm.from_string(
|
speculative_algo = SpeculativeAlgorithm.from_string(
|
||||||
get_server_args().speculative_algorithm
|
get_spec().speculative_algorithm
|
||||||
)
|
)
|
||||||
if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle():
|
if MOE_NVFP4_DISPATCH and not speculative_algo.is_eagle():
|
||||||
total_dispatch_payload_size_per_token = (
|
total_dispatch_payload_size_per_token = (
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ from typing import (
|
|||||||
import torch
|
import torch
|
||||||
import torch.nn.functional as F
|
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:
|
try:
|
||||||
from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx
|
from triton_kernels.matmul_ogs import GatherIndx, RoutingData, ScatterIndx
|
||||||
@@ -430,10 +430,9 @@ class TopK(MultiPlatformOp):
|
|||||||
assert num_expert_group is not None and topk_group is not None
|
assert num_expert_group is not None and topk_group is not None
|
||||||
|
|
||||||
self.layer_id = layer_id
|
self.layer_id = layer_id
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
self.enable_waterfill = (
|
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
|
self.waterfill_balancer = None
|
||||||
@@ -507,9 +506,8 @@ class TopK(MultiPlatformOp):
|
|||||||
# ===== TO BE REFACTORED ====
|
# ===== TO BE REFACTORED ====
|
||||||
elif get_moe_runner_backend().is_experimental_sgl_trtllm():
|
elif get_moe_runner_backend().is_experimental_sgl_trtllm():
|
||||||
try:
|
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:
|
except ValueError:
|
||||||
use_standard_for_lora = False
|
use_standard_for_lora = False
|
||||||
output_format = (
|
output_format = (
|
||||||
@@ -1362,9 +1360,9 @@ def _eplb_remap_enabled() -> bool:
|
|||||||
# there is no EPLB mapping, so the remap must be skipped.
|
# there is no EPLB mapping, so the remap must be skipped.
|
||||||
return False
|
return False
|
||||||
return (
|
return (
|
||||||
server_args.enable_eplb
|
get_exec().moe.enable_eplb
|
||||||
or server_args.init_expert_location != "trivial"
|
or get_exec().moe.init_expert_location != "trivial"
|
||||||
or server_args.ep_num_redundant_experts > 0
|
or get_exec().moe.ep_num_redundant_experts > 0
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_forward, get_parallel
|
from sglang.srt.runtime_context import get_exec, get_flags, get_forward, get_parallel
|
||||||
from sglang.srt.utils import is_cuda, is_npu
|
from sglang.srt.utils import is_cuda, is_npu
|
||||||
|
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
@@ -239,8 +239,8 @@ def get_deepep_output_dtype(self) -> DispatcherOutputDtype:
|
|||||||
|
|
||||||
# 0. Parse server argument.
|
# 0. Parse server argument.
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
if server_args and server_args.deepep_dispatcher_output_dtype != "auto":
|
if server_args and get_exec().moe.deepep_dispatcher_output_dtype != "auto":
|
||||||
return DispatcherOutputDtype(server_args.deepep_dispatcher_output_dtype)
|
return DispatcherOutputDtype(get_exec().moe.deepep_dispatcher_output_dtype)
|
||||||
|
|
||||||
# 1. Parse deprecated environment variables.
|
# 1. Parse deprecated environment variables.
|
||||||
if envs.SGLANG_DEEPEP_BF16_DISPATCH.get():
|
if envs.SGLANG_DEEPEP_BF16_DISPATCH.get():
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers import deep_gemm_wrapper
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
|
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
|
from sglang.srt.utils.common import torch_release
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -34,7 +34,6 @@ from sglang.kernels.ops.quantization.fp8_kernel import (
|
|||||||
w8a8_block_fp8_matmul_deepgemm,
|
w8a8_block_fp8_matmul_deepgemm,
|
||||||
w8a8_block_fp8_matmul_triton,
|
w8a8_block_fp8_matmul_triton,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
ceil_align,
|
ceil_align,
|
||||||
ceil_div,
|
ceil_div,
|
||||||
@@ -1844,7 +1843,7 @@ def apply_fp8_linear(
|
|||||||
if (
|
if (
|
||||||
input_scale is not None
|
input_scale is not None
|
||||||
and input_scale.numel() == 1
|
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 = (
|
qinput = (
|
||||||
(input_2d * input_scale.reciprocal())
|
(input_2d * input_scale.reciprocal())
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ from sglang.srt.layers.quantization.base_config import (
|
|||||||
QuantizeMethodBase,
|
QuantizeMethodBase,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.utils import is_layer_skipped
|
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 (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
is_cpu,
|
is_cpu,
|
||||||
@@ -333,7 +333,7 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4()
|
self.use_flashinfer = get_moe_runner_backend().is_flashinfer_mxfp4()
|
||||||
self.use_marlin = get_moe_runner_backend().is_marlin()
|
self.use_marlin = get_moe_runner_backend().is_marlin()
|
||||||
self.flashinfer_mxfp4_moe_precision = (
|
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 three FlashInfer
|
# When `flashinfer_mxfp4` is enabled, dispatch to one of three FlashInfer
|
||||||
# entry points depending on the GPU:
|
# 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.dp_attention import is_allocation_symmetric
|
||||||
from sglang.srt.layers.moe.utils import RoutingMethodType
|
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 (
|
from sglang.srt.utils import (
|
||||||
is_flashinfer_available,
|
is_flashinfer_available,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
@@ -51,7 +51,7 @@ class Mxfp4FlashinferTrtllmMoEMethod:
|
|||||||
self._fp8 = fp8_method
|
self._fp8 = fp8_method
|
||||||
self.prefix = prefix
|
self.prefix = prefix
|
||||||
self.flashinfer_mxfp4_moe_precision = (
|
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):
|
def create_moe_runner(self, layer, moe_runner_config):
|
||||||
|
|||||||
@@ -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.rotary_embedding.utils import apply_rotary_emb
|
||||||
from sglang.srt.layers.utils import MultiPlatformOp
|
from sglang.srt.layers.utils import MultiPlatformOp
|
||||||
from sglang.srt.platforms import current_platform
|
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 (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
get_bool_env_var,
|
get_bool_env_var,
|
||||||
@@ -129,7 +129,7 @@ class RotaryEmbedding(MultiPlatformOp):
|
|||||||
self._apply_rotary_emb_wrapped = apply_rotary_emb
|
self._apply_rotary_emb_wrapped = apply_rotary_emb
|
||||||
|
|
||||||
# XXX (MUSA): Implement sgl_kernel.rotary_embedding support for MUSA backend
|
# 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._forward_method = self.forward_native
|
||||||
self._apply_rotary_emb_wrapped = torch.compile(
|
self._apply_rotary_emb_wrapped = torch.compile(
|
||||||
dynamic=True,
|
dynamic=True,
|
||||||
@@ -153,7 +153,7 @@ class RotaryEmbedding(MultiPlatformOp):
|
|||||||
# create the cache on GPU for faster initialization. This may cause
|
# create the cache on GPU for faster initialization. This may cause
|
||||||
# a slight numerical difference between the HF implementation and ours.
|
# a slight numerical difference between the HF implementation and ours.
|
||||||
init_device = (
|
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 / (
|
inv_freq = 1.0 / (
|
||||||
base
|
base
|
||||||
@@ -164,7 +164,7 @@ class RotaryEmbedding(MultiPlatformOp):
|
|||||||
/ self.rotary_dim
|
/ 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()
|
inv_freq = inv_freq.cuda()
|
||||||
return inv_freq
|
return inv_freq
|
||||||
|
|
||||||
|
|||||||
@@ -18,7 +18,7 @@ from sglang.srt.layers.rotary_embedding.yarn import (
|
|||||||
yarn_get_mscale_simple,
|
yarn_get_mscale_simple,
|
||||||
yarn_linear_ramp_mask,
|
yarn_linear_ramp_mask,
|
||||||
)
|
)
|
||||||
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 (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
is_cuda,
|
is_cuda,
|
||||||
@@ -132,7 +132,7 @@ class MRotaryEmbedding(RotaryEmbedding):
|
|||||||
self.register_buffer("axis_map", axis_map, persistent=False)
|
self.register_buffer("axis_map", axis_map, persistent=False)
|
||||||
else:
|
else:
|
||||||
self.axis_map = None
|
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
|
self._forward_method = self.forward_native
|
||||||
|
|
||||||
def get_cos_sin_with_position(self, positions):
|
def get_cos_sin_with_position(self, positions):
|
||||||
|
|||||||
@@ -15,7 +15,7 @@ from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
|||||||
from sglang.srt.layers.logprob_processor import (
|
from sglang.srt.layers.logprob_processor import (
|
||||||
OutputLogprobProcessor,
|
OutputLogprobProcessor,
|
||||||
)
|
)
|
||||||
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.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||||
from sglang.srt.utils.async_probe import sanitize_nan_logits
|
from sglang.srt.utils.async_probe import sanitize_nan_logits
|
||||||
@@ -74,12 +74,14 @@ class Sampler(nn.Module):
|
|||||||
if is_dp_attention_enabled():
|
if is_dp_attention_enabled():
|
||||||
self.tp_sync_group = get_parallel().attn_tp_group.device_group
|
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.
|
# 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.
|
# 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_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()
|
self.output_logprob_processor = OutputLogprobProcessor()
|
||||||
|
|
||||||
@@ -260,7 +262,7 @@ class Sampler(nn.Module):
|
|||||||
positions=positions,
|
positions=positions,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
backend = get_server_args().sampling_backend
|
backend = get_exec().kernel.sampling_backend
|
||||||
if backend == "flashinfer":
|
if backend == "flashinfer":
|
||||||
assert (
|
assert (
|
||||||
sampling_info.sampling_seed is None
|
sampling_info.sampling_seed is None
|
||||||
@@ -540,7 +542,7 @@ def create_sampler(backend: Optional[str] = None) -> "Sampler":
|
|||||||
"""Create a sampler honoring custom backend registrations."""
|
"""Create a sampler honoring custom backend registrations."""
|
||||||
|
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
backend = backend or (server_args.sampling_backend if server_args else None)
|
backend = backend or (get_exec().kernel.sampling_backend if server_args else None)
|
||||||
|
|
||||||
if backend in _CUSTOM_SAMPLER_FACTORIES:
|
if backend in _CUSTOM_SAMPLER_FACTORIES:
|
||||||
sampler = _CUSTOM_SAMPLER_FACTORIES[backend]()
|
sampler = _CUSTOM_SAMPLER_FACTORIES[backend]()
|
||||||
|
|||||||
@@ -48,7 +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.cpu_monitor import start_cpu_monitor_thread
|
||||||
from sglang.srt.observability.req_time_stats import DPControllerReqTimeStats
|
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.observability.trace import process_tracing_init, trace_set_thread_info
|
||||||
from sglang.srt.runtime_context import publish
|
from sglang.srt.runtime_context import get_exec, publish
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
DP_ATTENTION_HANDSHAKE_PORT_DELTA,
|
||||||
PortArgs,
|
PortArgs,
|
||||||
@@ -232,7 +232,7 @@ class DataParallelController:
|
|||||||
sock_send(worker, obj)
|
sock_send(worker, obj)
|
||||||
|
|
||||||
def update_active_ranks(self, ranks: ActiveRanksOutput):
|
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:
|
if len(ranks.status) != self.max_dp_size:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; "
|
"[Elastic EP][DPC] active rank status len=%d != max_dp_size=%d; "
|
||||||
@@ -485,7 +485,7 @@ class DataParallelController:
|
|||||||
logger.debug("Worker port broadcast completed")
|
logger.debug("Worker port broadcast completed")
|
||||||
return worker_ports
|
return worker_ports
|
||||||
finally:
|
finally:
|
||||||
if self.server_args.elastic_ep_backend is None:
|
if get_exec().moe.elastic_ep_backend is None:
|
||||||
rep_socket.close()
|
rep_socket.close()
|
||||||
else:
|
else:
|
||||||
threading.Thread(
|
threading.Thread(
|
||||||
|
|||||||
@@ -33,7 +33,12 @@ from sglang.srt.managers.schedule_batch import (
|
|||||||
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
|
from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalStaticCache
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.multimodal.evs import EVSEmbeddingResult
|
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,
|
||||||
|
)
|
||||||
from sglang.srt.utils import flatten_nested_list, is_hip, is_npu, print_warning_once
|
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.srt.utils.stale_shm_cleanup import make_shm_name
|
||||||
from sglang.utils import logger
|
from sglang.utils import logger
|
||||||
@@ -931,7 +936,7 @@ def _adjust_embedding_length(
|
|||||||
f"tokens from multimodal embeddings."
|
f"tokens from multimodal embeddings."
|
||||||
)
|
)
|
||||||
if num_mm_tokens_in_input_ids < num_mm_tokens_in_embedding:
|
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:
|
if chunked_prefill_size != -1:
|
||||||
logger.warning(
|
logger.warning(
|
||||||
"You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill"
|
"You may want to avoid this issue by raising `chunked_prefill_size`, or disabling chunked prefill"
|
||||||
@@ -1295,7 +1300,7 @@ def general_mm_embed_routine(
|
|||||||
# encoder/ViT execution and multimodal feature placement, while
|
# encoder/ViT execution and multimodal feature placement, while
|
||||||
# the language model range below excludes both.
|
# the language model range below excludes both.
|
||||||
with torch.profiler.record_function("sglang.vlm.mm_embedding"):
|
with torch.profiler.record_function("sglang.vlm.mm_embedding"):
|
||||||
if server_args and server_args.enable_adaptive_dispatch_to_encoder:
|
if server_args and get_disagg().enable_adaptive_dispatch_to_encoder:
|
||||||
# Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches
|
# Split by precomputed vs non-precomputed so get_embedding_and_mask only sees uniform batches
|
||||||
input_embeds, other_info = _embed_mm_inputs_with_split(
|
input_embeds, other_info = _embed_mm_inputs_with_split(
|
||||||
mm_inputs_list=mm_inputs_list,
|
mm_inputs_list=mm_inputs_list,
|
||||||
@@ -1340,7 +1345,7 @@ def general_mm_embed_routine(
|
|||||||
feature = getattr(mm_item, "feature", None)
|
feature = getattr(mm_item, "feature", None)
|
||||||
if isinstance(feature, torch.Tensor) and feature.is_cuda:
|
if isinstance(feature, torch.Tensor) and feature.is_cuda:
|
||||||
mm_item.feature = feature.to("cpu", non_blocking=True)
|
mm_item.feature = feature.to("cpu", non_blocking=True)
|
||||||
if get_server_args().language_only:
|
if get_disagg().language_only:
|
||||||
precomputed_embeddings = getattr(
|
precomputed_embeddings = getattr(
|
||||||
mm_item, "precomputed_embeddings", None
|
mm_item, "precomputed_embeddings", None
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from sglang.srt.dllm.config import DllmConfig
|
from sglang.srt.dllm.config import DllmConfig
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
|
from sglang.srt.runtime_context import get_exec, get_schedule, get_serving, get_spec
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
Range,
|
Range,
|
||||||
ceil_align,
|
ceil_align,
|
||||||
@@ -1097,7 +1098,7 @@ class Req(ReqDllmMixin):
|
|||||||
"""Check if this request is prefill-only (no token generation needed)."""
|
"""Check if this request is prefill-only (no token generation needed)."""
|
||||||
# NOTE: when spec is enabled, prefill_only optimizations are disabled
|
# 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
|
return self.sampling_params.max_new_tokens == 0 and spec_alg is None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -1118,7 +1119,7 @@ class Req(ReqDllmMixin):
|
|||||||
def effective_kv_committed_len(self) -> int:
|
def effective_kv_committed_len(self) -> int:
|
||||||
# Report only the prompt prefix so thinking + answer fall into the
|
# Report only the prompt prefix so thinking + answer fall into the
|
||||||
# overallocated range and are reclaimed by release_kv_cache. #22373.
|
# 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 min(self.kv_committed_len, len(self.origin_input_ids))
|
||||||
return self.kv_committed_len
|
return self.kv_committed_len
|
||||||
|
|
||||||
@@ -2922,7 +2923,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if server_args.enable_mamba_extra_buffer():
|
if server_args.enable_mamba_extra_buffer():
|
||||||
mamba_track_interval = server_args.mamba_track_interval
|
mamba_track_interval = get_exec().mamba.mamba_track_interval
|
||||||
|
|
||||||
if len(self.reqs) == 0:
|
if len(self.reqs) == 0:
|
||||||
self.mamba_track_indices = torch.empty(
|
self.mamba_track_indices = torch.empty(
|
||||||
@@ -3168,8 +3169,8 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
continue
|
continue
|
||||||
else:
|
else:
|
||||||
pre_len = (
|
pre_len = (
|
||||||
pre_len - server_args.chunked_prefill_size
|
pre_len - get_schedule().chunked_prefill_size
|
||||||
if server_args.chunked_prefill_size > 0
|
if get_schedule().chunked_prefill_size > 0
|
||||||
else pre_len
|
else pre_len
|
||||||
)
|
)
|
||||||
self._evict_swa(req, pre_len)
|
self._evict_swa(req, pre_len)
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from array import array
|
|||||||
|
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor
|
from sglang.srt.managers.prefill_delayer import PrefillDelayerSinglePassExecutor
|
||||||
|
from sglang.srt.runtime_context import get_disagg
|
||||||
from sglang.srt.utils import get_bool_env_var
|
from sglang.srt.utils import get_bool_env_var
|
||||||
|
|
||||||
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
|
_ROUTING_KEY_POLICY_DEBUG_LOG = get_bool_env_var("SGLANG_ROUTING_KEY_POLICY_DEBUG_LOG")
|
||||||
@@ -56,7 +57,6 @@ from sglang.srt.mem_cache.multi_ended_allocator import (
|
|||||||
UnifiedMambaTokenToKVPoolAllocator,
|
UnifiedMambaTokenToKVPoolAllocator,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -195,7 +195,7 @@ class SchedulePolicy:
|
|||||||
if (
|
if (
|
||||||
not isinstance(policy, CacheAwarePolicy)
|
not isinstance(policy, CacheAwarePolicy)
|
||||||
and self.tree_cache.supports_fast_match_prefix()
|
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:
|
for r in waiting_queue:
|
||||||
match_prefix_for_req(self.tree_cache, r, include_req=True)
|
match_prefix_for_req(self.tree_cache, r, include_req=True)
|
||||||
|
|||||||
@@ -27,6 +27,20 @@ from functools import partial
|
|||||||
from http import HTTPStatus
|
from http import HTTPStatus
|
||||||
from typing import Any, Deque, Dict, List, Optional, Tuple, Union
|
from typing import Any, Deque, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
get_device,
|
||||||
|
get_disagg,
|
||||||
|
get_exec,
|
||||||
|
get_lora,
|
||||||
|
get_memory,
|
||||||
|
get_mm,
|
||||||
|
get_model,
|
||||||
|
get_observability,
|
||||||
|
get_schedule,
|
||||||
|
get_serving,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
|
|
||||||
from sglang.srt.utils.common import suppress_noisy_warnings # isort: skip
|
from sglang.srt.utils.common import suppress_noisy_warnings # isort: skip
|
||||||
|
|
||||||
suppress_noisy_warnings()
|
suppress_noisy_warnings()
|
||||||
@@ -482,9 +496,9 @@ class Scheduler(
|
|||||||
attn_tp_cpu_group=self.attn_tp_cpu_group,
|
attn_tp_cpu_group=self.attn_tp_cpu_group,
|
||||||
tp_cpu_group=self.tp_cpu_group,
|
tp_cpu_group=self.tp_cpu_group,
|
||||||
attn_cp_cpu_group=self.attn_cp_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(
|
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.pp_rank == 0
|
||||||
and self.ps.attn_tp_rank == 0
|
and self.ps.attn_tp_rank == 0
|
||||||
and self.ps.attn_cp_rank == 0
|
and self.ps.attn_cp_rank == 0
|
||||||
@@ -526,8 +540,8 @@ class Scheduler(
|
|||||||
self.init_hisparse_coordinator()
|
self.init_hisparse_coordinator()
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.server_args.disaggregation_mode == "decode"
|
get_disagg().disaggregation_mode == "decode"
|
||||||
and self.server_args.disaggregation_decode_enable_offload_kvcache
|
and get_disagg().disaggregation_decode_enable_offload_kvcache
|
||||||
):
|
):
|
||||||
self.decode_offload_manager = DecodeKVCacheOffloadManager(
|
self.decode_offload_manager = DecodeKVCacheOffloadManager(
|
||||||
req_to_token_pool=self.req_to_token_pool,
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
@@ -642,7 +656,7 @@ class Scheduler(
|
|||||||
|
|
||||||
self.dllm_config = ( # For diffusion LLM
|
self.dllm_config = ( # For diffusion LLM
|
||||||
DllmConfig.from_server_args(self.server_args)
|
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
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -671,10 +685,10 @@ class Scheduler(
|
|||||||
port_args=port_args,
|
port_args=port_args,
|
||||||
is_rank_zero=is_rank_zero,
|
is_rank_zero=is_rank_zero,
|
||||||
skip_tokenizer_init=self.server_args.skip_tokenizer_init,
|
skip_tokenizer_init=self.server_args.skip_tokenizer_init,
|
||||||
metrics_enabled=self.server_args.enable_metrics
|
metrics_enabled=get_observability().enable_metrics
|
||||||
and (
|
and (
|
||||||
self.ps.attn_tp_rank == 0
|
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(),
|
enable_scripted_runtime=envs.SGLANG_TEST_SCRIPTED_RUNTIME.get(),
|
||||||
)
|
)
|
||||||
@@ -693,7 +707,7 @@ class Scheduler(
|
|||||||
port_args,
|
port_args,
|
||||||
self.ps.dp_size,
|
self.ps.dp_size,
|
||||||
dp_rank,
|
dp_rank,
|
||||||
publish_interval=self.server_args.load_snapshot_publish_interval,
|
publish_interval=get_observability().load_snapshot_publish_interval,
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning("load snapshot writer init failed: %s", e)
|
logger.warning("load snapshot writer init failed: %s", e)
|
||||||
@@ -703,7 +717,7 @@ class Scheduler(
|
|||||||
self.ps.pp_rank == 0
|
self.ps.pp_rank == 0
|
||||||
and self.ps.attn_tp_rank == 0
|
and self.ps.attn_tp_rank == 0
|
||||||
and self.ps.attn_cp_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(
|
self.idle_sleeper = IdleSleeper(
|
||||||
sockets=[
|
sockets=[
|
||||||
@@ -737,22 +751,22 @@ class Scheduler(
|
|||||||
else:
|
else:
|
||||||
if self.model_config.is_multimodal:
|
if self.model_config.is_multimodal:
|
||||||
self.processor = get_processor(
|
self.processor = get_processor(
|
||||||
server_args.tokenizer_path,
|
get_serving().tokenizer_path,
|
||||||
tokenizer_mode=server_args.tokenizer_mode,
|
tokenizer_mode=get_serving().tokenizer_mode,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=get_model().trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=get_model().revision,
|
||||||
use_fast=not server_args.disable_fast_image_processor,
|
use_fast=not get_mm().disable_fast_image_processor,
|
||||||
tokenizer_backend=server_args.tokenizer_backend,
|
tokenizer_backend=get_serving().tokenizer_backend,
|
||||||
model_name=server_args.model_path,
|
model_name=get_model().model_path,
|
||||||
)
|
)
|
||||||
self.tokenizer = get_tokenizer_from_processor(self.processor)
|
self.tokenizer = get_tokenizer_from_processor(self.processor)
|
||||||
else:
|
else:
|
||||||
self.tokenizer = get_tokenizer(
|
self.tokenizer = get_tokenizer(
|
||||||
server_args.tokenizer_path,
|
get_serving().tokenizer_path,
|
||||||
tokenizer_mode=server_args.tokenizer_mode,
|
tokenizer_mode=get_serving().tokenizer_mode,
|
||||||
trust_remote_code=server_args.trust_remote_code,
|
trust_remote_code=get_model().trust_remote_code,
|
||||||
revision=server_args.revision,
|
revision=get_model().revision,
|
||||||
tokenizer_backend=server_args.tokenizer_backend,
|
tokenizer_backend=get_serving().tokenizer_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Load multimodal processor for M-RoPE fallback computation.
|
# Load multimodal processor for M-RoPE fallback computation.
|
||||||
@@ -774,9 +788,9 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Set reasoning_parser and think_end_id if --reasoning_parser is enabled
|
# 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(
|
reasoning_parser = ReasoningParser(
|
||||||
model_type=self.server_args.reasoning_parser,
|
model_type=get_serving().reasoning_parser,
|
||||||
stream_reasoning=False,
|
stream_reasoning=False,
|
||||||
tokenizer=self.tokenizer,
|
tokenizer=self.tokenizer,
|
||||||
)
|
)
|
||||||
@@ -847,7 +861,7 @@ class Scheduler(
|
|||||||
target_worker=self.tp_worker,
|
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):
|
# 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
|
# 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
|
# build_load_config reads server_args.load_format, so a bag-only
|
||||||
@@ -855,10 +869,10 @@ class Scheduler(
|
|||||||
# format.
|
# format.
|
||||||
self.server_args.override(
|
self.server_args.override(
|
||||||
"scheduler.draft_load_format",
|
"scheduler.draft_load_format",
|
||||||
load_format=self.server_args.speculative_draft_load_format,
|
load_format=get_spec().speculative_draft_load_format,
|
||||||
)
|
)
|
||||||
logger.info(
|
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)
|
DraftWorkerClass = self.spec_algorithm.create_worker(self.server_args)
|
||||||
@@ -925,8 +939,8 @@ class Scheduler(
|
|||||||
model_runner.post_capture_resize_kv_pool()
|
model_runner.post_capture_resize_kv_pool()
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.server_args.elastic_ep_backend is not None
|
get_exec().moe.elastic_ep_backend is not None
|
||||||
and self.server_args.ep_join_mode == "recover"
|
and get_exec().moe.ep_join_mode == "recover"
|
||||||
):
|
):
|
||||||
model_runner.post_capture_elastic_ep_recover()
|
model_runner.post_capture_elastic_ep_recover()
|
||||||
|
|
||||||
@@ -955,7 +969,7 @@ class Scheduler(
|
|||||||
# --min-free-slots-delay. Built independently of the prefill delayer.
|
# --min-free-slots-delay. Built independently of the prefill delayer.
|
||||||
self.min_free_slots_delayer: Optional[MinFreeSlotsDelayer] = None
|
self.min_free_slots_delayer: Optional[MinFreeSlotsDelayer] = None
|
||||||
min_free_slots = resolve_min_free_slots(
|
min_free_slots = resolve_min_free_slots(
|
||||||
self.server_args.min_free_slots_delay,
|
get_schedule().min_free_slots_delay,
|
||||||
self.max_running_requests,
|
self.max_running_requests,
|
||||||
is_dflash_family=self.spec_algorithm.is_dflash_family(),
|
is_dflash_family=self.spec_algorithm.is_dflash_family(),
|
||||||
)
|
)
|
||||||
@@ -1001,14 +1015,14 @@ class Scheduler(
|
|||||||
if self.ps.tp_rank == 0:
|
if self.ps.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
f"max_total_num_tokens={self.max_total_num_tokens}, "
|
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_prefill_tokens={self.max_prefill_tokens}, "
|
||||||
f"max_running_requests={self.max_running_requests}, "
|
f"max_running_requests={self.max_running_requests}, "
|
||||||
f"context_len={self.model_config.context_len}, "
|
f"context_len={self.model_config.context_len}, "
|
||||||
f"{'available_cpu_mem' if self.device == 'cpu' else 'available_gpu_mem'}={avail_mem:.2f} GB"
|
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(
|
self.metrics_collector.emit_constants(
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
# TODO: max_running_requests_under_SLO has no setter — dead chain.
|
# TODO: max_running_requests_under_SLO has no setter — dead chain.
|
||||||
@@ -1055,7 +1069,7 @@ class Scheduler(
|
|||||||
self._engine_paused = False
|
self._engine_paused = False
|
||||||
|
|
||||||
def init_chunked_prefill(self):
|
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 = (
|
uses_transformers_backend = (
|
||||||
get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS
|
get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS
|
||||||
)
|
)
|
||||||
@@ -1075,13 +1089,12 @@ class Scheduler(
|
|||||||
self.chunked_req = None
|
self.chunked_req = None
|
||||||
self._pending_chunked_abort_req = None
|
self._pending_chunked_abort_req = None
|
||||||
self.is_mixed_chunk = (
|
self.is_mixed_chunk = (
|
||||||
self.chunked_prefill_size is not None
|
self.chunked_prefill_size is not None and get_schedule().enable_mixed_chunk
|
||||||
and self.server_args.enable_mixed_chunk
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# Init the dynamic chunking predictor for PP
|
# Init the dynamic chunking predictor for PP
|
||||||
self.enable_dynamic_chunking = (
|
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:
|
if self.enable_dynamic_chunking:
|
||||||
try:
|
try:
|
||||||
@@ -1117,8 +1130,8 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
self.prefill_delayer: Optional[PrefillDelayer] = None
|
self.prefill_delayer: Optional[PrefillDelayer] = None
|
||||||
self.max_prefill_bs: int = 0
|
self.max_prefill_bs: int = 0
|
||||||
if self.server_args.enable_prefill_delayer:
|
if get_schedule().enable_prefill_delayer:
|
||||||
if self.server_args.disaggregation_mode == "decode":
|
if get_disagg().disaggregation_mode == "decode":
|
||||||
logger.info(
|
logger.info(
|
||||||
"Ignoring --enable-prefill-delayer on decode engine "
|
"Ignoring --enable-prefill-delayer on decode engine "
|
||||||
"(no prefill scheduling path; delayer would be a no-op)."
|
"(no prefill scheduling path; delayer would be a no-op)."
|
||||||
@@ -1135,15 +1148,15 @@ class Scheduler(
|
|||||||
if self.metrics_reporter.enable_metrics
|
if self.metrics_reporter.enable_metrics
|
||||||
else None
|
else None
|
||||||
),
|
),
|
||||||
max_delay_passes=self.server_args.prefill_delayer_max_delay_passes,
|
max_delay_passes=get_schedule().prefill_delayer_max_delay_passes,
|
||||||
token_usage_low_watermark=self.server_args.prefill_delayer_token_usage_low_watermark,
|
token_usage_low_watermark=get_schedule().prefill_delayer_token_usage_low_watermark,
|
||||||
device=self.tp_group.device,
|
device=self.tp_group.device,
|
||||||
)
|
)
|
||||||
|
|
||||||
# NOTE: preemption is enabled by default for priority scheduling.
|
# NOTE: preemption is enabled by default for priority scheduling.
|
||||||
self.enable_priority_preemption = (
|
self.enable_priority_preemption = (
|
||||||
self.enable_priority_scheduling
|
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(
|
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
|
||||||
@@ -1159,12 +1172,12 @@ class Scheduler(
|
|||||||
def init_watch_dog_memory_saver_input_blocker(self):
|
def init_watch_dog_memory_saver_input_blocker(self):
|
||||||
# Start watchdog thread
|
# Start watchdog thread
|
||||||
self.watchdog = create_scheduler_watchdog(
|
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
|
# Init memory saver, profiler and metric stats
|
||||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
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
|
# Init recv skipper and input blocker
|
||||||
@@ -1186,11 +1199,9 @@ class Scheduler(
|
|||||||
self.disagg_decode_prealloc_queue = None
|
self.disagg_decode_prealloc_queue = None
|
||||||
self.disagg_decode_transfer_queue = None
|
self.disagg_decode_transfer_queue = None
|
||||||
|
|
||||||
self.disaggregation_mode = DisaggregationMode(
|
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
|
||||||
self.server_args.disaggregation_mode
|
|
||||||
)
|
|
||||||
self.transfer_backend = TransferBackend(
|
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?
|
# 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?
|
||||||
@@ -1260,10 +1271,10 @@ class Scheduler(
|
|||||||
tp_size=self.ps.tp_size,
|
tp_size=self.ps.tp_size,
|
||||||
dp_size=self.server_args.dp_size,
|
dp_size=self.server_args.dp_size,
|
||||||
gpu_id=self.ps.gpu_id,
|
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,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
pp_rank=self.ps.pp_rank,
|
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,
|
transfer_backend=self.transfer_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1289,7 +1300,7 @@ class Scheduler(
|
|||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
tp_size=self.ps.tp_size,
|
tp_size=self.ps.tp_size,
|
||||||
gpu_id=self.ps.gpu_id,
|
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,
|
gloo_group=self.attn_tp_cpu_group,
|
||||||
max_total_num_tokens=self.max_total_num_tokens,
|
max_total_num_tokens=self.max_total_num_tokens,
|
||||||
scheduler=self,
|
scheduler=self,
|
||||||
@@ -1303,11 +1314,10 @@ class Scheduler(
|
|||||||
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
self.enable_staging = envs.SGLANG_DISAGG_STAGING_BUFFER.get()
|
||||||
|
|
||||||
# Init mm receiver for EPD disaggregation mode
|
# Init mm receiver for EPD disaggregation mode
|
||||||
if (
|
if get_disagg().language_only and get_disagg().encoder_transfer_backend in [
|
||||||
self.server_args.language_only
|
"zmq_to_scheduler",
|
||||||
and self.server_args.encoder_transfer_backend
|
"mooncake",
|
||||||
in ["zmq_to_scheduler", "mooncake"]
|
]:
|
||||||
):
|
|
||||||
self.mm_receiver = create_mm_receiver(
|
self.mm_receiver = create_mm_receiver(
|
||||||
self.server_args,
|
self.server_args,
|
||||||
dtype=self.model_config.dtype,
|
dtype=self.model_config.dtype,
|
||||||
@@ -1388,7 +1398,7 @@ class Scheduler(
|
|||||||
|
|
||||||
def init_deterministic_inference_config(self):
|
def init_deterministic_inference_config(self):
|
||||||
"""Initialize deterministic inference configuration for different attention backends."""
|
"""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
|
self.truncation_align_size = None
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -1794,10 +1804,10 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
|
|
||||||
def init_lora_drainer(self) -> None:
|
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.lora_drainer = LoRADrainer(
|
||||||
self.server_args.max_loras_per_batch,
|
get_lora().max_loras_per_batch,
|
||||||
self.server_args.lora_drain_wait_threshold,
|
get_lora().lora_drain_wait_threshold,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.lora_drainer = None
|
self.lora_drainer = None
|
||||||
@@ -1923,7 +1933,7 @@ class Scheduler(
|
|||||||
|
|
||||||
def init_kv_events_publisher(self) -> None:
|
def init_kv_events_publisher(self) -> None:
|
||||||
self.kv_events_publisher = SchedulerKvEventsPublisher(
|
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,
|
ps=self.ps,
|
||||||
attn_tp_rank=self.ps.attn_tp_rank,
|
attn_tp_rank=self.ps.attn_tp_rank,
|
||||||
attn_cp_rank=self.ps.attn_cp_rank,
|
attn_cp_rank=self.ps.attn_cp_rank,
|
||||||
@@ -2107,7 +2117,7 @@ class Scheduler(
|
|||||||
return image_inputs
|
return image_inputs
|
||||||
|
|
||||||
def _get_multimodal_inputs(self, mm_inputs_dict):
|
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)
|
return self._process_and_broadcast_mm_inputs(mm_inputs_dict)
|
||||||
else:
|
else:
|
||||||
return MultimodalInputs.from_processor_output(mm_inputs_dict)
|
return MultimodalInputs.from_processor_output(mm_inputs_dict)
|
||||||
@@ -2154,7 +2164,7 @@ class Scheduler(
|
|||||||
|
|
||||||
def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None:
|
def _maybe_namespace_elastic_radix_cache(self, req: Req) -> None:
|
||||||
if (
|
if (
|
||||||
self.server_args.elastic_ep_backend is None
|
get_exec().moe.elastic_ep_backend is None
|
||||||
or self.disable_radix_cache
|
or self.disable_radix_cache
|
||||||
or not self.tree_cache.is_tree_cache()
|
or not self.tree_cache.is_tree_cache()
|
||||||
):
|
):
|
||||||
@@ -2200,8 +2210,7 @@ class Scheduler(
|
|||||||
)
|
)
|
||||||
# Radix-native sessions use only the top-level session_id.
|
# Radix-native sessions use only the top-level session_id.
|
||||||
radix_native_session = (
|
radix_native_session = (
|
||||||
recv_req.session_id is not None
|
recv_req.session_id is not None and get_memory().enable_session_radix_cache
|
||||||
and self.server_args.enable_session_radix_cache
|
|
||||||
)
|
)
|
||||||
|
|
||||||
if session_id is None or radix_native_session:
|
if session_id is None or radix_native_session:
|
||||||
@@ -2213,7 +2222,7 @@ class Scheduler(
|
|||||||
|
|
||||||
if recv_req.bootstrap_port is None:
|
if recv_req.bootstrap_port is None:
|
||||||
# Use default bootstrap port
|
# 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(
|
req = Req(
|
||||||
recv_req.rid,
|
recv_req.rid,
|
||||||
@@ -2366,7 +2375,7 @@ class Scheduler(
|
|||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
return
|
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
|
# The ascend backend samples from logits directly and never builds the
|
||||||
# top-k/top-p support, so it cannot produce a sampling mask.
|
# top-k/top-p support, so it cannot produce a sampling mask.
|
||||||
error_msg = (
|
error_msg = (
|
||||||
@@ -2415,7 +2424,7 @@ class Scheduler(
|
|||||||
error_msg = validate_input_length(
|
error_msg = validate_input_length(
|
||||||
req,
|
req,
|
||||||
self.max_req_input_len,
|
self.max_req_input_len,
|
||||||
self.server_args.allow_auto_truncate,
|
get_serving().allow_auto_truncate,
|
||||||
)
|
)
|
||||||
if error_msg:
|
if error_msg:
|
||||||
req.set_finish_with_abort(error_msg)
|
req.set_finish_with_abort(error_msg)
|
||||||
@@ -2693,7 +2702,7 @@ class Scheduler(
|
|||||||
error_msg = validate_input_length(
|
error_msg = validate_input_length(
|
||||||
req,
|
req,
|
||||||
self.max_req_input_len,
|
self.max_req_input_len,
|
||||||
self.server_args.allow_auto_truncate,
|
get_serving().allow_auto_truncate,
|
||||||
)
|
)
|
||||||
if error_msg:
|
if error_msg:
|
||||||
self._add_request_to_queue(req)
|
self._add_request_to_queue(req)
|
||||||
@@ -2905,7 +2914,7 @@ class Scheduler(
|
|||||||
if (
|
if (
|
||||||
need_mlp_sync
|
need_mlp_sync
|
||||||
and not self.spec_algorithm.is_none()
|
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.
|
# 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:
|
# Before merging the new batch into running batch:
|
||||||
@@ -2979,7 +2988,7 @@ class Scheduler(
|
|||||||
for req in ready_grammar_requests:
|
for req in ready_grammar_requests:
|
||||||
self._add_request_to_queue(req)
|
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()
|
self.tree_cache.check_hicache_events()
|
||||||
|
|
||||||
if self.enable_priority_preemption or self.is_hybrid_swa:
|
if self.enable_priority_preemption or self.is_hybrid_swa:
|
||||||
@@ -3046,7 +3055,7 @@ class Scheduler(
|
|||||||
self.priority_scheduling_preemption_threshold,
|
self.priority_scheduling_preemption_threshold,
|
||||||
max_prefill_bs=self.max_prefill_bs,
|
max_prefill_bs=self.max_prefill_bs,
|
||||||
max_running_requests=self.max_running_requests,
|
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,
|
prefill_delayer_single_pass=prefill_delayer_single_pass,
|
||||||
dllm_config=self.dllm_config,
|
dllm_config=self.dllm_config,
|
||||||
waiting_queue_len=len(self.waiting_queue),
|
waiting_queue_len=len(self.waiting_queue),
|
||||||
@@ -3619,7 +3628,7 @@ class Scheduler(
|
|||||||
|
|
||||||
def _maybe_report_active_ranks(self) -> None:
|
def _maybe_report_active_ranks(self) -> None:
|
||||||
if not (
|
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
|
return
|
||||||
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||||
@@ -3924,7 +3933,7 @@ class Scheduler(
|
|||||||
ok, msg = self.tree_cache.attach_storage_backend(
|
ok, msg = self.tree_cache.attach_storage_backend(
|
||||||
storage_backend=recv_req.hicache_storage_backend,
|
storage_backend=recv_req.hicache_storage_backend,
|
||||||
storage_backend_extra_config_json=recv_req.hicache_storage_backend_extra_config_json,
|
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_storage_prefetch_policy=recv_req.hicache_storage_prefetch_policy,
|
||||||
hicache_write_policy=recv_req.hicache_write_policy,
|
hicache_write_policy=recv_req.hicache_write_policy,
|
||||||
)
|
)
|
||||||
@@ -4044,7 +4053,7 @@ class Scheduler(
|
|||||||
}
|
}
|
||||||
ret["effective_max_running_requests_per_dp"] = self.max_running_requests
|
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
|
from sglang.srt.elastic_ep.elastic_ep import ElasticEPStateManager
|
||||||
|
|
||||||
ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling()
|
ret["is_scaling_elastic_ep"] = ElasticEPStateManager.is_scaling()
|
||||||
@@ -4583,10 +4592,10 @@ class Scheduler(
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
def close_session(self, recv_req: CloseSessionReqInput):
|
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)
|
self.tree_cache.release_radix_session(recv_req.session_id)
|
||||||
if recv_req.session_id in self.session_controller or not (
|
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)
|
self.session_controller.close(recv_req)
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,13 @@ from sglang.srt.mem_cache.common import (
|
|||||||
maybe_cache_unfinished_req,
|
maybe_cache_unfinished_req,
|
||||||
release_kv_cache,
|
release_kv_cache,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_disagg,
|
||||||
|
get_exec,
|
||||||
|
get_memory,
|
||||||
|
get_observability,
|
||||||
|
get_server_args,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
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.indexer_topk import get_global_indexer_capturer
|
||||||
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
|
from sglang.srt.state_capturer.routed_experts import get_global_experts_capturer
|
||||||
@@ -84,7 +90,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
|
|
||||||
def process_batch_result_prebuilt(self, batch: ScheduleBatch):
|
def process_batch_result_prebuilt(self, batch: ScheduleBatch):
|
||||||
assert self.disaggregation_mode == DisaggregationMode.DECODE
|
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:
|
if use_free_group:
|
||||||
self.token_to_kv_pool_allocator.free_group_begin()
|
self.token_to_kv_pool_allocator.free_group_begin()
|
||||||
for req in batch.reqs:
|
for req in batch.reqs:
|
||||||
@@ -92,7 +98,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
req.update_finish_state()
|
req.update_finish_state()
|
||||||
if req.finished():
|
if req.finished():
|
||||||
req.time_stats.set_quick_finish_time()
|
req.time_stats.set_quick_finish_time()
|
||||||
if self.server_args.enable_hisparse:
|
if get_memory().enable_hisparse:
|
||||||
self.hisparse_coordinator.request_finished(req)
|
self.hisparse_coordinator.request_finished(req)
|
||||||
release_kv_cache(req, self.tree_cache)
|
release_kv_cache(req, self.tree_cache)
|
||||||
|
|
||||||
@@ -243,7 +249,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
req.time_stats.set_completion_time()
|
req.time_stats.set_completion_time()
|
||||||
elif not batch.decoding_reqs or req not in batch.decoding_reqs:
|
elif not batch.decoding_reqs or req not in batch.decoding_reqs:
|
||||||
maybe_cache_unfinished_req(req, self.tree_cache)
|
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.hisparse_coordinator.admit_request_into_staging(req)
|
||||||
|
|
||||||
self._maybe_collect_customized_info(i, req, logits_output)
|
self._maybe_collect_customized_info(i, req, logits_output)
|
||||||
@@ -756,7 +762,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
num_block_accept_tokens=result.num_block_accept_tokens,
|
num_block_accept_tokens=result.num_block_accept_tokens,
|
||||||
num_cap_tokens=result.num_cap_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(
|
self.metrics_collector.increment_decode_cuda_graph_pass(
|
||||||
value=can_run_cuda_graph
|
value=can_run_cuda_graph
|
||||||
)
|
)
|
||||||
@@ -939,7 +945,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
self._mamba_prefix_cache_update(req, batch, result, i)
|
self._mamba_prefix_cache_update(req, batch, result, i)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.server_args.disaggregation_decode_enable_offload_kvcache
|
get_disagg().disaggregation_decode_enable_offload_kvcache
|
||||||
and not req.finished()
|
and not req.finished()
|
||||||
):
|
):
|
||||||
self.decode_offload_manager.offload_kv_cache(req)
|
self.decode_offload_manager.offload_kv_cache(req)
|
||||||
@@ -959,12 +965,12 @@ class SchedulerBatchResultProcessor:
|
|||||||
self._maybe_collect_routed_experts(req)
|
self._maybe_collect_routed_experts(req)
|
||||||
self._maybe_collect_indexer_topk(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
|
# 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):
|
if not self.decode_offload_manager.offload_kv_cache(req):
|
||||||
self.decode_offload_manager.finalize_release_on_finish(req)
|
self.decode_offload_manager.finalize_release_on_finish(req)
|
||||||
else:
|
else:
|
||||||
if self.server_args.enable_hisparse:
|
if get_memory().enable_hisparse:
|
||||||
self.hisparse_coordinator.request_finished(req)
|
self.hisparse_coordinator.request_finished(req)
|
||||||
prepare_release = getattr(
|
prepare_release = getattr(
|
||||||
self.model_worker, "prepare_for_kv_cache_release", None
|
self.model_worker, "prepare_for_kv_cache_release", None
|
||||||
@@ -1063,7 +1069,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
other_idx
|
other_idx
|
||||||
].item() == -1 and mamba_lazy_spec_in_window(
|
].item() == -1 and mamba_lazy_spec_in_window(
|
||||||
req,
|
req,
|
||||||
server_args.mamba_track_interval,
|
get_exec().mamba.mamba_track_interval,
|
||||||
server_args.max_speculative_num_draft_tokens,
|
server_args.max_speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
@@ -1102,7 +1108,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
For spec decode, the boundary is detected by comparing the
|
For spec decode, the boundary is detected by comparing the
|
||||||
accepted seq_len range against interval boundaries.
|
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 batch.spec_algorithm.is_none():
|
||||||
if req.kv_committed_len % interval == 0:
|
if req.kv_committed_len % interval == 0:
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ from sglang.srt.model_executor.cuda_graph_config import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.observability.metrics_collector import DPCooperationInfo
|
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.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils.common import require_mlp_tp_gather
|
from sglang.srt.utils.common import require_mlp_tp_gather
|
||||||
@@ -385,7 +386,7 @@ class SchedulerDPAttnAdapter:
|
|||||||
get_idle_batch=self.get_idle_batch,
|
get_idle_batch=self.get_idle_batch,
|
||||||
disable_cuda_graph=cuda_graph_fully_disabled(),
|
disable_cuda_graph=cuda_graph_fully_disabled(),
|
||||||
require_mlp_tp_gather=require_mlp_tp_gather(self.server_args),
|
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,
|
offload_tags=self.offload_tags,
|
||||||
dwdp=self.server_args.dwdp_size > 1,
|
dwdp=self.server_args.dwdp_size > 1,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from sglang.srt.managers.load_snapshot import (
|
|||||||
QueueMetrics,
|
QueueMetrics,
|
||||||
SpeculativeMetrics,
|
SpeculativeMetrics,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_lora
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
@@ -155,7 +156,7 @@ class SchedulerLoadInquirer:
|
|||||||
)
|
)
|
||||||
|
|
||||||
lora = None
|
lora = None
|
||||||
if self.server_args.enable_lora:
|
if get_lora().enable_lora:
|
||||||
lora = LoRAMetrics(
|
lora = LoRAMetrics(
|
||||||
slots_used=stats.lora_pool_slots_used,
|
slots_used=stats.lora_pool_slots_used,
|
||||||
slots_total=stats.lora_pool_slots_total,
|
slots_total=stats.lora_pool_slots_total,
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import torch
|
|||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
|
from sglang.srt.runtime_context import get_exec
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
MIS_DELIMITER_TOKEN_ID,
|
MIS_DELIMITER_TOKEN_ID,
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
@@ -164,7 +165,7 @@ class SchedulerLogprobResultProcessor:
|
|||||||
delimiter token receive logprobs.
|
delimiter token receive logprobs.
|
||||||
"""
|
"""
|
||||||
return (
|
return (
|
||||||
self.server_args.enable_mis
|
get_exec().features.enable_mis
|
||||||
and req.is_prefill_only
|
and req.is_prefill_only
|
||||||
and req.multi_item_delimiter_indices is not None
|
and req.multi_item_delimiter_indices is not None
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ from sglang.srt.observability.metrics_collector import (
|
|||||||
SchedulerStats,
|
SchedulerStats,
|
||||||
compute_routing_key_stats,
|
compute_routing_key_stats,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_spec
|
||||||
from sglang.srt.utils.device_timer import DeviceTimer
|
from sglang.srt.utils.device_timer import DeviceTimer
|
||||||
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
from sglang.srt.utils.scheduler_status_logger import SchedulerStatusLogger
|
||||||
|
|
||||||
@@ -764,12 +765,10 @@ class SchedulerMetricsReporter:
|
|||||||
else:
|
else:
|
||||||
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
|
spec_accept_length = self.spec_num_accept_tokens / self.spec_num_forward_ct
|
||||||
num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct
|
num_correct_drafts = self.spec_num_accept_tokens - self.spec_num_forward_ct
|
||||||
if self.scheduler.server_args.speculative_num_draft_tokens:
|
if get_spec().speculative_num_draft_tokens:
|
||||||
draft_per_round = (
|
draft_per_round = get_spec().speculative_num_draft_tokens - 1
|
||||||
self.scheduler.server_args.speculative_num_draft_tokens - 1
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
draft_per_round = self.scheduler.server_args.speculative_num_steps or 0
|
draft_per_round = get_spec().speculative_num_steps or 0
|
||||||
total_draft_tokens = self.spec_num_forward_ct * draft_per_round
|
total_draft_tokens = self.spec_num_forward_ct * draft_per_round
|
||||||
spec_accept_rate = (
|
spec_accept_rate = (
|
||||||
num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0
|
num_correct_drafts / total_draft_tokens if total_draft_tokens > 0 else 0
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from sglang.srt.managers.schedule_batch import (
|
|||||||
Req,
|
Req,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache
|
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.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
|
|
||||||
@@ -153,7 +154,7 @@ class SchedulerOutputStreamer:
|
|||||||
return_sampling_mask=return_sampling_mask,
|
return_sampling_mask=return_sampling_mask,
|
||||||
spec_algorithm=self.spec_algorithm,
|
spec_algorithm=self.spec_algorithm,
|
||||||
disaggregation_mode=self.disaggregation_mode,
|
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,
|
default_force_stream_interval=DEFAULT_FORCE_STREAM_INTERVAL,
|
||||||
get_cached_tokens_details=self.get_cached_tokens_details,
|
get_cached_tokens_details=self.get_cached_tokens_details,
|
||||||
rust_server_mode=self.rust_server is not None,
|
rust_server_mode=self.rust_server is not None,
|
||||||
@@ -184,7 +185,7 @@ class SchedulerOutputStreamer:
|
|||||||
if (
|
if (
|
||||||
req.finished()
|
req.finished()
|
||||||
and self.ps.attn_tp_rank == 0
|
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()
|
req.log_time_stats()
|
||||||
|
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
|
from sglang.srt.managers.io_struct import ProfileReq, ProfileReqOutput, ProfileReqType
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||||
from sglang.srt.platforms import current_platform
|
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 import is_mps, is_npu
|
||||||
from sglang.srt.utils.profile_merger import ProfileMerger
|
from sglang.srt.utils.profile_merger import ProfileMerger
|
||||||
from sglang.srt.utils.profile_utils import ProfileManager
|
from sglang.srt.utils.profile_utils import ProfileManager
|
||||||
@@ -257,7 +257,7 @@ class SchedulerProfilerManager:
|
|||||||
self.profile_in_progress = True
|
self.profile_in_progress = True
|
||||||
|
|
||||||
if "CUDA_PROFILER" in activities:
|
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()
|
torch.cuda.cudart().cudaProfilerStart()
|
||||||
self.profile_in_progress = True
|
self.profile_in_progress = True
|
||||||
|
|
||||||
@@ -368,7 +368,7 @@ class SchedulerProfilerManager:
|
|||||||
torch.cuda.memory._record_memory_history(enabled=None)
|
torch.cuda.memory._record_memory_history(enabled=None)
|
||||||
|
|
||||||
if "CUDA_PROFILER" in self.profiler_activities:
|
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()
|
torch.cuda.cudart().cudaProfilerStop()
|
||||||
|
|
||||||
merge_message = self._merge_profile_traces()
|
merge_message = self._merge_profile_traces()
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from sglang.srt.managers.mm_utils import (
|
|||||||
has_shm_features,
|
has_shm_features,
|
||||||
unwrap_shm_features,
|
unwrap_shm_features,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_disagg
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
broadcast_pyobj,
|
broadcast_pyobj,
|
||||||
point_to_point_pyobj,
|
point_to_point_pyobj,
|
||||||
@@ -231,8 +232,8 @@ class SchedulerRequestReceiver:
|
|||||||
# Process MM requests under EPD-disaggregation mode
|
# Process MM requests under EPD-disaggregation mode
|
||||||
if (
|
if (
|
||||||
self.ps.pp_rank == 0
|
self.ps.pp_rank == 0
|
||||||
and self.server_args.language_only
|
and get_disagg().language_only
|
||||||
and self.server_args.encoder_transfer_backend
|
and get_disagg().encoder_transfer_backend
|
||||||
in ["zmq_to_scheduler", "mooncake"]
|
in ["zmq_to_scheduler", "mooncake"]
|
||||||
):
|
):
|
||||||
recv_reqs, abort_reqs = self.mm_receiver.process_waiting_requests(recv_reqs)
|
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,
|
PPProxyTensors,
|
||||||
)
|
)
|
||||||
from sglang.srt.observability.req_time_stats import set_time_batch
|
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.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
|
from sglang.srt.utils import DynamicGradMode, broadcast_pyobj, point_to_point_pyobj
|
||||||
from sglang.srt.utils.common import get_device_module, is_xpu
|
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()
|
self.decode_offload_manager.check_offload_progress()
|
||||||
|
|
||||||
if rmbs[next_mb_id] is not None:
|
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_transfer_queue.queue)
|
||||||
+ len(self.disagg_decode_prealloc_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)
|
queue_size += len(self.decode_offload_manager.ongoing_offload)
|
||||||
|
|
||||||
if server_is_idle and queue_size == 0:
|
if server_is_idle and queue_size == 0:
|
||||||
|
|||||||
@@ -47,6 +47,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
PPProxyTensors,
|
PPProxyTensors,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
|
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.server_args import ServerArgs
|
||||||
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
|
from sglang.srt.utils import MultiprocessingSerializer, broadcast_pyobj, set_random_seed
|
||||||
from sglang.srt.utils.hf_transformers_utils import (
|
from sglang.srt.utils.hf_transformers_utils import (
|
||||||
@@ -408,14 +409,14 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
self.model_config = ModelConfig.from_server_args(
|
self.model_config = ModelConfig.from_server_args(
|
||||||
self.server_args,
|
self.server_args,
|
||||||
model_path=(
|
model_path=(
|
||||||
self.server_args.model_path
|
get_model().model_path
|
||||||
if not self.is_draft_worker
|
if not self.is_draft_worker
|
||||||
else self.server_args.speculative_draft_model_path
|
else get_spec().speculative_draft_model_path
|
||||||
),
|
),
|
||||||
model_revision=(
|
model_revision=(
|
||||||
self.server_args.revision
|
get_model().revision
|
||||||
if not self.is_draft_worker
|
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,
|
is_draft_model=self.is_draft_worker,
|
||||||
context_length=self.context_length,
|
context_length=self.context_length,
|
||||||
@@ -426,7 +427,7 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
|
|
||||||
self._model_runner = ModelRunner(
|
self._model_runner = ModelRunner(
|
||||||
model_config=self.model_config,
|
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,
|
gpu_id=self.gpu_id,
|
||||||
ps=self.ps,
|
ps=self.ps,
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
@@ -442,11 +443,11 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
self.model_runner_list.append(self.model_runner)
|
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(
|
self.model_runner_list.append(
|
||||||
ModelRunner(
|
ModelRunner(
|
||||||
model_config=self.model_config,
|
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,
|
gpu_id=self.gpu_id,
|
||||||
ps=self.ps,
|
ps=self.ps,
|
||||||
nccl_port=self.nccl_port,
|
nccl_port=self.nccl_port,
|
||||||
@@ -462,7 +463,7 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
def _init_dllm_algorithm(self):
|
def _init_dllm_algorithm(self):
|
||||||
from sglang.srt.dllm.algorithm.base import DllmAlgorithm
|
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)
|
self.dllm_algorithm = DllmAlgorithm.from_server_args(self.server_args)
|
||||||
else:
|
else:
|
||||||
self.dllm_algorithm = None
|
self.dllm_algorithm = None
|
||||||
@@ -488,9 +489,9 @@ class TpModelWorker(BaseTpWorker):
|
|||||||
)
|
)
|
||||||
return (
|
return (
|
||||||
self.model_runner.max_total_num_tokens,
|
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.model_runner.max_running_requests,
|
||||||
self.server_args.max_queued_requests,
|
get_schedule().max_queued_requests,
|
||||||
max_req_len,
|
max_req_len,
|
||||||
max_req_len - 5,
|
max_req_len - 5,
|
||||||
self.random_seed,
|
self.random_seed,
|
||||||
|
|||||||
@@ -23,6 +23,7 @@ from sglang.srt.observability.metrics_collector import (
|
|||||||
RadixCacheMetricsCollector,
|
RadixCacheMetricsCollector,
|
||||||
resolve_collector_class,
|
resolve_collector_class,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_observability
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
@@ -238,8 +239,8 @@ class BasePrefixCache(ABC, PrefixCacheTrait):
|
|||||||
|
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
labels = {"cache_type": self.__class__.__name__}
|
labels = {"cache_type": self.__class__.__name__}
|
||||||
if server_args.extra_metric_labels:
|
if get_observability().extra_metric_labels:
|
||||||
labels.update(server_args.extra_metric_labels)
|
labels.update(get_observability().extra_metric_labels)
|
||||||
radix_cache_cls = resolve_collector_class(
|
radix_cache_cls = resolve_collector_class(
|
||||||
server_args,
|
server_args,
|
||||||
STAT_LOGGER_ROLE_RADIX_CACHE,
|
STAT_LOGGER_ROLE_RADIX_CACHE,
|
||||||
|
|||||||
@@ -16,7 +16,12 @@ 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.allocator.swa import SWATokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
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_schedule,
|
||||||
|
get_server_args,
|
||||||
|
get_serving,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.utils.common import ceil_align
|
from sglang.srt.utils.common import ceil_align
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -179,12 +184,12 @@ def _release_overallocated_kv_indices(
|
|||||||
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
|
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
|
||||||
) -> None:
|
) -> None:
|
||||||
global_server_args = get_server_args()
|
global_server_args = get_server_args()
|
||||||
page_size = global_server_args.page_size
|
page_size = get_schedule().page_size
|
||||||
spec_algo = global_server_args.speculative_algorithm
|
spec_algo = get_spec().speculative_algorithm
|
||||||
|
|
||||||
# strip_thinking_cache intentionally reports output tokens as overallocated
|
# strip_thinking_cache intentionally reports output tokens as overallocated
|
||||||
# so they fall into the free path below (#22373).
|
# 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 (
|
assert (
|
||||||
start_p == end_p
|
start_p == end_p
|
||||||
), f"Unexpected overallocated KV cache, {req.kv_committed_len=}, {req.kv.kv_allocated_len=}"
|
), 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.base_swa_memory_pool import BaseSWAKVPool
|
||||||
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
|
from sglang.srt.mem_cache.deepseek_v4_compress_state import CompressStatePool
|
||||||
from sglang.srt.mem_cache.memory_pool import KVCache
|
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, get_spec
|
||||||
from sglang.srt.utils import ceil_div, is_hip
|
from sglang.srt.utils import ceil_div, is_hip
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -276,7 +276,7 @@ class DeepSeekV4IndexerPool(KVCache):
|
|||||||
end_layer,
|
end_layer,
|
||||||
)
|
)
|
||||||
self.index_head_dim = index_head_dim
|
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()
|
self._create_buffer()
|
||||||
|
|
||||||
@@ -577,8 +577,8 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
self.c128_kv_pool = None
|
self.c128_kv_pool = None
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
spec_extra = (
|
spec_extra = (
|
||||||
(server_args.speculative_num_draft_tokens - 1)
|
(get_spec().speculative_num_draft_tokens - 1)
|
||||||
if server_args.speculative_algorithm is not None
|
if get_spec().speculative_algorithm is not None
|
||||||
else 0
|
else 0
|
||||||
)
|
)
|
||||||
self.unified_kv_pool = DeepSeekV4UnifiedKVPool(
|
self.unified_kv_pool = DeepSeekV4UnifiedKVPool(
|
||||||
@@ -659,7 +659,7 @@ class DeepSeekV4TokenToKVPool(BaseSWAKVPool):
|
|||||||
|
|
||||||
def get_ring_size(self, compress_ratio: int) -> int:
|
def get_ring_size(self, compress_ratio: int) -> int:
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
is_speculative = server_args.speculative_algorithm is not None
|
is_speculative = get_spec().speculative_algorithm is not None
|
||||||
return get_compress_state_ring_size(compress_ratio, is_speculative)
|
return get_compress_state_ring_size(compress_ratio, is_speculative)
|
||||||
|
|
||||||
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
|
def translate_loc_from_full_to_swa(self, kv_indices: torch.Tensor):
|
||||||
|
|||||||
@@ -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.mem_cache.swa_memory_pool import SWAKVPool
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import get_model, get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
get_context,
|
||||||
|
get_disagg,
|
||||||
|
get_exec,
|
||||||
|
get_memory,
|
||||||
|
get_parallel,
|
||||||
|
get_schedule,
|
||||||
|
get_spec,
|
||||||
|
)
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
@@ -184,6 +192,7 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator]
|
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator]
|
||||||
memory_pool_config: Optional[MemoryPoolConfig]
|
memory_pool_config: Optional[MemoryPoolConfig]
|
||||||
draft_model_idx: Optional[int] = None
|
draft_model_idx: Optional[int] = None
|
||||||
|
kv_cache_dtype_str: Optional[str] = None
|
||||||
mambaish_config: Optional[Any] = field(init=False)
|
mambaish_config: Optional[Any] = field(init=False)
|
||||||
hybrid_gdn_config: Optional[Any] = field(init=False)
|
hybrid_gdn_config: Optional[Any] = field(init=False)
|
||||||
is_inkling_mtp_draft: bool = field(init=False)
|
is_inkling_mtp_draft: bool = field(init=False)
|
||||||
@@ -211,7 +220,7 @@ class KVCacheConfigurator:
|
|||||||
def _build_fp4_quant_method(self, *, num_layers: int):
|
def _build_fp4_quant_method(self, *, num_layers: int):
|
||||||
if not is_float4_e2m1fn_x2(self.kv_cache_dtype):
|
if not is_float4_e2m1fn_x2(self.kv_cache_dtype):
|
||||||
return None
|
return None
|
||||||
quant_name = resolve_kv_cache_quant(get_model().kv_cache_dtype)
|
quant_name = resolve_kv_cache_quant(self.kv_cache_dtype_str)
|
||||||
if quant_name is None:
|
if quant_name is None:
|
||||||
return None
|
return None
|
||||||
quant_method = get_kv_cache_quant_method(
|
quant_method = get_kv_cache_quant_method(
|
||||||
@@ -314,8 +323,8 @@ class KVCacheConfigurator:
|
|||||||
# from one byte buffer, then return. Gated to the target worker
|
# 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).
|
# (req_to_token_pool is None); supports hybrid Mamba and hybrid SWA (not DSV4).
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_unified_memory
|
get_memory().enable_unified_memory
|
||||||
and self.server_args.disaggregation_mode == "null"
|
and get_disagg().disaggregation_mode == "null"
|
||||||
and req_to_token_pool is None
|
and req_to_token_pool is None
|
||||||
):
|
):
|
||||||
if self.mambaish_config is not None:
|
if self.mambaish_config is not None:
|
||||||
@@ -364,13 +373,13 @@ class KVCacheConfigurator:
|
|||||||
# TARGET_VERIFY, so their pools skip the per-step intermediate
|
# TARGET_VERIFY, so their pools skip the per-step intermediate
|
||||||
# (SpeculativeState) buffers only the target pool consumes.
|
# (SpeculativeState) buffers only the target pool consumes.
|
||||||
req_to_token_pool = req_to_token_pool.clone_with_new_mamba(
|
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,
|
mamba_spec_state_size=sizes.max_running_requests,
|
||||||
cache_params=self.mambaish_config.mamba2_cache_params,
|
cache_params=self.mambaish_config.mamba2_cache_params,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
draft_model_idx=self.draft_model_idx,
|
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
|
# Initialize token_to_kv_pool
|
||||||
@@ -400,7 +409,7 @@ class KVCacheConfigurator:
|
|||||||
# unsupported pool families before allocation. Keep this guard here so
|
# unsupported pool families before allocation. Keep this guard here so
|
||||||
# future pool-selection refactors fail at boot instead of on first use.
|
# future pool-selection refactors fail at boot instead of on first use.
|
||||||
if (
|
if (
|
||||||
self.server_args.prefill_only_disable_kv_cache
|
get_schedule().prefill_only_disable_kv_cache
|
||||||
and not self.is_draft_worker
|
and not self.is_draft_worker
|
||||||
and not isinstance(token_to_kv_pool, NoOpMHATokenToKVPool)
|
and not isinstance(token_to_kv_pool, NoOpMHATokenToKVPool)
|
||||||
):
|
):
|
||||||
@@ -435,8 +444,8 @@ class KVCacheConfigurator:
|
|||||||
assert self.page_size >= 1, f"page_size must be >= 1, got {self.page_size}"
|
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.
|
# Mirror the non-shared path's extra_max_context_len computation.
|
||||||
extra_max_context_len = 4
|
extra_max_context_len = 4
|
||||||
if self.server_args.speculative_num_draft_tokens is not None:
|
if get_spec().speculative_num_draft_tokens is not None:
|
||||||
extra_max_context_len += self.server_args.speculative_num_draft_tokens
|
extra_max_context_len += get_spec().speculative_num_draft_tokens
|
||||||
|
|
||||||
mamba_layer_ids = [
|
mamba_layer_ids = [
|
||||||
i
|
i
|
||||||
@@ -471,14 +480,14 @@ class KVCacheConfigurator:
|
|||||||
model_context_len=self.model_config.context_len,
|
model_context_len=self.model_config.context_len,
|
||||||
extra_max_context_len=extra_max_context_len,
|
extra_max_context_len=extra_max_context_len,
|
||||||
max_total_num_tokens=max_total_num_tokens,
|
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,
|
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(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
speculative_num_draft_tokens=self.server_args.speculative_num_draft_tokens,
|
speculative_num_draft_tokens=get_spec().speculative_num_draft_tokens,
|
||||||
disable_overlap_schedule=self.server_args.disable_overlap_schedule,
|
disable_overlap_schedule=get_schedule().disable_overlap_schedule,
|
||||||
need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"),
|
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
|
||||||
mamba_full_memory_ratio=self.server_args.mamba_full_memory_ratio,
|
mamba_full_memory_ratio=get_schedule().mamba_full_memory_ratio,
|
||||||
# Overlap mode: the allocator's `free` drops a wait_stream(forward_stream)
|
# Overlap mode: the allocator's `free` drops a wait_stream(forward_stream)
|
||||||
# barrier so eager compaction serializes after the in-flight forward's
|
# barrier so eager compaction serializes after the in-flight forward's
|
||||||
# v2p/KV reads. Near-no-op in normal mode.
|
# v2p/KV reads. Near-no-op in normal mode.
|
||||||
@@ -511,13 +520,13 @@ class KVCacheConfigurator:
|
|||||||
), "unified memory pool does not support MLA-SWA hybrid yet"
|
), "unified memory pool does not support MLA-SWA hybrid yet"
|
||||||
# Mirror the non-shared path's extra_max_context_len computation.
|
# Mirror the non-shared path's extra_max_context_len computation.
|
||||||
extra_max_context_len = 4
|
extra_max_context_len = 4
|
||||||
if self.server_args.speculative_num_draft_tokens is not None:
|
if get_spec().speculative_num_draft_tokens is not None:
|
||||||
extra_max_context_len += self.server_args.speculative_num_draft_tokens
|
extra_max_context_len += get_spec().speculative_num_draft_tokens
|
||||||
req_to_token_pool = ReqToTokenPool(
|
req_to_token_pool = ReqToTokenPool(
|
||||||
size=max_num_reqs,
|
size=max_num_reqs,
|
||||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||||
device=self.device,
|
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)
|
head_num = self.model_config.get_num_kv_heads(get_parallel().attn_tp_size)
|
||||||
@@ -567,8 +576,8 @@ class KVCacheConfigurator:
|
|||||||
full_attention_layer_ids=full_attention_layer_ids,
|
full_attention_layer_ids=full_attention_layer_ids,
|
||||||
full_max_total_num_tokens=full_max_total_num_tokens,
|
full_max_total_num_tokens=full_max_total_num_tokens,
|
||||||
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
swa_max_total_num_tokens=swa_max_total_num_tokens,
|
||||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
need_sort=self.server_args.disaggregation_mode in ("decode", "prefill"),
|
need_sort=get_disagg().disaggregation_mode in ("decode", "prefill"),
|
||||||
# Overlap mode: same wait_stream(forward_stream) rationale as
|
# Overlap mode: same wait_stream(forward_stream) rationale as
|
||||||
# `_init_unified_mamba_pools`.
|
# `_init_unified_mamba_pools`.
|
||||||
forward_stream=self.forward_stream,
|
forward_stream=self.forward_stream,
|
||||||
@@ -588,7 +597,7 @@ class KVCacheConfigurator:
|
|||||||
is_dsv4_model: bool,
|
is_dsv4_model: bool,
|
||||||
current_platform,
|
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
|
return
|
||||||
|
|
||||||
unsupported_pool_family = None
|
unsupported_pool_family = None
|
||||||
@@ -623,9 +632,9 @@ class KVCacheConfigurator:
|
|||||||
def _build_req_to_token_pool(self, *, max_num_reqs: int) -> ReqToTokenPool:
|
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)
|
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
|
# 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:
|
if self.mambaish_config:
|
||||||
req_to_token_pool = self._build_hybrid_mamba_decode_req_pool(
|
req_to_token_pool = self._build_hybrid_mamba_decode_req_pool(
|
||||||
max_num_reqs=max_num_reqs,
|
max_num_reqs=max_num_reqs,
|
||||||
@@ -665,7 +674,7 @@ class KVCacheConfigurator:
|
|||||||
size=max_num_reqs,
|
size=max_num_reqs,
|
||||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||||
device=self.device,
|
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,
|
cache_params=self.mambaish_config.mamba2_cache_params,
|
||||||
mamba_layer_ids=(
|
mamba_layer_ids=(
|
||||||
[
|
[
|
||||||
@@ -675,16 +684,16 @@ class KVCacheConfigurator:
|
|||||||
]
|
]
|
||||||
),
|
),
|
||||||
speculative_num_draft_tokens=self.server_args.max_speculative_num_draft_tokens,
|
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(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
pre_alloc_size=pre_alloc_size,
|
pre_alloc_size=pre_alloc_size,
|
||||||
enable_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
|
||||||
mamba_size=self.server_args.max_mamba_cache_size,
|
mamba_size=get_schedule().max_mamba_cache_size,
|
||||||
start_layer=self.layer_info.start_layer,
|
start_layer=self.layer_info.start_layer,
|
||||||
linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len,
|
linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len,
|
||||||
mamba_envelope_layout=self.server_args.enable_page_major_kv_layout,
|
mamba_envelope_layout=get_memory().enable_page_major_kv_layout,
|
||||||
enable_gdn_replayssm_spec=(
|
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
|
and self.hybrid_gdn_config is not None
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -708,7 +717,7 @@ class KVCacheConfigurator:
|
|||||||
size=max_num_reqs,
|
size=max_num_reqs,
|
||||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||||
device=self.device,
|
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,
|
pre_alloc_size=pre_alloc_size,
|
||||||
)
|
)
|
||||||
return req_to_token_pool
|
return req_to_token_pool
|
||||||
@@ -721,11 +730,11 @@ class KVCacheConfigurator:
|
|||||||
) -> ReqToTokenPool:
|
) -> ReqToTokenPool:
|
||||||
req_to_token_pool = HybridReqToTokenPool(
|
req_to_token_pool = HybridReqToTokenPool(
|
||||||
size=max_num_reqs,
|
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,
|
mamba_spec_state_size=max_num_reqs,
|
||||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||||
device=self.device,
|
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,
|
cache_params=self.mambaish_config.mamba2_cache_params,
|
||||||
mamba_layer_ids=(
|
mamba_layer_ids=(
|
||||||
[
|
[
|
||||||
@@ -737,14 +746,14 @@ class KVCacheConfigurator:
|
|||||||
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
enable_mamba_extra_buffer=self.server_args.enable_mamba_extra_buffer(),
|
||||||
enable_mamba_extra_buffer_lazy=self.server_args.enable_mamba_extra_buffer_lazy(),
|
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_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_overlap_schedule=not self.server_args.disable_overlap_schedule,
|
enable_overlap_schedule=not get_schedule().disable_overlap_schedule,
|
||||||
start_layer=self.layer_info.start_layer,
|
start_layer=self.layer_info.start_layer,
|
||||||
enable_linear_replayssm=self.server_args.enable_linear_replayssm,
|
enable_linear_replayssm=get_exec().mamba.enable_linear_replayssm,
|
||||||
linear_replayssm_cache_len=self.server_args.linear_replayssm_cache_len,
|
linear_replayssm_cache_len=get_exec().mamba.linear_replayssm_cache_len,
|
||||||
mamba_envelope_layout=self.server_args.enable_page_major_kv_layout,
|
mamba_envelope_layout=get_memory().enable_page_major_kv_layout,
|
||||||
enable_gdn_replayssm_spec=(
|
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
|
and self.hybrid_gdn_config is not None
|
||||||
),
|
),
|
||||||
)
|
)
|
||||||
@@ -770,7 +779,7 @@ class KVCacheConfigurator:
|
|||||||
size=max_num_reqs,
|
size=max_num_reqs,
|
||||||
max_context_len=self.model_config.context_len + extra_max_context_len,
|
max_context_len=self.model_config.context_len + extra_max_context_len,
|
||||||
device=self.device,
|
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
|
return req_to_token_pool
|
||||||
|
|
||||||
@@ -786,7 +795,7 @@ class KVCacheConfigurator:
|
|||||||
# selected by swapping in the PageMajorMHATokenToKVPool subclass. The
|
# selected by swapping in the PageMajorMHATokenToKVPool subclass. The
|
||||||
# default keeps upstream's per-layer layout. The Mamba state pool is routed
|
# 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.
|
# 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 = (
|
mha_pool_class = (
|
||||||
PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool
|
PageMajorMHATokenToKVPool if enable_page_major else MHATokenToKVPool
|
||||||
)
|
)
|
||||||
@@ -894,7 +903,7 @@ class KVCacheConfigurator:
|
|||||||
c128_state_dtype: Optional[torch.dtype],
|
c128_state_dtype: Optional[torch.dtype],
|
||||||
req_to_token_pool: ReqToTokenPool,
|
req_to_token_pool: ReqToTokenPool,
|
||||||
) -> KVCache:
|
) -> KVCache:
|
||||||
swa_page_size = self.server_args.page_size
|
swa_page_size = get_schedule().page_size
|
||||||
if not _is_npu:
|
if not _is_npu:
|
||||||
assert swa_page_size == 256, "In paged swa mode, page_size must be 256."
|
assert swa_page_size == 256, "In paged swa mode, page_size must be 256."
|
||||||
|
|
||||||
@@ -928,12 +937,12 @@ class KVCacheConfigurator:
|
|||||||
# sliding eviction in ``ScheduleBatch._evict_swa``.
|
# sliding eviction in ``ScheduleBatch._evict_swa``.
|
||||||
c4_state_pool_size = npu_state_pool_size(
|
c4_state_pool_size = npu_state_pool_size(
|
||||||
ratio=4,
|
ratio=4,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
max_num_reqs=max_running_requests,
|
max_num_reqs=max_running_requests,
|
||||||
)
|
)
|
||||||
c128_state_pool_size = npu_state_pool_size(
|
c128_state_pool_size = npu_state_pool_size(
|
||||||
ratio=128,
|
ratio=128,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
max_num_reqs=max_running_requests,
|
max_num_reqs=max_running_requests,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
@@ -951,7 +960,7 @@ class KVCacheConfigurator:
|
|||||||
c128_size=c128_max_total_num_tokens,
|
c128_size=c128_max_total_num_tokens,
|
||||||
c4_state_pool_size=c4_state_pool_size,
|
c4_state_pool_size=c4_state_pool_size,
|
||||||
c128_state_pool_size=c128_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,
|
swa_page_size=swa_page_size,
|
||||||
sliding_window=self.model_config.window_size,
|
sliding_window=self.model_config.window_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
@@ -962,11 +971,11 @@ class KVCacheConfigurator:
|
|||||||
indexer_head_dim=self.model_config.index_head_dim,
|
indexer_head_dim=self.model_config.index_head_dim,
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
compression_ratios=compression_ratios,
|
||||||
start_layer=self.layer_info.start_layer,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_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=(
|
online_mtp_max_draft_tokens=(
|
||||||
self.server_args.max_speculative_num_draft_tokens or 0
|
self.server_args.max_speculative_num_draft_tokens or 0
|
||||||
),
|
),
|
||||||
@@ -977,7 +986,7 @@ class KVCacheConfigurator:
|
|||||||
PoolCls = current_platform.get_dsa_kv_pool_cls()
|
PoolCls = current_platform.get_dsa_kv_pool_cls()
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||||
@@ -988,7 +997,7 @@ class KVCacheConfigurator:
|
|||||||
kv_cache_dtype=self.kv_cache_dtype,
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
server_args=self.server_args,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
|
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
|
||||||
@@ -1001,14 +1010,14 @@ class KVCacheConfigurator:
|
|||||||
PoolCls = current_platform.get_mla_kv_pool_cls()
|
PoolCls = current_platform.get_mla_kv_pool_cls()
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
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),
|
index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None),
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
)
|
)
|
||||||
@@ -1018,13 +1027,13 @@ class KVCacheConfigurator:
|
|||||||
PoolCls = current_platform.get_mha_kv_pool_cls()
|
PoolCls = current_platform.get_mha_kv_pool_cls()
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
)
|
)
|
||||||
@@ -1055,7 +1064,7 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool = SWAKVPool(
|
token_to_kv_pool = SWAKVPool(
|
||||||
size=full_max_total_num_tokens,
|
size=full_max_total_num_tokens,
|
||||||
size_swa=swa_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,
|
dtype=self.kv_cache_dtype,
|
||||||
post_capture_active=self.post_capture_kv_active,
|
post_capture_active=self.post_capture_kv_active,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
@@ -1077,14 +1086,14 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
token_to_kv_pool = NPUMLATokenToKVPool(
|
token_to_kv_pool = NPUMLATokenToKVPool(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
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),
|
index_head_dim=(self.model_config.index_head_dim if is_dsa_model else None),
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
)
|
)
|
||||||
@@ -1097,13 +1106,13 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
token_to_kv_pool = NPUMHATokenToKVPool(
|
token_to_kv_pool = NPUMHATokenToKVPool(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
)
|
)
|
||||||
@@ -1117,7 +1126,7 @@ class KVCacheConfigurator:
|
|||||||
dsa_cp_layer_shard_size,
|
dsa_cp_layer_shard_size,
|
||||||
) = get_glm_dsa_cp_layer_shard_info(self)
|
) = get_glm_dsa_cp_layer_shard_info(self)
|
||||||
pool_kwargs = {}
|
pool_kwargs = {}
|
||||||
if self.server_args.enable_hisparse:
|
if get_memory().enable_hisparse:
|
||||||
PoolCls = HiSparseDSATokenToKVPool
|
PoolCls = HiSparseDSATokenToKVPool
|
||||||
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
from sglang.srt.mem_cache.sparsity import parse_hisparse_config
|
||||||
|
|
||||||
@@ -1137,7 +1146,7 @@ class KVCacheConfigurator:
|
|||||||
PoolCls = DSATokenToKVPool
|
PoolCls = DSATokenToKVPool
|
||||||
token_to_kv_pool = PoolCls(
|
token_to_kv_pool = PoolCls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||||
@@ -1148,7 +1157,7 @@ class KVCacheConfigurator:
|
|||||||
kv_cache_dtype=self.kv_cache_dtype,
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
server_args=self.server_args,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
|
index_head_dim=get_dsa_index_head_dim(self.model_config.hf_config),
|
||||||
@@ -1159,13 +1168,13 @@ class KVCacheConfigurator:
|
|||||||
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
def _build_mla_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||||
token_to_kv_pool = MLATokenToKVPoolFP4(
|
token_to_kv_pool = MLATokenToKVPoolFP4(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
)
|
)
|
||||||
@@ -1174,13 +1183,13 @@ class KVCacheConfigurator:
|
|||||||
def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
def _build_mla_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||||
token_to_kv_pool = MLATokenToKVPool(
|
token_to_kv_pool = MLATokenToKVPool(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
kv_lora_rank=self.model_config.kv_lora_rank,
|
kv_lora_rank=self.model_config.kv_lora_rank,
|
||||||
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
qk_rope_head_dim=self.model_config.qk_rope_head_dim,
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
)
|
)
|
||||||
@@ -1207,7 +1216,7 @@ class KVCacheConfigurator:
|
|||||||
}
|
}
|
||||||
swa_pool_class = (
|
swa_pool_class = (
|
||||||
MHATokenToKVPoolMXFP8
|
MHATokenToKVPoolMXFP8
|
||||||
if get_model().kv_cache_dtype == "mxfp8"
|
if self.kv_cache_dtype_str == "mxfp8"
|
||||||
else mha_pool_class
|
else mha_pool_class
|
||||||
)
|
)
|
||||||
swa_attention_layer_ids = self.model_config.swa_attention_layer_ids
|
swa_attention_layer_ids = self.model_config.swa_attention_layer_ids
|
||||||
@@ -1237,7 +1246,7 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool = SWAKVPool(
|
token_to_kv_pool = SWAKVPool(
|
||||||
size=full_max_total_num_tokens,
|
size=full_max_total_num_tokens,
|
||||||
size_swa=size_swa,
|
size_swa=size_swa,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
post_capture_active=self.post_capture_kv_active,
|
post_capture_active=self.post_capture_kv_active,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
@@ -1245,7 +1254,7 @@ class KVCacheConfigurator:
|
|||||||
swa_attention_layer_ids=swa_attention_layer_ids,
|
swa_attention_layer_ids=swa_attention_layer_ids,
|
||||||
full_attention_layer_ids=full_attention_layer_ids,
|
full_attention_layer_ids=full_attention_layer_ids,
|
||||||
device=self.device,
|
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,
|
token_to_kv_pool_class=swa_pool_class,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
@@ -1260,7 +1269,7 @@ class KVCacheConfigurator:
|
|||||||
)
|
)
|
||||||
token_to_kv_pool = MiniMaxSparseKVPool(
|
token_to_kv_pool = MiniMaxSparseKVPool(
|
||||||
size=max_total_num_tokens,
|
size=max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
index_dtype=self.model_dtype,
|
index_dtype=self.model_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
@@ -1270,7 +1279,7 @@ class KVCacheConfigurator:
|
|||||||
sparse_layer_ids=sparse_layer_ids,
|
sparse_layer_ids=sparse_layer_ids,
|
||||||
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
|
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
)
|
)
|
||||||
@@ -1305,11 +1314,11 @@ class KVCacheConfigurator:
|
|||||||
# buffers) for the full-attention layers, same as the SWA branch.
|
# buffers) for the full-attention layers, same as the SWA branch.
|
||||||
full_pool_class = (
|
full_pool_class = (
|
||||||
MHATokenToKVPoolMXFP8
|
MHATokenToKVPoolMXFP8
|
||||||
if get_model().kv_cache_dtype == "mxfp8" and not self.use_mla_backend
|
if self.kv_cache_dtype_str == "mxfp8" and not self.use_mla_backend
|
||||||
else mha_pool_class
|
else mha_pool_class
|
||||||
)
|
)
|
||||||
token_to_kv_pool = HybridLinearKVPool(
|
token_to_kv_pool = HybridLinearKVPool(
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
size=max_total_num_tokens,
|
size=max_total_num_tokens,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
@@ -1318,8 +1327,8 @@ class KVCacheConfigurator:
|
|||||||
full_attention_layer_ids=full_attention_layer_ids,
|
full_attention_layer_ids=full_attention_layer_ids,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
mamba_pool=req_to_token_pool.mamba_pool,
|
mamba_pool=req_to_token_pool.mamba_pool,
|
||||||
enable_memory_saver=self.server_args.enable_memory_saver,
|
enable_memory_saver=get_exec().features.enable_memory_saver,
|
||||||
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
|
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
|
||||||
use_mla=self.use_mla_backend,
|
use_mla=self.use_mla_backend,
|
||||||
start_layer=self.layer_info.start_layer,
|
start_layer=self.layer_info.start_layer,
|
||||||
full_kv_pool_class=full_pool_class,
|
full_kv_pool_class=full_pool_class,
|
||||||
@@ -1332,30 +1341,30 @@ class KVCacheConfigurator:
|
|||||||
def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
def _build_mha_fp4_kv_pool(self, *, max_total_num_tokens: int) -> KVCache:
|
||||||
token_to_kv_pool = MHATokenToKVPoolFP4(
|
token_to_kv_pool = MHATokenToKVPoolFP4(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
v_head_dim=self.model_config.v_head_dim,
|
v_head_dim=self.model_config.v_head_dim,
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
enable_alt_stream=not self.server_args.enable_pdmux,
|
enable_alt_stream=not get_disagg().enable_pdmux,
|
||||||
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
|
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
|
||||||
)
|
)
|
||||||
return token_to_kv_pool
|
return token_to_kv_pool
|
||||||
|
|
||||||
def _build_mha_kv_pool(
|
def _build_mha_kv_pool(
|
||||||
self, *, max_total_num_tokens: int, mha_pool_class: type, quant_method=None
|
self, *, max_total_num_tokens: int, mha_pool_class: type, quant_method=None
|
||||||
) -> KVCache:
|
) -> KVCache:
|
||||||
if get_model().kv_cache_dtype == "mxfp8":
|
if self.kv_cache_dtype_str == "mxfp8":
|
||||||
pool_cls = MHATokenToKVPoolMXFP8
|
pool_cls = MHATokenToKVPoolMXFP8
|
||||||
else:
|
else:
|
||||||
pool_cls = (
|
pool_cls = (
|
||||||
NoOpMHATokenToKVPool
|
NoOpMHATokenToKVPool
|
||||||
if self.server_args.prefill_only_disable_kv_cache
|
if get_schedule().prefill_only_disable_kv_cache
|
||||||
else mha_pool_class
|
else mha_pool_class
|
||||||
)
|
)
|
||||||
pool_kwargs = {}
|
pool_kwargs = {}
|
||||||
@@ -1365,18 +1374,18 @@ class KVCacheConfigurator:
|
|||||||
pool_kwargs["post_capture_active"] = self.post_capture_kv_active
|
pool_kwargs["post_capture_active"] = self.post_capture_kv_active
|
||||||
token_to_kv_pool = pool_cls(
|
token_to_kv_pool = pool_cls(
|
||||||
max_total_num_tokens,
|
max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
|
||||||
head_dim=self.model_config.head_dim,
|
head_dim=self.model_config.head_dim,
|
||||||
v_head_dim=self.model_config.v_head_dim,
|
v_head_dim=self.model_config.v_head_dim,
|
||||||
layer_num=self.layer_info.num_effective_layers,
|
layer_num=self.layer_info.num_effective_layers,
|
||||||
device=self.device,
|
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,
|
start_layer=self.layer_info.start_layer,
|
||||||
end_layer=self.layer_info.end_layer,
|
end_layer=self.layer_info.end_layer,
|
||||||
enable_alt_stream=not self.server_args.enable_pdmux,
|
enable_alt_stream=not get_disagg().enable_pdmux,
|
||||||
enable_kv_cache_copy=(self.server_args.speculative_algorithm is not None),
|
enable_kv_cache_copy=(get_spec().speculative_algorithm is not None),
|
||||||
**pool_kwargs,
|
**pool_kwargs,
|
||||||
)
|
)
|
||||||
return token_to_kv_pool
|
return token_to_kv_pool
|
||||||
@@ -1391,13 +1400,13 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator],
|
token_to_kv_pool_allocator: Optional[BaseTokenToKVPoolAllocator],
|
||||||
) -> BaseTokenToKVPoolAllocator:
|
) -> BaseTokenToKVPoolAllocator:
|
||||||
# Initialize token_to_kv_pool_allocator
|
# 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 token_to_kv_pool_allocator is None:
|
||||||
if current_platform.is_out_of_tree():
|
if current_platform.is_out_of_tree():
|
||||||
AllocatorCls = current_platform.get_paged_allocator_cls()
|
AllocatorCls = current_platform.get_paged_allocator_cls()
|
||||||
token_to_kv_pool_allocator = AllocatorCls(
|
token_to_kv_pool_allocator = AllocatorCls(
|
||||||
sizes.max_total_num_tokens,
|
sizes.max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
@@ -1422,7 +1431,7 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool_allocator = swa_allocator_cls(
|
token_to_kv_pool_allocator = swa_allocator_cls(
|
||||||
sizes.full_max_total_num_tokens,
|
sizes.full_max_total_num_tokens,
|
||||||
sizes.swa_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,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
@@ -1435,7 +1444,7 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
|
||||||
sizes.max_total_num_tokens,
|
sizes.max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
@@ -1445,7 +1454,7 @@ class KVCacheConfigurator:
|
|||||||
if self.is_hybrid_swa and sizes.full_max_total_num_tokens == 0:
|
if self.is_hybrid_swa and sizes.full_max_total_num_tokens == 0:
|
||||||
token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = PureSWATokenToKVPoolAllocator(
|
||||||
sizes.swa_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,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
@@ -1455,14 +1464,14 @@ class KVCacheConfigurator:
|
|||||||
token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
||||||
sizes.full_max_total_num_tokens,
|
sizes.full_max_total_num_tokens,
|
||||||
sizes.swa_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,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
need_sort=need_sort,
|
need_sort=need_sort,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
if self.server_args.enable_hisparse:
|
if get_memory().enable_hisparse:
|
||||||
from sglang.srt.mem_cache.sparsity import (
|
from sglang.srt.mem_cache.sparsity import (
|
||||||
parse_hisparse_config,
|
parse_hisparse_config,
|
||||||
)
|
)
|
||||||
@@ -1470,7 +1479,7 @@ class KVCacheConfigurator:
|
|||||||
hisparse_cfg = parse_hisparse_config(self.server_args)
|
hisparse_cfg = parse_hisparse_config(self.server_args)
|
||||||
token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = HiSparseTokenToKVPoolAllocator(
|
||||||
sizes.max_total_num_tokens,
|
sizes.max_total_num_tokens,
|
||||||
page_size=self.server_args.page_size,
|
page_size=get_schedule().page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=token_to_kv_pool,
|
kvcache=token_to_kv_pool,
|
||||||
@@ -1478,8 +1487,7 @@ class KVCacheConfigurator:
|
|||||||
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
host_to_device_ratio=hisparse_cfg.host_to_device_ratio,
|
||||||
)
|
)
|
||||||
elif (
|
elif (
|
||||||
self.server_args.page_size == 1
|
get_schedule().page_size == 1 and self.server_args.dcp_size == 1
|
||||||
and self.server_args.dcp_size == 1
|
|
||||||
):
|
):
|
||||||
token_to_kv_pool_allocator = TokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = TokenToKVPoolAllocator(
|
||||||
sizes.max_total_num_tokens,
|
sizes.max_total_num_tokens,
|
||||||
@@ -1491,7 +1499,7 @@ class KVCacheConfigurator:
|
|||||||
else:
|
else:
|
||||||
token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = PagedTokenToKVPoolAllocator(
|
||||||
sizes.max_total_num_tokens * self.server_args.dcp_size,
|
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,
|
* self.server_args.dcp_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
@@ -1499,7 +1507,7 @@ class KVCacheConfigurator:
|
|||||||
need_sort=need_sort,
|
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."
|
assert self.is_hybrid_swa, "DeepSeek V4 HiSparse requires SWA mode."
|
||||||
token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator(
|
token_to_kv_pool_allocator = DeepSeekV4HiSparseTokenToKVPoolAllocator(
|
||||||
token_to_kv_pool_allocator
|
token_to_kv_pool_allocator
|
||||||
@@ -1551,7 +1559,7 @@ class KVCacheConfigurator:
|
|||||||
cpu_group=get_world_group().cpu_group,
|
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:
|
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.
|
# Mamba state is a fixed pre-capture allocation, so it can't ride the ~0 post-capture slack.
|
||||||
slack_gb = max(
|
slack_gb = max(
|
||||||
@@ -1575,7 +1583,7 @@ class KVCacheConfigurator:
|
|||||||
)
|
)
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"Loaded weights leave no GPU memory for the KV cache under "
|
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"Raise --mem-fraction-static above "
|
||||||
f"{suggested_mem_fraction_static:.3f} "
|
f"{suggested_mem_fraction_static:.3f} "
|
||||||
f"(minimum viable = 1 - available/pre = "
|
f"(minimum viable = 1 - available/pre = "
|
||||||
@@ -1586,7 +1594,7 @@ class KVCacheConfigurator:
|
|||||||
return int(rest_memory * (1 << 30)) # return in bytes
|
return int(rest_memory * (1 << 30)) # return in bytes
|
||||||
|
|
||||||
def _calculate_mamba_ratio(self) -> int:
|
def _calculate_mamba_ratio(self) -> int:
|
||||||
if self.server_args.disable_radix_cache:
|
if get_memory().disable_radix_cache:
|
||||||
return 1
|
return 1
|
||||||
|
|
||||||
skip_decode_lock = envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.get()
|
skip_decode_lock = envs.SGLANG_OPT_MAMBA_SKIP_DECODE_LOCK.get()
|
||||||
@@ -1598,7 +1606,7 @@ class KVCacheConfigurator:
|
|||||||
if self.server_args.enable_mamba_extra_buffer():
|
if self.server_args.enable_mamba_extra_buffer():
|
||||||
# ping-pong buffer size is 2 when overlap schedule is on, 1 otherwise.
|
# 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.
|
# 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():
|
if self.server_args.enable_mamba_extra_buffer_lazy():
|
||||||
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY
|
additional_ratio = MAMBA_CACHE_V2_ADDITIONAL_RATIO_OVERLAP_LAZY
|
||||||
else:
|
else:
|
||||||
@@ -1622,7 +1630,7 @@ class KVCacheConfigurator:
|
|||||||
Page alignment is handled by the configurator, not here.
|
Page alignment is handled by the configurator, not here.
|
||||||
If constraints change the value, the configurator re-runs and re-aligns.
|
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
|
# Apply user-specified upper bound
|
||||||
if user_limit is not None:
|
if user_limit is not None:
|
||||||
@@ -1652,7 +1660,7 @@ class KVCacheConfigurator:
|
|||||||
estimated = int(token_capacity / self.model_config.context_len * 512)
|
estimated = int(token_capacity / self.model_config.context_len * 512)
|
||||||
estimated = max(min(estimated, 4096), 2048)
|
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:
|
if max_num_reqs is not None:
|
||||||
requested_per_worker = max_num_reqs // self.ps.attn_dp_size
|
requested_per_worker = max_num_reqs // self.ps.attn_dp_size
|
||||||
max_num_reqs = min(requested_per_worker, token_capacity // 2)
|
max_num_reqs = min(requested_per_worker, token_capacity // 2)
|
||||||
@@ -1663,13 +1671,13 @@ class KVCacheConfigurator:
|
|||||||
if self.mambaish_config is not None:
|
if self.mambaish_config is not None:
|
||||||
ratio = self._calculate_mamba_ratio()
|
ratio = self._calculate_mamba_ratio()
|
||||||
max_num_reqs = min(
|
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:
|
if max_num_reqs <= 0:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"Hybrid (mamba/linear-attention) state cache is too small to serve "
|
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"mamba_ratio={ratio}, resulting max_num_reqs={max_num_reqs}. "
|
||||||
f"Try: (1) reduce --max-running-requests, "
|
f"Try: (1) reduce --max-running-requests, "
|
||||||
f"(2) increase --mem-fraction-static, or "
|
f"(2) increase --mem-fraction-static, or "
|
||||||
@@ -1699,7 +1707,7 @@ class KVCacheConfigurator:
|
|||||||
)
|
)
|
||||||
configurator = create_memory_pool_configurator(self)
|
configurator = create_memory_pool_configurator(self)
|
||||||
config = configurator.finalize_with_max_running_requests(config)
|
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
|
return config
|
||||||
|
|
||||||
def config_from_budget(
|
def config_from_budget(
|
||||||
@@ -1715,14 +1723,14 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
configurator = create_memory_pool_configurator(self)
|
configurator = create_memory_pool_configurator(self)
|
||||||
config = configurator.calculate_pool_sizes(
|
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)
|
max_tokens = self._apply_token_constraints(config.max_total_num_tokens)
|
||||||
if cap_tokens is not None:
|
if cap_tokens is not None:
|
||||||
max_tokens = min(max_tokens, cap_tokens)
|
max_tokens = min(max_tokens, cap_tokens)
|
||||||
if max_tokens != config.max_total_num_tokens:
|
if max_tokens != config.max_total_num_tokens:
|
||||||
config = configurator.calculate_pool_sizes_from_max_tokens(
|
config = configurator.calculate_pool_sizes_from_max_tokens(
|
||||||
max_tokens, self.server_args.page_size
|
max_tokens, get_schedule().page_size
|
||||||
)
|
)
|
||||||
return config
|
return config
|
||||||
|
|
||||||
@@ -1735,13 +1743,14 @@ class KVCacheConfigurator:
|
|||||||
# The ring is allocated per slot but is not part of mamba_cache_per_req;
|
# The ring is allocated per slot but is not part of mamba_cache_per_req;
|
||||||
# the solve must charge it too or num_slots is over-provisioned.
|
# the solve must charge it too or num_slots is over-provisioned.
|
||||||
replayssm_active = (
|
replayssm_active = (
|
||||||
server_args.enable_gdn_replayssm_spec and self.hybrid_gdn_config is not None
|
get_exec().mamba.enable_gdn_replayssm_spec
|
||||||
|
and self.hybrid_gdn_config is not None
|
||||||
)
|
)
|
||||||
if replayssm_active:
|
if replayssm_active:
|
||||||
record_len = (
|
record_len = (
|
||||||
server_args.max_speculative_num_draft_tokens
|
server_args.max_speculative_num_draft_tokens
|
||||||
if server_args.max_speculative_num_draft_tokens is not None
|
if server_args.max_speculative_num_draft_tokens is not None
|
||||||
else server_args.linear_replayssm_cache_len
|
else get_exec().mamba.linear_replayssm_cache_len
|
||||||
)
|
)
|
||||||
replayssm_ring_per_req = (
|
replayssm_ring_per_req = (
|
||||||
config.mamba2_cache_params.replayssm_ring_bytes_per_req(
|
config.mamba2_cache_params.replayssm_ring_bytes_per_req(
|
||||||
@@ -1751,45 +1760,45 @@ class KVCacheConfigurator:
|
|||||||
else:
|
else:
|
||||||
replayssm_ring_per_req = 0
|
replayssm_ring_per_req = 0
|
||||||
if has_spec_dec:
|
if has_spec_dec:
|
||||||
assert server_args.speculative_num_draft_tokens is not None
|
assert get_spec().speculative_num_draft_tokens is not None
|
||||||
assert server_args.max_running_requests is not None
|
assert get_schedule().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
|
# Use explicitly set max_mamba_cache_size
|
||||||
server_args.override(
|
get_context().override(
|
||||||
"mamba_pool.per_dp_shard",
|
"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,
|
// self.ps.attn_dp_size,
|
||||||
)
|
)
|
||||||
# Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
|
# Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
|
||||||
if has_spec_dec and not replayssm_active:
|
if has_spec_dec and not replayssm_active:
|
||||||
ratio = self._calculate_mamba_ratio()
|
ratio = self._calculate_mamba_ratio()
|
||||||
capped_reqs = min(
|
capped_reqs = min(
|
||||||
server_args.max_running_requests // self.ps.attn_dp_size,
|
get_schedule().max_running_requests // self.ps.attn_dp_size,
|
||||||
server_args.max_mamba_cache_size // ratio,
|
get_schedule().max_mamba_cache_size // ratio,
|
||||||
)
|
)
|
||||||
intermediate_size = (
|
intermediate_size = (
|
||||||
config.mamba2_cache_params.mamba_cache_per_req
|
config.mamba2_cache_params.mamba_cache_per_req
|
||||||
* (capped_reqs + 1)
|
* (capped_reqs + 1)
|
||||||
* server_args.speculative_num_draft_tokens
|
* get_spec().speculative_num_draft_tokens
|
||||||
)
|
)
|
||||||
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
||||||
elif (
|
elif (
|
||||||
server_args.disable_radix_cache
|
get_memory().disable_radix_cache
|
||||||
and server_args.max_running_requests is not None
|
and get_schedule().max_running_requests is not None
|
||||||
):
|
):
|
||||||
# Use explicitly set max_running_requests when radix cache is disabled
|
# Use explicitly set max_running_requests when radix cache is disabled
|
||||||
server_args.override(
|
get_context().override(
|
||||||
"mamba_pool.from_max_running_requests",
|
"mamba_pool.from_max_running_requests",
|
||||||
max_mamba_cache_size=server_args.max_running_requests
|
max_mamba_cache_size=get_schedule().max_running_requests
|
||||||
// self.ps.attn_dp_size,
|
// self.ps.attn_dp_size,
|
||||||
)
|
)
|
||||||
# Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
|
# Reserve intermediate memory based on capped max_num_reqs (+1 padding slot)
|
||||||
if has_spec_dec and not replayssm_active:
|
if has_spec_dec and not replayssm_active:
|
||||||
intermediate_size = (
|
intermediate_size = (
|
||||||
config.mamba2_cache_params.mamba_cache_per_req
|
config.mamba2_cache_params.mamba_cache_per_req
|
||||||
* (server_args.max_mamba_cache_size + 1)
|
* (get_schedule().max_mamba_cache_size + 1)
|
||||||
* server_args.speculative_num_draft_tokens
|
* get_spec().speculative_num_draft_tokens
|
||||||
)
|
)
|
||||||
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
||||||
else:
|
else:
|
||||||
@@ -1802,16 +1811,16 @@ class KVCacheConfigurator:
|
|||||||
# (K + 1) * per_req + (K / ratio + 1) * D * per_req = mamba_budget_bytes
|
# (K + 1) * per_req + (K / ratio + 1) * D * per_req = mamba_budget_bytes
|
||||||
mamba_budget = (
|
mamba_budget = (
|
||||||
total_rest_memory
|
total_rest_memory
|
||||||
* server_args.mamba_full_memory_ratio
|
* get_schedule().mamba_full_memory_ratio
|
||||||
/ (1 + server_args.mamba_full_memory_ratio)
|
/ (1 + get_schedule().mamba_full_memory_ratio)
|
||||||
)
|
)
|
||||||
mamba_budget_bytes = mamba_budget * (1 << 30)
|
mamba_budget_bytes = mamba_budget * (1 << 30)
|
||||||
|
|
||||||
if has_spec_dec and not replayssm_active:
|
if has_spec_dec and not replayssm_active:
|
||||||
ratio = self._calculate_mamba_ratio()
|
ratio = self._calculate_mamba_ratio()
|
||||||
D = server_args.speculative_num_draft_tokens
|
D = get_spec().speculative_num_draft_tokens
|
||||||
# Joint solve: main_state + intermediate = mamba_budget
|
# Joint solve: main_state + intermediate = mamba_budget
|
||||||
server_args.override(
|
get_context().override(
|
||||||
"mamba_pool.memory_budget_spec",
|
"mamba_pool.memory_budget_spec",
|
||||||
max_mamba_cache_size=int(
|
max_mamba_cache_size=int(
|
||||||
(mamba_budget_bytes - per_req * (1 + D))
|
(mamba_budget_bytes - per_req * (1 + D))
|
||||||
@@ -1821,14 +1830,14 @@ class KVCacheConfigurator:
|
|||||||
# Intermediate memory is included in mamba_budget, subtract it
|
# Intermediate memory is included in mamba_budget, subtract it
|
||||||
# so the return value only has main_state subtracted from total
|
# so the return value only has main_state subtracted from total
|
||||||
capped_reqs = min(
|
capped_reqs = min(
|
||||||
server_args.max_running_requests // self.ps.attn_dp_size,
|
get_schedule().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 + 1) * D
|
intermediate_size = per_req * (capped_reqs + 1) * D
|
||||||
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
total_rest_memory = total_rest_memory - (intermediate_size / (1 << 30))
|
||||||
else:
|
else:
|
||||||
per_slot = per_req + replayssm_ring_per_req
|
per_slot = per_req + replayssm_ring_per_req
|
||||||
server_args.override(
|
get_context().override(
|
||||||
"mamba_pool.memory_budget",
|
"mamba_pool.memory_budget",
|
||||||
max_mamba_cache_size=int(
|
max_mamba_cache_size=int(
|
||||||
(mamba_budget_bytes - per_slot) // per_slot
|
(mamba_budget_bytes - per_slot) // per_slot
|
||||||
@@ -1839,10 +1848,10 @@ class KVCacheConfigurator:
|
|||||||
# A non-positive value means GPU memory is insufficient for the requested
|
# A non-positive value means GPU memory is insufficient for the requested
|
||||||
# configuration. Fail fast with actionable advice instead of silently
|
# configuration. Fail fast with actionable advice instead of silently
|
||||||
# producing garbled output at runtime.
|
# producing garbled output at runtime.
|
||||||
if server_args.max_mamba_cache_size <= 0:
|
if get_schedule().max_mamba_cache_size <= 0:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"Not enough GPU memory for hybrid (mamba/linear-attention) state cache. "
|
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"(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"mamba_cache_per_req={config.mamba2_cache_params.mamba_cache_per_req / (1 << 20):.2f} MB). "
|
||||||
f"Try: (1) reduce --max-running-requests, "
|
f"Try: (1) reduce --max-running-requests, "
|
||||||
@@ -1853,7 +1862,7 @@ class KVCacheConfigurator:
|
|||||||
|
|
||||||
# +1: the pool's padding slot
|
# +1: the pool's padding slot
|
||||||
mamba_state_memory = (
|
mamba_state_memory = (
|
||||||
(server_args.max_mamba_cache_size + 1)
|
(get_schedule().max_mamba_cache_size + 1)
|
||||||
* (config.mamba2_cache_params.mamba_cache_per_req + replayssm_ring_per_req)
|
* (config.mamba2_cache_params.mamba_cache_per_req + replayssm_ring_per_req)
|
||||||
/ (1 << 30)
|
/ (1 << 30)
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
||||||
from sglang.srt.mem_cache.storage.flexkv.flexkv_connector import FlexKVConnector
|
from sglang.srt.mem_cache.storage.flexkv.flexkv_connector import FlexKVConnector
|
||||||
|
from sglang.srt.runtime_context import get_spec
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.configs.model_config import ModelConfig
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
@@ -393,7 +394,7 @@ class FlexKVRadixCache(RadixCache):
|
|||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_server_args
|
||||||
|
|
||||||
global_server_args = get_server_args()
|
global_server_args = get_server_args()
|
||||||
topk = global_server_args.speculative_eagle_topk
|
topk = get_spec().speculative_eagle_topk
|
||||||
enable_kv_committed_len = topk is None or topk == 1
|
enable_kv_committed_len = topk is None or topk == 1
|
||||||
if enable_kv_committed_len:
|
if enable_kv_committed_len:
|
||||||
kv_committed_len = req.kv_committed_len
|
kv_committed_len = req.kv_committed_len
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ from sglang.srt.mem_cache.base_prefix_cache import (
|
|||||||
MatchResult,
|
MatchResult,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixCache, RadixKey, TreeNode
|
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, get_spec
|
||||||
from sglang.srt.utils import create_device_stream, device_stream_context
|
from sglang.srt.utils import create_device_stream, device_stream_context
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -109,7 +109,7 @@ class LMCRadixCache(RadixCache):
|
|||||||
):
|
):
|
||||||
super().__init__(params)
|
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()
|
kvcache = self.token_to_kv_pool_allocator.get_kvcache()
|
||||||
connector_kwargs = dict(
|
connector_kwargs = dict(
|
||||||
@@ -448,7 +448,7 @@ class LMCRadixCache(RadixCache):
|
|||||||
return
|
return
|
||||||
|
|
||||||
global_server_args = get_server_args()
|
global_server_args = get_server_args()
|
||||||
topk = global_server_args.speculative_eagle_topk
|
topk = get_spec().speculative_eagle_topk
|
||||||
enable_kv_committed_len = topk is None or topk == 1
|
enable_kv_committed_len = topk is None or topk == 1
|
||||||
if enable_kv_committed_len:
|
if enable_kv_committed_len:
|
||||||
kv_committed_len = req.kv_committed_len
|
kv_committed_len = req.kv_committed_len
|
||||||
|
|||||||
@@ -35,7 +35,7 @@ from sglang.srt.mem_cache.unified_cache.components.tree_component import (
|
|||||||
TreeComponent,
|
TreeComponent,
|
||||||
get_and_increase_time_counter,
|
get_and_increase_time_counter,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_exec, get_server_args
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
@@ -66,7 +66,7 @@ class MambaComponent(TreeComponent):
|
|||||||
), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {params.page_size}"
|
), f"MambaComponent requires page_size=1 when mamba_extra_buffer is disabled, got {params.page_size}"
|
||||||
super().__init__(cache, params)
|
super().__init__(cache, params)
|
||||||
self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
|
self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
|
||||||
self.mamba_max_states_per_path = get_server_args().mamba_max_states_per_path
|
self.mamba_max_states_per_path = get_exec().mamba.mamba_max_states_per_path
|
||||||
# HiCache state
|
# HiCache state
|
||||||
self._mamba_pool_host = None # set to host mamba pool when HiCache enabled
|
self._mamba_pool_host = None # set to host mamba pool when HiCache enabled
|
||||||
|
|
||||||
|
|||||||
@@ -37,7 +37,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
|
||||||
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
|
from sglang.srt.model_executor.runner_utils.capture_mode import model_capture_mode
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel
|
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
empty_context,
|
empty_context,
|
||||||
log_info_on_rank0,
|
log_info_on_rank0,
|
||||||
@@ -1027,9 +1027,9 @@ class CPUGraphRunner:
|
|||||||
retrieve_next_token=None,
|
retrieve_next_token=None,
|
||||||
retrieve_next_sibling=None,
|
retrieve_next_sibling=None,
|
||||||
retrieve_cum_len=None,
|
retrieve_cum_len=None,
|
||||||
spec_steps=self.model_runner.server_args.speculative_num_steps,
|
spec_steps=get_spec().speculative_num_steps,
|
||||||
topk=self.model_runner.server_args.speculative_eagle_topk,
|
topk=self.model_runner.server_args.speculative_eagle_topk,
|
||||||
draft_token_num=self.model_runner.server_args.speculative_num_draft_tokens,
|
draft_token_num=get_spec().speculative_num_draft_tokens,
|
||||||
capture_hidden_mode=CaptureHiddenMode.FULL,
|
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||||
seq_lens_sum=None,
|
seq_lens_sum=None,
|
||||||
seq_lens_cpu=None,
|
seq_lens_cpu=None,
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ cuda_graph_config, and the --cuda-graph-config JSON CLI parser.
|
|||||||
|
|
||||||
Module-level imports are pure stdlib — no torch / sglang.srt deps — so
|
Module-level imports are pure stdlib — no torch / sglang.srt deps — so
|
||||||
ServerArgs can import everything here without pulling in backend
|
ServerArgs can import everything here without pulling in backend
|
||||||
classes. check_cuda_graph_backend lazy-imports get_server_args
|
classes. check_cuda_graph_backend lazy-imports the config accessor
|
||||||
inside the function body to preserve that invariant.
|
inside the function body to preserve that invariant.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -179,15 +179,14 @@ def _diff_phase(actual: PhaseConfig, baseline: PhaseConfig) -> Dict[str, Any]:
|
|||||||
|
|
||||||
def check_cuda_graph_backend(phase: str, backend: str) -> bool:
|
def check_cuda_graph_backend(phase: str, backend: str) -> bool:
|
||||||
"""True if cuda_graph_config[phase].backend == backend on the
|
"""True if cuda_graph_config[phase].backend == backend on the
|
||||||
global server args. Returns False if the global server args have not
|
published config. Returns False if the config has not been published
|
||||||
been initialized yet (e.g. unit tests, early startup)."""
|
yet (e.g. unit tests, early startup)."""
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_exec
|
||||||
|
|
||||||
try:
|
try:
|
||||||
server_args = get_server_args()
|
cfg = get_exec().graph.cuda_graph_config
|
||||||
except ValueError:
|
except ValueError:
|
||||||
return False
|
return False
|
||||||
cfg = server_args.cuda_graph_config
|
|
||||||
if cfg is None or phase not in Phase.ALL:
|
if cfg is None or phase not in Phase.ALL:
|
||||||
return False
|
return False
|
||||||
return getattr(cfg, phase).backend == backend
|
return getattr(cfg, phase).backend == backend
|
||||||
|
|||||||
@@ -51,7 +51,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
from sglang.srt.model_executor.forward_batch_deepseek_mha_mixin import (
|
||||||
ForwardBatchDeepSeekMHAMixin,
|
ForwardBatchDeepSeekMHAMixin,
|
||||||
)
|
)
|
||||||
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 (
|
from sglang.srt.utils import (
|
||||||
is_cuda,
|
is_cuda,
|
||||||
is_hip,
|
is_hip,
|
||||||
@@ -965,7 +965,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# --enable-mis: every request must carry delimiter indices (the score
|
# --enable-mis: every request must carry delimiter indices (the score
|
||||||
# endpoint always produces MIS-structured requests; consumers index
|
# endpoint always produces MIS-structured requests; consumers index
|
||||||
# without None-checking).
|
# 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
|
r.multi_item_delimiter_indices is not None for r in batch.reqs
|
||||||
):
|
):
|
||||||
assert all(
|
assert all(
|
||||||
@@ -1134,7 +1134,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# batch_size * [3 * seq_len]
|
# batch_size * [3 * seq_len]
|
||||||
batch_size = self.seq_lens_cpu.shape[0]
|
batch_size = self.seq_lens_cpu.shape[0]
|
||||||
mrope_positions_list = [[]] * batch_size
|
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):
|
for batch_idx in range(batch_size):
|
||||||
mm_input = batch.multimodal_inputs[batch_idx]
|
mm_input = batch.multimodal_inputs[batch_idx]
|
||||||
if self.forward_mode.is_decode():
|
if self.forward_mode.is_decode():
|
||||||
|
|||||||
@@ -162,9 +162,13 @@ from sglang.srt.model_executor.runner import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.platforms import current_platform
|
from sglang.srt.platforms import current_platform
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_context,
|
||||||
|
get_exec,
|
||||||
get_global_dwdp_manager,
|
get_global_dwdp_manager,
|
||||||
|
get_lora,
|
||||||
|
get_model,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_schedule,
|
||||||
set_global_dwdp_manager,
|
set_global_dwdp_manager,
|
||||||
)
|
)
|
||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
@@ -322,7 +326,7 @@ class ModelRunner:
|
|||||||
self.init_threads_binding()
|
self.init_threads_binding()
|
||||||
|
|
||||||
# Set float32 matmul precision
|
# Set float32 matmul precision
|
||||||
if get_server_args().enable_tf32_matmul:
|
if get_exec().features.enable_tf32_matmul:
|
||||||
torch.set_float32_matmul_precision("high")
|
torch.set_float32_matmul_precision("high")
|
||||||
|
|
||||||
# Set device early so that TransferEngine init (e.g. Ascend NPU)
|
# Set device early so that TransferEngine init (e.g. Ascend NPU)
|
||||||
@@ -399,7 +403,7 @@ class ModelRunner:
|
|||||||
|
|
||||||
def _initialize_elastic_ep_joiner(self) -> None:
|
def _initialize_elastic_ep_joiner(self) -> None:
|
||||||
if not (
|
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_scale_joiner
|
and self.server_args.is_ep_scale_joiner
|
||||||
):
|
):
|
||||||
return
|
return
|
||||||
@@ -473,7 +477,7 @@ class ModelRunner:
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
gpu_id=self.gpu_id,
|
gpu_id=self.gpu_id,
|
||||||
model_config=self.model_config,
|
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,
|
get_model=lambda: self.model,
|
||||||
update_model_fields=self.update_model_fields,
|
update_model_fields=self.update_model_fields,
|
||||||
recapture_cuda_graph=self.init_decode_cuda_graph,
|
recapture_cuda_graph=self.init_decode_cuda_graph,
|
||||||
@@ -527,6 +531,7 @@ class ModelRunner:
|
|||||||
model_config=self.model_config,
|
model_config=self.model_config,
|
||||||
server_args=self.server_args,
|
server_args=self.server_args,
|
||||||
kv_cache_dtype=self.kv_cache_dtype,
|
kv_cache_dtype=self.kv_cache_dtype,
|
||||||
|
kv_cache_dtype_str=self.kv_cache_dtype_str,
|
||||||
model_dtype=self.dtype,
|
model_dtype=self.dtype,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
sliding_window_size=self.sliding_window_size,
|
sliding_window_size=self.sliding_window_size,
|
||||||
@@ -550,7 +555,7 @@ class ModelRunner:
|
|||||||
def init_mindspore_runner(self):
|
def init_mindspore_runner(self):
|
||||||
# Init the mindspore runner
|
# Init the mindspore runner
|
||||||
# for now, there is only some communication initialization work
|
# 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
|
from sglang.srt.model_executor.mindspore_runner import init_ms_distributed
|
||||||
|
|
||||||
init_ms_distributed(
|
init_ms_distributed(
|
||||||
@@ -607,7 +612,7 @@ class ModelRunner:
|
|||||||
|
|
||||||
def init_memory_saver_adapter(self):
|
def init_memory_saver_adapter(self):
|
||||||
self.memory_saver_adapter = TorchMemorySaverAdapter.create(
|
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):
|
def maybe_init_remote_instance_transfer_engine(self):
|
||||||
@@ -643,7 +648,7 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def maybe_init_lplb_solvers(self):
|
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)
|
init_lplb_solvers(model_config=self.model_config)
|
||||||
|
|
||||||
def maybe_init_eplb_manager(self):
|
def maybe_init_eplb_manager(self):
|
||||||
@@ -657,12 +662,12 @@ class ModelRunner:
|
|||||||
get_expert_backup_client=lambda: self.expert_backup_client,
|
get_expert_backup_client=lambda: self.expert_backup_client,
|
||||||
get_weight_updater=lambda: self.weight_updater,
|
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
|
else None
|
||||||
)
|
)
|
||||||
|
|
||||||
def maybe_init_elastic_ep(self):
|
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)
|
ElasticEPStateManager.init(self.server_args)
|
||||||
|
|
||||||
def init_token_oracle(self):
|
def init_token_oracle(self):
|
||||||
@@ -681,8 +686,8 @@ class ModelRunner:
|
|||||||
get_model=lambda: self.model,
|
get_model=lambda: self.model,
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
self.server_args.enable_elastic_expert_backup
|
get_exec().moe.enable_elastic_expert_backup
|
||||||
and self.server_args.elastic_ep_backend is not None
|
and get_exec().moe.elastic_ep_backend is not None
|
||||||
)
|
)
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
@@ -691,17 +696,17 @@ class ModelRunner:
|
|||||||
# In layered loading, torchao may have been applied
|
# In layered loading, torchao may have been applied
|
||||||
torchao_applied = getattr(self.model, "torchao_applied", False)
|
torchao_applied = getattr(self.model, "torchao_applied", False)
|
||||||
if not torchao_applied:
|
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)
|
supports_torch_tp = getattr(self.model, "supports_torch_tp", False)
|
||||||
if self.ps.tp_size > 1 and supports_torch_tp:
|
if self.ps.tp_size > 1 and supports_torch_tp:
|
||||||
self.apply_torch_tp()
|
self.apply_torch_tp()
|
||||||
|
|
||||||
def maybe_init_lora_manager(self):
|
def maybe_init_lora_manager(self):
|
||||||
if self.server_args.enable_lora:
|
if get_lora().enable_lora:
|
||||||
self.init_lora_manager()
|
self.init_lora_manager()
|
||||||
|
|
||||||
def maybe_enable_batch_invariant_mode(self):
|
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
|
from sglang.srt.batch_invariant_ops import enable_batch_invariant_mode
|
||||||
|
|
||||||
enable_batch_invariant_mode()
|
enable_batch_invariant_mode()
|
||||||
@@ -996,8 +1001,8 @@ class ModelRunner:
|
|||||||
remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine,
|
remote_instance_weight_transporter_engine=self.remote_instance_weight_transporter.engine,
|
||||||
remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id,
|
remote_instance_weight_transporter_session_id=self.remote_instance_weight_transporter.session_id,
|
||||||
draft_model_idx=self.draft_model_idx,
|
draft_model_idx=self.draft_model_idx,
|
||||||
weight_cache_mode=self.server_args.weight_cache_mode,
|
weight_cache_mode=get_model().weight_cache_mode,
|
||||||
weight_cache_socket=self.server_args.weight_cache_socket,
|
weight_cache_socket=get_model().weight_cache_socket,
|
||||||
)
|
)
|
||||||
|
|
||||||
# If the weight cache is enabled, override the load format to IPC_CACHE
|
# If the weight cache is enabled, override the load format to IPC_CACHE
|
||||||
@@ -1038,7 +1043,7 @@ class ModelRunner:
|
|||||||
get_offloader().post_init()
|
get_offloader().post_init()
|
||||||
|
|
||||||
# Register model for layerwise NVTX profiling if enabled
|
# 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 = PytHooks()
|
||||||
pyt_hooks.register_hooks(self.model, module_prefix="model")
|
pyt_hooks.register_hooks(self.model, module_prefix="model")
|
||||||
|
|
||||||
@@ -1095,7 +1100,7 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
|
|
||||||
dist_barrier_after_load(
|
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,
|
tp_rank=self.ps.tp_rank,
|
||||||
is_ep_joiner=self.server_args.is_ep_joiner,
|
is_ep_joiner=self.server_args.is_ep_joiner,
|
||||||
)
|
)
|
||||||
@@ -1115,16 +1120,16 @@ class ModelRunner:
|
|||||||
self.lora_manager = LoRAManager(
|
self.lora_manager = LoRAManager(
|
||||||
base_model=self.model,
|
base_model=self.model,
|
||||||
base_hf_config=self.model_config.hf_config,
|
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,
|
load_config=self.load_config,
|
||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
server_args=self.server_args,
|
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_size=self.ps.tp_size,
|
||||||
tp_rank=self.ps.tp_rank,
|
tp_rank=self.ps.tp_rank,
|
||||||
max_lora_rank=self.server_args.max_lora_rank,
|
max_lora_rank=get_lora().max_lora_rank,
|
||||||
target_modules=self.server_args.lora_target_modules,
|
target_modules=get_lora().lora_target_modules,
|
||||||
lora_paths=self.server_args.lora_paths,
|
lora_paths=get_lora().lora_paths,
|
||||||
)
|
)
|
||||||
if not cuda_graph_fully_disabled():
|
if not cuda_graph_fully_disabled():
|
||||||
init_lora_cuda_graph_moe_buffers(
|
init_lora_cuda_graph_moe_buffers(
|
||||||
@@ -1157,29 +1162,10 @@ class ModelRunner:
|
|||||||
else:
|
else:
|
||||||
return self.max_total_num_tokens
|
return self.max_total_num_tokens
|
||||||
|
|
||||||
def _record_kv_cache_dtype(self, resolved: str) -> None:
|
|
||||||
# the weight-resolved kv-cache dtype is written to the config
|
|
||||||
# bags via get_context().override, so get_model().kv_cache_dtype readers
|
|
||||||
# see it. server_args stays the pristine RAW record -- configure_kv_cache
|
|
||||||
# _dtype reads it as the resolver INPUT. A draft / mock runner whose
|
|
||||||
# server_args is not the published object keeps the private-bag write.
|
|
||||||
from sglang.srt.runtime_context import get_context
|
|
||||||
|
|
||||||
if get_context()._server_args is self.server_args:
|
|
||||||
get_context().override(
|
|
||||||
"ModelRunner.configure_kv_cache_dtype", kv_cache_dtype=resolved
|
|
||||||
)
|
|
||||||
else:
|
|
||||||
self.server_args.override(
|
|
||||||
"ModelRunner.configure_kv_cache_dtype", kv_cache_dtype=resolved
|
|
||||||
)
|
|
||||||
|
|
||||||
def configure_kv_cache_dtype(self):
|
def configure_kv_cache_dtype(self):
|
||||||
spec_algorithm = getattr(self, "spec_algorithm", None)
|
spec_algorithm = getattr(self, "spec_algorithm", None)
|
||||||
resolved_kv_cache_dtype, self.kv_cache_dtype = (
|
resolved_kv_cache_dtype, self.kv_cache_dtype = (
|
||||||
kv_cache_dtype.configure_kv_cache_dtype(
|
kv_cache_dtype.configure_kv_cache_dtype(
|
||||||
# RAW user intent = resolver INPUT; server_args stays pristine
|
|
||||||
# so read it here -- not the resolved get_model() bag.
|
|
||||||
server_args_kv_cache_dtype=self.server_args.kv_cache_dtype,
|
server_args_kv_cache_dtype=self.server_args.kv_cache_dtype,
|
||||||
model=getattr(self, "model", None),
|
model=getattr(self, "model", None),
|
||||||
model_dtype=getattr(self, "dtype", torch.bfloat16),
|
model_dtype=getattr(self, "dtype", torch.bfloat16),
|
||||||
@@ -1201,8 +1187,6 @@ class ModelRunner:
|
|||||||
if resolved_kv_cache_dtype is not None
|
if resolved_kv_cache_dtype is not None
|
||||||
else self.server_args.kv_cache_dtype
|
else self.server_args.kv_cache_dtype
|
||||||
)
|
)
|
||||||
if resolved_kv_cache_dtype is not None:
|
|
||||||
self._record_kv_cache_dtype(resolved_kv_cache_dtype)
|
|
||||||
|
|
||||||
def _get_attention_backend(self, init_new_workspace: bool = False):
|
def _get_attention_backend(self, init_new_workspace: bool = False):
|
||||||
return get_attention_backend(
|
return get_attention_backend(
|
||||||
@@ -1391,7 +1375,7 @@ class ModelRunner:
|
|||||||
)
|
)
|
||||||
output.expert_distribution_metrics = recorder_outputs.get("metrics")
|
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 (
|
if (
|
||||||
not self.is_draft_worker
|
not self.is_draft_worker
|
||||||
and (experts_capturer := get_global_experts_capturer()) is not None
|
and (experts_capturer := get_global_experts_capturer()) is not None
|
||||||
@@ -1421,7 +1405,7 @@ class ModelRunner:
|
|||||||
self.msprobe_debugger.stop()
|
self.msprobe_debugger.stop()
|
||||||
self.msprobe_debugger.step()
|
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()
|
self.maybe_join_ep_ranks()
|
||||||
|
|
||||||
return output
|
return output
|
||||||
@@ -1852,7 +1836,7 @@ class ModelRunner:
|
|||||||
local_timeout = (
|
local_timeout = (
|
||||||
state.pending_since is not None
|
state.pending_since is not None
|
||||||
and time.monotonic() - state.pending_since
|
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))
|
timeout = state.active_ranks.new_tensor(int(local_timeout))
|
||||||
dist.all_reduce(timeout, op=dist.ReduceOp.MAX, group=dist.group.WORLD)
|
dist.all_reduce(timeout, op=dist.ReduceOp.MAX, group=dist.group.WORLD)
|
||||||
@@ -1922,7 +1906,7 @@ class ModelRunner:
|
|||||||
load_config: LoadConfig,
|
load_config: LoadConfig,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.model = new_model
|
self.model = new_model
|
||||||
self.server_args.override(
|
get_context().override(
|
||||||
"model_runner.update_model_fields",
|
"model_runner.update_model_fields",
|
||||||
model_path=model_path,
|
model_path=model_path,
|
||||||
load_format=load_format,
|
load_format=load_format,
|
||||||
|
|||||||
+3
-2
@@ -11,6 +11,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
|||||||
RemoteInstanceWeightLoaderBackend,
|
RemoteInstanceWeightLoaderBackend,
|
||||||
register_memory_region,
|
register_memory_region,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import get_model
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
from sglang.srt.utils.network import NetworkAddress, get_local_ip_auto
|
||||||
|
|
||||||
@@ -58,7 +59,7 @@ class RemoteInstanceWeightTransporter:
|
|||||||
# ModelExpress owns TransferEngine memory registration and metadata
|
# ModelExpress owns TransferEngine memory registration and metadata
|
||||||
# publishing for backend=modelexpress. Re-registering here would
|
# publishing for backend=modelexpress. Re-registering here would
|
||||||
# overlap the same weight buffers.
|
# overlap the same weight buffers.
|
||||||
and self.server_args.remote_instance_weight_loader_backend
|
and get_model().remote_instance_weight_loader_backend
|
||||||
!= RemoteInstanceWeightLoaderBackend.MODELEXPRESS
|
!= RemoteInstanceWeightLoaderBackend.MODELEXPRESS
|
||||||
and self.engine is not None
|
and self.engine is not None
|
||||||
and self.weight_info is None
|
and self.weight_info is None
|
||||||
@@ -84,7 +85,7 @@ class RemoteInstanceWeightTransporter:
|
|||||||
else:
|
else:
|
||||||
bootstrap_host = "127.0.0.1"
|
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)
|
bootstrap_na = NetworkAddress(bootstrap_host, bootstrap_port)
|
||||||
url = f"{bootstrap_na.to_url()}/register_transfer_engine_info"
|
url = f"{bootstrap_na.to_url()}/register_transfer_engine_info"
|
||||||
|
|
||||||
|
|||||||
@@ -33,7 +33,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.mem_cache.allocation_sizing import get_alloc_len_per_decode
|
from sglang.srt.mem_cache.allocation_sizing import get_alloc_len_per_decode
|
||||||
from sglang.srt.mem_cache.deepseek_v4_memory_pool import get_compress_state_ring_size
|
from sglang.srt.mem_cache.deepseek_v4_memory_pool import get_compress_state_ring_size
|
||||||
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import DSATokenToKVPool
|
||||||
from sglang.srt.runtime_context import get_model, get_parallel
|
from sglang.srt.runtime_context import get_parallel
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
ceil_align,
|
ceil_align,
|
||||||
ceil_div,
|
ceil_div,
|
||||||
@@ -119,6 +119,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, kvc: KVCacheConfigurator):
|
def __init__(self, kvc: KVCacheConfigurator):
|
||||||
|
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
|
||||||
# Determine effective number of layers for KV cache
|
# Determine effective number of layers for KV cache
|
||||||
if mambaish := mambaish_config(kvc.model_config):
|
if mambaish := mambaish_config(kvc.model_config):
|
||||||
effective_layer_ids = [
|
effective_layer_ids = [
|
||||||
@@ -304,7 +305,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
)
|
)
|
||||||
# FP4 prefill uses one shared FP8 dequant workspace across layers.
|
# FP4 prefill uses one shared FP8 dequant workspace across layers.
|
||||||
cell_size += n * k * 2 * kv_size
|
cell_size += n * k * 2 * kv_size
|
||||||
elif get_model().kv_cache_dtype == "mxfp8":
|
elif self.kv_cache_dtype_str == "mxfp8":
|
||||||
scale_block_size = 32
|
scale_block_size = 32
|
||||||
n = model_config.get_num_kv_heads(tp_size)
|
n = model_config.get_num_kv_heads(tp_size)
|
||||||
cell_size += (
|
cell_size += (
|
||||||
@@ -339,6 +340,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, kvc: KVCacheConfigurator):
|
def __init__(self, kvc: KVCacheConfigurator):
|
||||||
|
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
|
||||||
model_config = kvc.model_config
|
model_config = kvc.model_config
|
||||||
kv_cache_dtype = kvc.kv_cache_dtype
|
kv_cache_dtype = kvc.kv_cache_dtype
|
||||||
kv_size = torch._utils._element_size(kv_cache_dtype)
|
kv_size = torch._utils._element_size(kv_cache_dtype)
|
||||||
@@ -368,7 +370,7 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
|
|||||||
* kv_size
|
* kv_size
|
||||||
)
|
)
|
||||||
|
|
||||||
if get_model().kv_cache_dtype == "mxfp8":
|
if self.kv_cache_dtype_str == "mxfp8":
|
||||||
scale_block_size = 32
|
scale_block_size = 32
|
||||||
self._full_per_token += (
|
self._full_per_token += (
|
||||||
model_config.get_num_kv_heads(tp_size)
|
model_config.get_num_kv_heads(tp_size)
|
||||||
@@ -501,6 +503,7 @@ class SWAChunkCapPoolConfigurator(HybridSWAPoolConfigurator):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, kvc: KVCacheConfigurator):
|
def __init__(self, kvc: KVCacheConfigurator):
|
||||||
|
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
|
||||||
super().__init__(kvc)
|
super().__init__(kvc)
|
||||||
assert self._full_layers_num > 0
|
assert self._full_layers_num > 0
|
||||||
|
|
||||||
@@ -613,6 +616,7 @@ class DSV4PoolConfigurator(MemoryPoolConfigurator):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, kvc: KVCacheConfigurator):
|
def __init__(self, kvc: KVCacheConfigurator):
|
||||||
|
self.kv_cache_dtype_str = kvc.kv_cache_dtype_str
|
||||||
cfg = kvc.model_config
|
cfg = kvc.model_config
|
||||||
self.qk_nope_head_dim = cfg.qk_nope_head_dim
|
self.qk_nope_head_dim = cfg.qk_nope_head_dim
|
||||||
self.qk_rope_head_dim = cfg.qk_rope_head_dim
|
self.qk_rope_head_dim = cfg.qk_rope_head_dim
|
||||||
|
|||||||
@@ -91,7 +91,7 @@ from sglang.srt.model_executor.runner_utils.deepep_adapter import (
|
|||||||
DeepEPCudaGraphRunnerAdapter,
|
DeepEPCudaGraphRunnerAdapter,
|
||||||
)
|
)
|
||||||
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
from sglang.srt.multiplex.pdmux_context import get_current_stream_idx, get_stream_groups
|
||||||
from sglang.srt.runtime_context import get_flags, get_parallel
|
from sglang.srt.runtime_context import get_flags, get_parallel, get_spec
|
||||||
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
from sglang.srt.speculative.ragged_verify import resolve_ragged_verify_layout
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
empty_context,
|
empty_context,
|
||||||
@@ -246,12 +246,12 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
self.is_dllm = self.dllm_config is not None
|
self.is_dllm = self.dllm_config is not None
|
||||||
self.attn_backend = attn_backend or model_runner.attn_backend
|
self.attn_backend = attn_backend or model_runner.attn_backend
|
||||||
self.speculative_num_steps = (
|
self.speculative_num_steps = (
|
||||||
model_runner.server_args.speculative_num_steps
|
get_spec().speculative_num_steps
|
||||||
if speculative_num_steps is None
|
if speculative_num_steps is None
|
||||||
else speculative_num_steps
|
else speculative_num_steps
|
||||||
)
|
)
|
||||||
self.speculative_num_draft_tokens = (
|
self.speculative_num_draft_tokens = (
|
||||||
model_runner.server_args.speculative_num_draft_tokens
|
get_spec().speculative_num_draft_tokens
|
||||||
if speculative_num_draft_tokens is None
|
if speculative_num_draft_tokens is None
|
||||||
else speculative_num_draft_tokens
|
else speculative_num_draft_tokens
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -47,7 +47,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
|||||||
get_remote_instance_transfer_engine_info_per_rank,
|
get_remote_instance_transfer_engine_info_per_rank,
|
||||||
register_memory_region,
|
register_memory_region,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_exec, get_model, get_server_args
|
||||||
from sglang.srt.utils import get_available_gpu_memory
|
from sglang.srt.utils import get_available_gpu_memory
|
||||||
|
|
||||||
# Try to import accelerate (optional dependency)
|
# Try to import accelerate (optional dependency)
|
||||||
@@ -495,10 +495,10 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
hf_folder = model_name_or_path
|
hf_folder = model_name_or_path
|
||||||
|
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
if server_args and server_args.model_checksum is not None:
|
if server_args and get_model().model_checksum is not None:
|
||||||
from sglang.srt.utils.model_file_verifier import verify
|
from sglang.srt.utils.model_file_verifier import verify
|
||||||
|
|
||||||
checksums_source = server_args.model_checksum or model_name_or_path
|
checksums_source = get_model().model_checksum or model_name_or_path
|
||||||
verify(model_path=hf_folder, checksums_source=checksums_source)
|
verify(model_path=hf_folder, checksums_source=checksums_source)
|
||||||
|
|
||||||
hf_weights_files: List[str] = []
|
hf_weights_files: List[str] = []
|
||||||
@@ -581,11 +581,11 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
)
|
)
|
||||||
elif use_safetensors:
|
elif use_safetensors:
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
weight_loader_disable_mmap = server_args.weight_loader_disable_mmap
|
weight_loader_disable_mmap = get_model().weight_loader_disable_mmap
|
||||||
weight_loader_prefetch = server_args.weight_loader_prefetch_checkpoints
|
weight_loader_prefetch = get_model().weight_loader_prefetch_checkpoints
|
||||||
prefetch_num_threads = server_args.weight_loader_prefetch_num_threads
|
prefetch_num_threads = get_model().weight_loader_prefetch_num_threads
|
||||||
weight_loader_drop_cache_after_load = (
|
weight_loader_drop_cache_after_load = (
|
||||||
server_args.weight_loader_drop_cache_after_load
|
get_model().weight_loader_drop_cache_after_load
|
||||||
)
|
)
|
||||||
|
|
||||||
# Prefetch and multi-threaded loading both read the same shards,
|
# Prefetch and multi-threaded loading both read the same shards,
|
||||||
@@ -879,9 +879,8 @@ class LayeredModelLoader(DefaultModelLoader):
|
|||||||
device_config: DeviceConfig,
|
device_config: DeviceConfig,
|
||||||
) -> nn.Module:
|
) -> nn.Module:
|
||||||
from sglang.srt.layers.torchao_utils import apply_torchao_config_to_model
|
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)
|
target_device = torch.device(device_config.device)
|
||||||
quant_config = _get_quantization_config(model_config, self.load_config)
|
quant_config = _get_quantization_config(model_config, self.load_config)
|
||||||
|
|
||||||
@@ -1751,13 +1750,13 @@ class PreshardedModelLoader(DefaultModelLoader):
|
|||||||
"moe_dense_tp_size": server_args.moe_dense_tp_size,
|
"moe_dense_tp_size": server_args.moe_dense_tp_size,
|
||||||
"moe_dp_size": server_args.moe_dp_size,
|
"moe_dp_size": server_args.moe_dp_size,
|
||||||
"enable_dp_lm_head": server_args.enable_dp_lm_head,
|
"enable_dp_lm_head": server_args.enable_dp_lm_head,
|
||||||
"enable_fp32_lm_head": server_args.enable_fp32_lm_head,
|
"enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head,
|
||||||
"quantization": model_config.quantization,
|
"quantization": model_config.quantization,
|
||||||
"model_dtype": str(model_config.dtype),
|
"model_dtype": str(model_config.dtype),
|
||||||
"ep_num_redundant_experts": server_args.ep_num_redundant_experts,
|
"ep_num_redundant_experts": get_exec().moe.ep_num_redundant_experts,
|
||||||
"enable_eplb": server_args.enable_eplb,
|
"enable_eplb": get_exec().moe.enable_eplb,
|
||||||
"init_expert_location": self._normalize_init_expert_location(
|
"init_expert_location": self._normalize_init_expert_location(
|
||||||
server_args.init_expert_location
|
get_exec().moe.init_expert_location
|
||||||
),
|
),
|
||||||
"structural_signature": self._compute_structural_signature(model_config),
|
"structural_signature": self._compute_structural_signature(model_config),
|
||||||
}
|
}
|
||||||
@@ -3934,10 +3933,10 @@ class RunaiModelStreamerLoader(BaseModelLoader):
|
|||||||
)
|
)
|
||||||
|
|
||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
if server_args and server_args.model_checksum is not None:
|
if server_args and get_model().model_checksum is not None:
|
||||||
from sglang.srt.utils.model_file_verifier import verify
|
from sglang.srt.utils.model_file_verifier import verify
|
||||||
|
|
||||||
checksums_source = server_args.model_checksum or model_name_or_path
|
checksums_source = get_model().model_checksum or model_name_or_path
|
||||||
verify(model_path=hf_folder, checksums_source=checksums_source)
|
verify(model_path=hf_folder, checksums_source=checksums_source)
|
||||||
|
|
||||||
hf_weights_files = list_safetensors(path=hf_folder)
|
hf_weights_files = list_safetensors(path=hf_folder)
|
||||||
|
|||||||
@@ -78,6 +78,7 @@ from sglang.srt.models.utils import (
|
|||||||
enable_fused_set_kv_buffer,
|
enable_fused_set_kv_buffer,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
get_forward,
|
get_forward,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
@@ -209,7 +210,7 @@ class BailingMoESparseMoeBlock(nn.Module):
|
|||||||
self.router_dtype = torch.bfloat16
|
self.router_dtype = torch.bfloat16
|
||||||
|
|
||||||
# TODO global_server_args.ep_num_redundant_experts is used for eplb, not supported now
|
# 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
|
# check group topk
|
||||||
self.num_expert_group = getattr(config, "n_group", 0)
|
self.num_expert_group = getattr(config, "n_group", 0)
|
||||||
self.topk_group = getattr(config, "topk_group", 0)
|
self.topk_group = getattr(config, "topk_group", 0)
|
||||||
@@ -223,9 +224,7 @@ class BailingMoESparseMoeBlock(nn.Module):
|
|||||||
self.num_expert_group = self.topk_group = None
|
self.num_expert_group = self.topk_group = None
|
||||||
self.use_grouped_topk = False
|
self.use_grouped_topk = False
|
||||||
|
|
||||||
self.num_experts = (
|
self.num_experts = config.num_experts + get_exec().moe.ep_num_redundant_experts
|
||||||
config.num_experts + get_server_args().ep_num_redundant_experts
|
|
||||||
)
|
|
||||||
|
|
||||||
self.gate = BailingMoEGate(
|
self.gate = BailingMoEGate(
|
||||||
config=config,
|
config=config,
|
||||||
|
|||||||
@@ -59,6 +59,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.deepseek_v2 import DeepseekV2AttentionMLA, DeepseekV2MLP, _is_hip
|
||||||
from sglang.srt.models.utils import WeightsMapper
|
from sglang.srt.models.utils import WeightsMapper
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
get_device,
|
||||||
get_forward,
|
get_forward,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
@@ -529,7 +530,7 @@ class BailingMoELinearAttention(nn.Module):
|
|||||||
base=self.rope_theta,
|
base=self.rope_theta,
|
||||||
rope_scaling=config.rope_scaling,
|
rope_scaling=config.rope_scaling,
|
||||||
is_neox_style=True,
|
is_neox_style=True,
|
||||||
device=get_server_args().device,
|
device=get_device().device,
|
||||||
dtype=torch.float32,
|
dtype=torch.float32,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -690,7 +691,7 @@ class BailingMoEAttention(nn.Module):
|
|||||||
max_position=self.max_position_embeddings,
|
max_position=self.max_position_embeddings,
|
||||||
base=self.rope_theta,
|
base=self.rope_theta,
|
||||||
rope_scaling=config.rope_scaling,
|
rope_scaling=config.rope_scaling,
|
||||||
device=get_server_args().device,
|
device=get_device().device,
|
||||||
)
|
)
|
||||||
self.attn = RadixAttention(
|
self.attn = RadixAttention(
|
||||||
self.num_heads,
|
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.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
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
|
from sglang.srt.utils import add_prefix
|
||||||
|
|
||||||
BertConfig = None
|
BertConfig = None
|
||||||
@@ -365,9 +365,7 @@ class BertModel(nn.Module):
|
|||||||
quant_config=quant_config,
|
quant_config=quant_config,
|
||||||
prefix=add_prefix("encoder", prefix),
|
prefix=add_prefix("encoder", prefix),
|
||||||
)
|
)
|
||||||
pooling_type = (
|
pooling_type = PoolingType.CLS if get_model().is_embedding else PoolingType.LAST
|
||||||
PoolingType.CLS if get_server_args().is_embedding else PoolingType.LAST
|
|
||||||
)
|
|
||||||
self.pooler = (
|
self.pooler = (
|
||||||
BertPooler(config)
|
BertPooler(config)
|
||||||
if self.use_bert_pooler
|
if self.use_bert_pooler
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_methods
|
|||||||
AttnForwardMethod,
|
AttnForwardMethod,
|
||||||
)
|
)
|
||||||
from sglang.srt.models.deepseek_common.utils import _is_hip
|
from sglang.srt.models.deepseek_common.utils import _is_hip
|
||||||
from sglang.srt.runtime_context import get_exec, get_server_args
|
from sglang.srt.runtime_context import get_exec
|
||||||
from sglang.srt.utils import is_sm100_or_sm110_supported, use_intel_amx_backend
|
from sglang.srt.utils import is_sm100_or_sm110_supported, use_intel_amx_backend
|
||||||
|
|
||||||
MHA_ONE_SHOT_SUPPORTED_BACKENDS = ["fa3", "flashinfer", "flashmla"]
|
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):
|
def handle_attention_fa3(attn, forward_batch):
|
||||||
# when deterministic inference is enabled, use 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)
|
return _dispatch_mla_subtype(attn, forward_batch)
|
||||||
else:
|
else:
|
||||||
return _handle_attention_backend(attn, forward_batch, "fa3")
|
return _handle_attention_backend(attn, forward_batch, "fa3")
|
||||||
@@ -194,7 +194,7 @@ def handle_attention_triton(attn, forward_batch):
|
|||||||
return AttnForwardMethod.MLA
|
return AttnForwardMethod.MLA
|
||||||
|
|
||||||
# when deterministic inference is enabled, use 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)
|
return _dispatch_mla_subtype(attn, forward_batch)
|
||||||
|
|
||||||
if (
|
if (
|
||||||
|
|||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user