config: business code no longer reads the published ServerArgs (#34081)
This commit is contained in:
@@ -49,7 +49,9 @@ from sglang.srt.hardware_backend.mlx.kv_cache import (
|
|||||||
uses_sliding_window_attention,
|
uses_sliding_window_attention,
|
||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
mamba_cache_chunk_size,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -281,7 +283,7 @@ class MlxModelRunner:
|
|||||||
):
|
):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
chunk_size = get_server_args().mamba_cache_chunk_size
|
chunk_size = mamba_cache_chunk_size()
|
||||||
track_len = prefix_len + (new_token_count // chunk_size) * chunk_size
|
track_len = prefix_len + (new_token_count // chunk_size) * chunk_size
|
||||||
branching_len = getattr(req, "mamba_branching_seqlen", None)
|
branching_len = getattr(req, "mamba_branching_seqlen", None)
|
||||||
if (
|
if (
|
||||||
|
|||||||
@@ -40,11 +40,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
|
|||||||
is_in_tc_piecewise_cuda_graph,
|
is_in_tc_piecewise_cuda_graph,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_pp_size,
|
||||||
get_device,
|
get_device,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
get_server_args,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.state_capturer.indexer_topk import (
|
from sglang.srt.state_capturer.indexer_topk import (
|
||||||
maybe_capture_indexer_topk,
|
maybe_capture_indexer_topk,
|
||||||
@@ -249,7 +249,7 @@ class Indexer(DSANPUIndexerMixin, BaseFusedOp):
|
|||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
self.sm_count = deep_gemm.get_num_sms()
|
self.sm_count = deep_gemm.get_num_sms()
|
||||||
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8)
|
self.half_device_sm_count = ceil_align(self.sm_count // 2, 8)
|
||||||
pp_size = get_server_args().pp_size
|
pp_size = configured_pp_size()
|
||||||
self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank
|
self.logits_with_pp_recv = pp_size > 1 and not get_pp_group().is_last_rank
|
||||||
else:
|
else:
|
||||||
self.logits_with_pp_recv = False
|
self.logits_with_pp_recv = False
|
||||||
|
|||||||
@@ -12,7 +12,10 @@ from sglang.srt.model_executor.runner_backend_utils.breakable_cuda_graph 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_parallel, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
process_model_config,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip
|
from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip
|
||||||
from sglang.srt.utils.common import ceil_align, ceil_div
|
from sglang.srt.utils.common import ceil_align, ceil_div
|
||||||
|
|
||||||
@@ -113,7 +116,7 @@ def is_dsa_enable_prefill_cp():
|
|||||||
return False
|
return False
|
||||||
from sglang.srt.configs.model_config import is_deepseek_dsa, is_deepseek_v4
|
from sglang.srt.configs.model_config import is_deepseek_dsa, is_deepseek_v4
|
||||||
|
|
||||||
hf_config = get_server_args().get_model_config().hf_config
|
hf_config = process_model_config().hf_config
|
||||||
return is_deepseek_dsa(hf_config) or is_deepseek_v4(hf_config)
|
return is_deepseek_dsa(hf_config) or is_deepseek_v4(hf_config)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -25,7 +25,11 @@ 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_exec, get_memory, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_memory,
|
||||||
|
mamba_cache_chunk_size,
|
||||||
|
)
|
||||||
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
|
||||||
|
|
||||||
@@ -297,8 +301,8 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
lens_to_track = (
|
lens_to_track = (
|
||||||
forward_batch.mamba_track_seqlens - forward_batch.extend_prefix_lens
|
forward_batch.mamba_track_seqlens - forward_batch.extend_prefix_lens
|
||||||
)
|
)
|
||||||
mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
|
chunk_size = mamba_cache_chunk_size()
|
||||||
aligned_len = (lens_to_track // mamba_cache_chunk_size) * mamba_cache_chunk_size
|
aligned_len = (lens_to_track // chunk_size) * chunk_size
|
||||||
start_indices = query_start_loc[:-1] + aligned_len - conv_state_len
|
start_indices = query_start_loc[:-1] + aligned_len - conv_state_len
|
||||||
start_indices = start_indices[forward_batch.mamba_track_mask]
|
start_indices = start_indices[forward_batch.mamba_track_mask]
|
||||||
|
|
||||||
@@ -316,7 +320,7 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
"""src/dst indices to track SSM states for prefix caching: aligned seqs
|
"""src/dst indices to track SSM states for prefix caching: aligned seqs
|
||||||
cache last_recurrent_state, unaligned cache intermediate `h` at the last
|
cache last_recurrent_state, unaligned cache intermediate `h` at the last
|
||||||
chunk boundary."""
|
chunk boundary."""
|
||||||
mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
|
chunk_size = mamba_cache_chunk_size()
|
||||||
# CPU to avoid kernel launches for the masking ops
|
# CPU to avoid kernel launches for the masking ops
|
||||||
mamba_track_mask = forward_batch.mamba_track_mask.cpu()
|
mamba_track_mask = forward_batch.mamba_track_mask.cpu()
|
||||||
extend_seq_lens = forward_batch.extend_seq_lens.cpu()
|
extend_seq_lens = forward_batch.extend_seq_lens.cpu()
|
||||||
@@ -326,9 +330,9 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
prefix_lens = forward_batch.extend_prefix_lens.cpu()
|
prefix_lens = forward_batch.extend_prefix_lens.cpu()
|
||||||
|
|
||||||
if isinstance(self, Mamba2AttnBackend):
|
if isinstance(self, Mamba2AttnBackend):
|
||||||
num_h_states = extend_seq_lens // mamba_cache_chunk_size
|
num_h_states = extend_seq_lens // chunk_size
|
||||||
else:
|
else:
|
||||||
num_h_states = (extend_seq_lens - 1) // mamba_cache_chunk_size + 1
|
num_h_states = (extend_seq_lens - 1) // chunk_size + 1
|
||||||
|
|
||||||
track_ssm_src_offset = torch.zeros_like(num_h_states)
|
track_ssm_src_offset = torch.zeros_like(num_h_states)
|
||||||
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
|
track_ssm_src_offset[1:] = torch.cumsum(num_h_states[:-1], dim=0)
|
||||||
@@ -338,17 +342,17 @@ class MambaAttnBackendBase(AttentionBackend):
|
|||||||
offset_masked = track_ssm_src_offset[mamba_track_mask]
|
offset_masked = track_ssm_src_offset[mamba_track_mask]
|
||||||
dst_masked = mamba_track_indices[mamba_track_mask]
|
dst_masked = mamba_track_indices[mamba_track_mask]
|
||||||
|
|
||||||
is_aligned = (lens_masked % mamba_cache_chunk_size) == 0
|
is_aligned = (lens_masked % chunk_size) == 0
|
||||||
|
|
||||||
# Aligned: last_recurrent_state from ssm_states.
|
# Aligned: last_recurrent_state from ssm_states.
|
||||||
track_ssm_final_src = mamba_cache_indices[mamba_track_mask][is_aligned]
|
track_ssm_final_src = mamba_cache_indices[mamba_track_mask][is_aligned]
|
||||||
track_ssm_final_dst = dst_masked[is_aligned]
|
track_ssm_final_dst = dst_masked[is_aligned]
|
||||||
|
|
||||||
# Unaligned: intermediate state from h.
|
# Unaligned: intermediate state from h.
|
||||||
# TODO: handle mamba_cache_chunk_size % page size != 0
|
# TODO: handle chunk_size % page size != 0
|
||||||
not_aligned = ~is_aligned
|
not_aligned = ~is_aligned
|
||||||
track_ssm_h_src = offset_masked[not_aligned] + (
|
track_ssm_h_src = offset_masked[not_aligned] + (
|
||||||
lens_masked[not_aligned] // mamba_cache_chunk_size
|
lens_masked[not_aligned] // chunk_size
|
||||||
)
|
)
|
||||||
track_ssm_h_dst = dst_masked[not_aligned]
|
track_ssm_h_dst = dst_masked[not_aligned]
|
||||||
|
|
||||||
|
|||||||
@@ -62,7 +62,11 @@ 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_exec, get_server_args, get_spec
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
get_spec,
|
||||||
|
mamba_cache_chunk_size,
|
||||||
|
)
|
||||||
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -99,7 +103,7 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
|
|||||||
# [n_layers, n_slots, conv_kernel - 1, conv_dim].
|
# [n_layers, n_slots, conv_kernel - 1, conv_dim].
|
||||||
self._mamba_cache = self.req_to_token_pool.mamba_pool.mamba_cache
|
self._mamba_cache = self.req_to_token_pool.mamba_pool.mamba_cache
|
||||||
self.conv_state_len: int = self.conv_states_shape[2]
|
self.conv_state_len: int = self.conv_states_shape[2]
|
||||||
self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
|
self.mamba_cache_chunk_size = mamba_cache_chunk_size()
|
||||||
# A plain table lookup is recordable; the unified pool's translate is an
|
# A plain table lookup is recordable; the unified pool's translate is an
|
||||||
# allocator lookup and must stay in the out-of-graph replay prep.
|
# allocator lookup and must stay in the out-of-graph replay prep.
|
||||||
self._slot_gather_recordable = (
|
self._slot_gather_recordable = (
|
||||||
|
|||||||
@@ -19,7 +19,9 @@ import torch
|
|||||||
from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
|
from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
|
||||||
LinearAttnKernelBase,
|
LinearAttnKernelBase,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
mamba_cache_chunk_size,
|
||||||
|
)
|
||||||
from sglang.srt.utils import is_cuda
|
from sglang.srt.utils import is_cuda
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -50,7 +52,7 @@ def maybe_build_flashinfer_checkpoint_plan(
|
|||||||
):
|
):
|
||||||
return
|
return
|
||||||
|
|
||||||
checkpoint_every_n_tokens = get_server_args().mamba_cache_chunk_size
|
checkpoint_every_n_tokens = mamba_cache_chunk_size()
|
||||||
extend_seq_lens = forward_batch.extend_seq_lens.to(device="cpu", dtype=torch.int64)
|
extend_seq_lens = forward_batch.extend_seq_lens.to(device="cpu", dtype=torch.int64)
|
||||||
track_mask = forward_batch.mamba_track_mask.to(device="cpu", dtype=torch.bool)
|
track_mask = forward_batch.mamba_track_mask.to(device="cpu", dtype=torch.bool)
|
||||||
relative_track_lens = forward_batch.mamba_track_seqlens.to(
|
relative_track_lens = forward_batch.mamba_track_seqlens.to(
|
||||||
|
|||||||
@@ -37,6 +37,7 @@ from sglang.srt.layers.attention.linear.kernels.kda_triton import TritonKDAKerne
|
|||||||
from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
|
from sglang.srt.layers.attention.linear.kernels.kernel_backend import (
|
||||||
LinearAttnKernelBase,
|
LinearAttnKernelBase,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.runtime_context import mamba_cache_chunk_size
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -194,11 +195,8 @@ class PtxKDAKernel(LinearAttnKernelBase):
|
|||||||
# which may be larger (for example 256 for Nemotron-H). Returning
|
# which may be larger (for example 256 for Nemotron-H). Returning
|
||||||
# the raw 64-token rows in that case would silently select the
|
# the raw 64-token rows in that case would silently select the
|
||||||
# wrong boundary and compute wrong offsets for later sequences.
|
# wrong boundary and compute wrong offsets for later sequences.
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
intermediate_stride_supported = (
|
intermediate_stride_supported = mamba_cache_chunk_size() == _CHUNK
|
||||||
get_server_args().mamba_cache_chunk_size == _CHUNK
|
|
||||||
)
|
|
||||||
seq_lens = (
|
seq_lens = (
|
||||||
[int(length) for length in seq_lens_cpu] if seq_lens_cpu is not None else []
|
[int(length) for length in seq_lens_cpu] if seq_lens_cpu is not None else []
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -15,7 +15,10 @@ from sglang.srt.layers.attention.linear.linear_metadata import (
|
|||||||
from sglang.srt.layers.radix_attention import RadixAttention
|
from sglang.srt.layers.radix_attention import RadixAttention
|
||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
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_parallel, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
mamba_cache_chunk_size,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -299,7 +302,7 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
if h_dst is None or h_dst.numel() == 0:
|
if h_dst is None or h_dst.numel() == 0:
|
||||||
return None, None
|
return None, None
|
||||||
|
|
||||||
mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
|
chunk_size = mamba_cache_chunk_size()
|
||||||
num_prefills = metadata.num_prefills
|
num_prefills = metadata.num_prefills
|
||||||
track_mask = forward_batch.mamba_track_mask[:num_prefills]
|
track_mask = forward_batch.mamba_track_mask[:num_prefills]
|
||||||
extend_lens = forward_batch.extend_seq_lens[:num_prefills]
|
extend_lens = forward_batch.extend_seq_lens[:num_prefills]
|
||||||
@@ -307,9 +310,7 @@ class LightningAttentionBackend(MambaAttnBackendBase):
|
|||||||
track_seqlens = forward_batch.mamba_track_seqlens[:num_prefills]
|
track_seqlens = forward_batch.mamba_track_seqlens[:num_prefills]
|
||||||
|
|
||||||
lens_to_track = track_seqlens - prefix_lens
|
lens_to_track = track_seqlens - prefix_lens
|
||||||
boundary_lens = (
|
boundary_lens = (lens_to_track // chunk_size) * chunk_size
|
||||||
lens_to_track // mamba_cache_chunk_size
|
|
||||||
) * mamba_cache_chunk_size
|
|
||||||
track_rows = (track_mask & (boundary_lens < extend_lens)).nonzero(
|
track_rows = (track_mask & (boundary_lens < extend_lens)).nonzero(
|
||||||
as_tuple=True
|
as_tuple=True
|
||||||
)[0]
|
)[0]
|
||||||
|
|||||||
@@ -230,12 +230,15 @@ class ContextParallelStrategy(ABC):
|
|||||||
|
|
||||||
|
|
||||||
def _is_dsa_active() -> bool:
|
def _is_dsa_active() -> bool:
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import get_parallel, get_server_args
|
||||||
|
|
||||||
sa = get_server_args()
|
# `_is_dsa_model_arch` is set nowhere in the tree, so this predicate is
|
||||||
|
# inert today (the getattr default makes it False). Kept verbatim rather
|
||||||
|
# than "fixed" here, because deciding what it should name is the CP path's
|
||||||
|
# call; the ratchet exempts it with that reason.
|
||||||
return bool(
|
return bool(
|
||||||
getattr(sa, "enable_prefill_cp", False)
|
get_parallel().enable_prefill_cp
|
||||||
and getattr(sa, "_is_dsa_model_arch", False)
|
and getattr(get_server_args(), "_is_dsa_model_arch", False)
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -288,7 +291,7 @@ def get_cp_strategy() -> Optional[ContextParallelStrategy]:
|
|||||||
server_args = get_server_args()
|
server_args = get_server_args()
|
||||||
except ValueError:
|
except ValueError:
|
||||||
return None
|
return None
|
||||||
if server_args is not None and getattr(server_args, "enable_prefill_cp", False):
|
if server_args is not None and get_parallel().enable_prefill_cp:
|
||||||
init_cp_strategy(server_args)
|
init_cp_strategy(server_args)
|
||||||
return _STRATEGY
|
return _STRATEGY
|
||||||
|
|
||||||
|
|||||||
@@ -27,7 +27,11 @@ from sglang.srt.distributed import (
|
|||||||
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
|
||||||
use_symmetric_memory,
|
use_symmetric_memory,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_flags, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_attn_cp_size,
|
||||||
|
configured_moe_dp_size,
|
||||||
|
get_flags,
|
||||||
|
)
|
||||||
from sglang.srt.utils import get_bool_env_var, is_hip
|
from sglang.srt.utils import get_bool_env_var, is_hip
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -986,8 +990,7 @@ def is_enable_moe_cp_allgather() -> bool:
|
|||||||
(``parallel_state.py``), so the live sizes are equal and the comparison would
|
(``parallel_state.py``), so the live sizes are equal and the comparison would
|
||||||
always be false.
|
always be false.
|
||||||
"""
|
"""
|
||||||
server_args = get_server_args()
|
return configured_attn_cp_size() > configured_moe_dp_size()
|
||||||
return server_args.attn_cp_size > server_args.moe_dp_size
|
|
||||||
|
|
||||||
|
|
||||||
def moe_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor):
|
def moe_cp_all_gather_into_tensor(output: torch.Tensor, input: torch.Tensor):
|
||||||
|
|||||||
@@ -67,16 +67,17 @@ def _get_state() -> Optional[_State]:
|
|||||||
|
|
||||||
if get_device_sm() not in (100, 103):
|
if get_device_sm() not in (100, 103):
|
||||||
return None
|
return None
|
||||||
server_args = ctx.get_server_args()
|
symm_mem = ctx.get_exec().comm.enable_symm_mem
|
||||||
if server_args.enable_symm_mem or server_args.moe_a2a_backend != "none":
|
a2a_backend = ctx.get_exec().moe.moe_a2a_backend
|
||||||
|
if symm_mem or a2a_backend != "none":
|
||||||
logger.info(
|
logger.info(
|
||||||
"K3 all-reduce fusion auto-probe: skipping "
|
"K3 all-reduce fusion auto-probe: skipping "
|
||||||
"(enable_symm_mem=%s, moe_a2a_backend=%s; under symm-mem the "
|
"(enable_symm_mem=%s, moe_a2a_backend=%s; under symm-mem the "
|
||||||
"allocator contexts conflict, and under EP a2a the model's "
|
"allocator contexts conflict, and under EP a2a the model's "
|
||||||
"symm-pool allocation contract does not hold on every AR "
|
"symm-pool allocation contract does not hold on every AR "
|
||||||
"call-site. Set SGLANG_K3_AR_FUSION=1 to force.)",
|
"call-site. Set SGLANG_K3_AR_FUSION=1 to force.)",
|
||||||
server_args.enable_symm_mem,
|
symm_mem,
|
||||||
server_args.moe_a2a_backend,
|
a2a_backend,
|
||||||
)
|
)
|
||||||
return None
|
return None
|
||||||
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
|
from sglang.srt.distributed.device_communicators.custom_all_reduce_v2 import (
|
||||||
@@ -131,11 +132,11 @@ _BUFS: List[_Buffer] = []
|
|||||||
@cache_once
|
@cache_once
|
||||||
def _max_buffer_rows() -> int:
|
def _max_buffer_rows() -> int:
|
||||||
"""Rows to reserve per buffer: the largest batch the server args allow."""
|
"""Rows to reserve per buffer: the largest batch the server args allow."""
|
||||||
server_args = ctx.get_server_args()
|
schedule = ctx.get_schedule()
|
||||||
chunked = server_args.chunked_prefill_size
|
chunked = schedule.chunked_prefill_size
|
||||||
if chunked is not None and chunked > 0:
|
if chunked is not None and chunked > 0:
|
||||||
return int(chunked)
|
return int(chunked)
|
||||||
return int(server_args.max_prefill_tokens or 0)
|
return int(schedule.max_prefill_tokens or 0)
|
||||||
|
|
||||||
|
|
||||||
def _create_buffer(
|
def _create_buffer(
|
||||||
|
|||||||
@@ -12,7 +12,10 @@ from sglang.srt.layers.moe.moe_runner.base import (
|
|||||||
register_fused_func,
|
register_fused_func,
|
||||||
)
|
)
|
||||||
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
from sglang.srt.model_executor.cuda_graph_config import cuda_graph_fully_disabled
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import (
|
||||||
|
cutedsl_moe_max_num_tokens,
|
||||||
|
get_parallel,
|
||||||
|
)
|
||||||
from sglang.srt.utils.common import log_info_on_rank0, print_warning_once
|
from sglang.srt.utils.common import log_info_on_rank0, print_warning_once
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -256,14 +259,11 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None:
|
|||||||
"Install with: pip install flashinfer"
|
"Install with: pip install flashinfer"
|
||||||
) from e
|
) from e
|
||||||
|
|
||||||
from sglang.srt.runtime_context import get_server_args
|
|
||||||
|
|
||||||
assert layer.intermediate_size_per_partition > 0, (
|
assert layer.intermediate_size_per_partition > 0, (
|
||||||
f"CuteDSL MoE: intermediate_size_per_partition must be > 0, "
|
f"CuteDSL MoE: intermediate_size_per_partition must be > 0, "
|
||||||
f"got {layer.intermediate_size_per_partition}. Check EP/TP configuration."
|
f"got {layer.intermediate_size_per_partition}. Check EP/TP configuration."
|
||||||
)
|
)
|
||||||
|
|
||||||
server_args = get_server_args()
|
|
||||||
# CuteDSL wrapper preallocates CG buffers used by any captured graph
|
# CuteDSL wrapper preallocates CG buffers used by any captured graph
|
||||||
# that routes through this MoE — decode and prefill alike.
|
# that routes through this MoE — decode and prefill alike.
|
||||||
use_cuda_graph = not cuda_graph_fully_disabled()
|
use_cuda_graph = not cuda_graph_fully_disabled()
|
||||||
@@ -277,9 +277,7 @@ def ensure_cutedsl_wrapper(layer: torch.nn.Module) -> None:
|
|||||||
else:
|
else:
|
||||||
# Standard allgather path: the MoE sees up to dp_size local forwards
|
# Standard allgather path: the MoE sees up to dp_size local forwards
|
||||||
# gathered together, so scale the per-rank forward bound by dp_size.
|
# gathered together, so scale the per-rank forward bound by dp_size.
|
||||||
max_num_tokens = (
|
max_num_tokens = get_parallel().dp_size * cutedsl_moe_max_num_tokens()
|
||||||
get_parallel().dp_size * server_args.cutedsl_moe_max_num_tokens()
|
|
||||||
)
|
|
||||||
top_k = layer.top_k if layer.top_k is not None else layer.moe_runner_config.top_k
|
top_k = layer.top_k if layer.top_k is not None else layer.moe_runner_config.top_k
|
||||||
# inference_mode(False) ensures the wrapper's pre-allocated CUDA-graph
|
# inference_mode(False) ensures the wrapper's pre-allocated CUDA-graph
|
||||||
# buffers are normal tensors. This call typically happens inside
|
# buffers are normal tensors. This call typically happens inside
|
||||||
|
|||||||
@@ -15,7 +15,10 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
from sglang.srt.layers.moe import get_moe_a2a_backend
|
from sglang.srt.layers.moe import get_moe_a2a_backend
|
||||||
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, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_parallel,
|
||||||
|
uses_mla_backend,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -69,8 +72,7 @@ def is_prefill_cp_in_seq_split():
|
|||||||
|
|
||||||
|
|
||||||
def is_mla_prefill_cp_enabled() -> bool:
|
def is_mla_prefill_cp_enabled() -> bool:
|
||||||
sa = get_server_args()
|
return get_parallel().enable_prefill_context_parallel and uses_mla_backend()
|
||||||
return get_parallel().enable_prefill_context_parallel and sa.use_mla_backend()
|
|
||||||
|
|
||||||
|
|
||||||
def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None):
|
def mla_use_prefill_cp(forward_batch, mla_enable_prefill_cp=None):
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_schedule,
|
get_schedule,
|
||||||
get_serving,
|
get_serving,
|
||||||
get_spec,
|
get_spec,
|
||||||
|
mamba_cache_chunk_size,
|
||||||
mamba_extra_buffer_enabled,
|
mamba_extra_buffer_enabled,
|
||||||
mamba_extra_buffer_lazy_enabled,
|
mamba_extra_buffer_lazy_enabled,
|
||||||
)
|
)
|
||||||
@@ -118,7 +119,7 @@ from sglang.srt.observability.req_time_stats import (
|
|||||||
DPControllerReqTimeStats,
|
DPControllerReqTimeStats,
|
||||||
SchedulerReqTimeStats,
|
SchedulerReqTimeStats,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel, get_server_args
|
from sglang.srt.runtime_context import get_parallel
|
||||||
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 SamplingParams
|
from sglang.srt.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.server_args import ServerArgs
|
from sglang.srt.server_args import ServerArgs
|
||||||
@@ -2616,21 +2617,20 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
self,
|
self,
|
||||||
req: Req,
|
req: Req,
|
||||||
) -> _MambaRadixCacheV2TrackEntry:
|
) -> _MambaRadixCacheV2TrackEntry:
|
||||||
server_args = get_server_args()
|
chunk_size = mamba_cache_chunk_size()
|
||||||
mamba_cache_chunk_size = server_args.mamba_cache_chunk_size
|
|
||||||
|
|
||||||
def _force_track_h(i: int) -> int:
|
def _force_track_h(i: int) -> int:
|
||||||
assert i % mamba_cache_chunk_size == 0
|
assert i % chunk_size == 0
|
||||||
# There are 3 cases for mamba_track_seqlen passed to mamba_track_seqlens_cpu:
|
# There are 3 cases for mamba_track_seqlen passed to mamba_track_seqlens_cpu:
|
||||||
# 1) aligned with mamba_cache_chunk_size-> retrieve from last_recurrent_state
|
# 1) aligned with chunk_size-> retrieve from last_recurrent_state
|
||||||
# a) is the last position -> retrieve from last_recurrent_state
|
# a) is the last position -> retrieve from last_recurrent_state
|
||||||
# b) is NOT the last position -> retrieve from h
|
# b) is NOT the last position -> retrieve from h
|
||||||
# 2) unaligned with mamba_cache_chunk_size -> retrieve from h
|
# 2) unaligned with chunk_size -> retrieve from h
|
||||||
# Currently, the math calculation only supports case 1a and 2. So for 1b, we need to add 1
|
# Currently, the math calculation only supports case 1a and 2. So for 1b, we need to add 1
|
||||||
# to force the math calculation to retrieve the correct mamba state from h.
|
# to force the math calculation to retrieve the correct mamba state from h.
|
||||||
return i + 1
|
return i + 1
|
||||||
|
|
||||||
mask = req.extend_range.length >= mamba_cache_chunk_size
|
mask = req.extend_range.length >= chunk_size
|
||||||
track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item()
|
track_index = req.mamba_ping_pong_track_buffer[req.mamba_next_track_idx].item()
|
||||||
mamba_track_seqlen = -1
|
mamba_track_seqlen = -1
|
||||||
if mask:
|
if mask:
|
||||||
@@ -2647,18 +2647,16 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# mamba radix cache to track which seqlen this mamba state should store at.
|
# mamba radix cache to track which seqlen this mamba state should store at.
|
||||||
mamba_track_seqlen_aligned = (
|
mamba_track_seqlen_aligned = (
|
||||||
len(req.prefix_indices)
|
len(req.prefix_indices)
|
||||||
+ (req.extend_range.length // mamba_cache_chunk_size)
|
+ (req.extend_range.length // chunk_size) * chunk_size
|
||||||
* mamba_cache_chunk_size
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# mamba_track_fla_chunk_aligned is the aligned seqlen based on mamba_cache_chunk_size
|
# mamba_track_fla_chunk_aligned is the aligned seqlen based on chunk_size
|
||||||
# If mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned, which can be true when
|
# If mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned, which can be true when
|
||||||
# page_size > mamba_cache_chunk_size, we need to force the math calculation to retrieve the correct mamba state from h
|
# page_size > chunk_size, we need to force the math calculation to retrieve the correct mamba state from h
|
||||||
# by _force_track_h()
|
# by _force_track_h()
|
||||||
mamba_track_fla_chunk_aligned = (
|
mamba_track_fla_chunk_aligned = (
|
||||||
len(req.prefix_indices)
|
len(req.prefix_indices)
|
||||||
+ (req.extend_range.length // mamba_cache_chunk_size)
|
+ (req.extend_range.length // chunk_size) * chunk_size
|
||||||
* mamba_cache_chunk_size
|
|
||||||
)
|
)
|
||||||
if mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned:
|
if mamba_track_fla_chunk_aligned != mamba_track_seqlen_aligned:
|
||||||
# We want to track mamba_track_seqlen_aligned, and it's not the last position,
|
# We want to track mamba_track_seqlen_aligned, and it's not the last position,
|
||||||
@@ -2679,7 +2677,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# is within the current extend batch.
|
# is within the current extend batch.
|
||||||
branching_seqlen_aligned_mask = (
|
branching_seqlen_aligned_mask = (
|
||||||
req.mamba_branching_seqlen - len(req.prefix_indices)
|
req.mamba_branching_seqlen - len(req.prefix_indices)
|
||||||
) % mamba_cache_chunk_size == 0
|
) % chunk_size == 0
|
||||||
if (
|
if (
|
||||||
req.mamba_branching_seqlen > len(req.prefix_indices)
|
req.mamba_branching_seqlen > len(req.prefix_indices)
|
||||||
and req.mamba_branching_seqlen < mamba_track_seqlen
|
and req.mamba_branching_seqlen < mamba_track_seqlen
|
||||||
|
|||||||
@@ -37,8 +37,8 @@ from sglang.srt.runtime_context import (
|
|||||||
get_exec,
|
get_exec,
|
||||||
get_memory,
|
get_memory,
|
||||||
get_observability,
|
get_observability,
|
||||||
get_server_args,
|
|
||||||
mamba_extra_buffer_lazy_enabled,
|
mamba_extra_buffer_lazy_enabled,
|
||||||
|
max_speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
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
|
||||||
@@ -1163,7 +1163,6 @@ class SchedulerBatchResultProcessor:
|
|||||||
keep_written_by_this_step = (
|
keep_written_by_this_step = (
|
||||||
crossed and planned_pos == req.mamba_next_track_idx
|
crossed and planned_pos == req.mamba_next_track_idx
|
||||||
)
|
)
|
||||||
server_args = get_server_args()
|
|
||||||
other_idx = 1 - req.mamba_next_track_idx
|
other_idx = 1 - req.mamba_next_track_idx
|
||||||
# Recompute the in-flight verify's plan (kv_committed_len is
|
# Recompute the in-flight verify's plan (kv_committed_len is
|
||||||
# frozen since its prepare, so the recompute is exact).
|
# frozen since its prepare, so the recompute is exact).
|
||||||
@@ -1172,7 +1171,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
].item() == -1 and mamba_lazy_spec_in_window(
|
].item() == -1 and mamba_lazy_spec_in_window(
|
||||||
req,
|
req,
|
||||||
get_exec().mamba.mamba_track_interval,
|
get_exec().mamba.mamba_track_interval,
|
||||||
server_args.max_speculative_num_draft_tokens,
|
max_speculative_num_draft_tokens(),
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
planned_pos is None
|
planned_pos is None
|
||||||
|
|||||||
@@ -50,7 +50,9 @@ from sglang.srt.mem_cache.multi_ended_allocator import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.radix_cache import RadixKey
|
from sglang.srt.mem_cache.radix_cache import RadixKey
|
||||||
from sglang.srt.mem_cache.utils import split_node_hash_value
|
from sglang.srt.mem_cache.utils import split_node_hash_value
|
||||||
from sglang.srt.runtime_context import get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
mamba_cache_chunk_size,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
@@ -450,7 +452,7 @@ class MambaRadixCache(KVCacheEventMixin, BasePrefixCache):
|
|||||||
)
|
)
|
||||||
self.req_to_token_pool: HybridReqToTokenPool = params.req_to_token_pool
|
self.req_to_token_pool: HybridReqToTokenPool = params.req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = params.token_to_kv_pool_allocator
|
||||||
self.mamba_cache_chunk_size = get_server_args().mamba_cache_chunk_size
|
self.mamba_cache_chunk_size = mamba_cache_chunk_size()
|
||||||
|
|
||||||
self.page_size = params.page_size
|
self.page_size = params.page_size
|
||||||
self.disable = params.disable
|
self.disable = params.disable
|
||||||
|
|||||||
@@ -35,7 +35,10 @@ 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_exec, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
get_exec,
|
||||||
|
mamba_cache_chunk_size,
|
||||||
|
)
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.managers.schedule_batch import Req
|
from sglang.srt.managers.schedule_batch import Req
|
||||||
@@ -65,7 +68,7 @@ class MambaComponent(TreeComponent):
|
|||||||
params.page_size == 1
|
params.page_size == 1
|
||||||
), 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 = mamba_cache_chunk_size()
|
||||||
self.mamba_max_states_per_path = get_exec().mamba.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
|
||||||
|
|||||||
@@ -48,6 +48,7 @@ from sglang.srt.model_loader.remote_instance_weight_loader_utils import (
|
|||||||
register_memory_region,
|
register_memory_region,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_moe_dp_size,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_model,
|
get_model,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
@@ -1776,14 +1777,13 @@ class PreshardedModelLoader(DefaultModelLoader):
|
|||||||
return 1
|
return 1
|
||||||
|
|
||||||
parallel = get_parallel()
|
parallel = get_parallel()
|
||||||
server_args = get_server_args()
|
|
||||||
return {
|
return {
|
||||||
"tp": _safe(lambda: parallel.tp_size),
|
"tp": _safe(lambda: parallel.tp_size),
|
||||||
"dp": _safe(lambda: parallel.moe_dp_size),
|
"dp": _safe(lambda: parallel.moe_dp_size),
|
||||||
"ep": _safe(lambda: parallel.moe_ep_size),
|
"ep": _safe(lambda: parallel.moe_ep_size),
|
||||||
"pp": _safe(lambda: parallel.pp_size),
|
"pp": _safe(lambda: parallel.pp_size),
|
||||||
"moe_dense_tp_size": parallel.moe_dense_tp_size,
|
"moe_dense_tp_size": parallel.moe_dense_tp_size,
|
||||||
"moe_dp_size": server_args.moe_dp_size,
|
"moe_dp_size": configured_moe_dp_size(),
|
||||||
"enable_dp_lm_head": parallel.enable_dp_lm_head,
|
"enable_dp_lm_head": parallel.enable_dp_lm_head,
|
||||||
"enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head,
|
"enable_fp32_lm_head": get_exec().features.enable_fp32_lm_head,
|
||||||
"quantization": model_config.quantization,
|
"quantization": model_config.quantization,
|
||||||
|
|||||||
@@ -193,12 +193,12 @@ from sglang.srt.models.deepseek_common.utils import (
|
|||||||
is_wint4afp8_or_wint4a16_config,
|
is_wint4afp8_or_wint4a16_config,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
attention_backends,
|
||||||
get_device,
|
get_device,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_forward,
|
get_forward,
|
||||||
get_model,
|
get_model,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
|
||||||
get_spec,
|
get_spec,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
@@ -1999,8 +1999,7 @@ class DeepseekV2AttentionMLA(
|
|||||||
# Determine attention backend name for current forward batch: prefer the
|
# Determine attention backend name for current forward batch: prefer the
|
||||||
# name stamped per-runner on the backend object, else resolve from server args.
|
# name stamped per-runner on the backend object, else resolve from server args.
|
||||||
backend = get_attn_backend()
|
backend = get_attn_backend()
|
||||||
server_args = get_server_args()
|
default_prefill_str, default_decode_str = attention_backends()
|
||||||
default_prefill_str, default_decode_str = server_args.get_attention_backends()
|
|
||||||
prefill_backend_str = (
|
prefill_backend_str = (
|
||||||
backend.prefill_attention_backend_str or default_prefill_str
|
backend.prefill_attention_backend_str or default_prefill_str
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -46,10 +46,9 @@ from sglang.srt.multimodal.mm_utils import (
|
|||||||
run_dp_sharded_mrope_vision_model,
|
run_dp_sharded_mrope_vision_model,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_tp_size,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_mm,
|
get_mm,
|
||||||
get_parallel,
|
|
||||||
get_server_args,
|
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import add_prefix, is_cuda, is_npu
|
from sglang.srt.utils import add_prefix, is_cuda, is_npu
|
||||||
|
|
||||||
@@ -731,13 +730,10 @@ class KimiK25ForConditionalGeneration(nn.Module):
|
|||||||
acknowledges the entire TP group so the bounded IPC pool remains
|
acknowledges the entire TP group so the bounded IPC pool remains
|
||||||
recyclable.
|
recyclable.
|
||||||
"""
|
"""
|
||||||
parallel = get_parallel()
|
# Same source as MmItemMemoryPool.try_to_recycle(), which waits on
|
||||||
server_args = get_server_args()
|
# configured_tp_size(): the live world size agrees once dist is up,
|
||||||
# Match MmItemMemoryPool.try_to_recycle(), which waits for the
|
# but a refcount that disagrees with the waiter would strand items.
|
||||||
# server TP size rather than the attention subgroup size.
|
ipc_consumer_count = max(configured_tp_size(), 1)
|
||||||
ipc_consumer_count = max(
|
|
||||||
getattr(server_args, "tp_size", parallel.attn_tp_size), 1
|
|
||||||
)
|
|
||||||
device_index = device.index
|
device_index = device.index
|
||||||
if device.type == "cuda" and device_index is None:
|
if device.type == "cuda" and device_index is None:
|
||||||
device_index = torch.cuda.current_device()
|
device_index = torch.cuda.current_device()
|
||||||
|
|||||||
@@ -111,7 +111,12 @@ from sglang.srt.multimodal.kimi_k3_image_processing import (
|
|||||||
to_chw_uint8,
|
to_chw_uint8,
|
||||||
)
|
)
|
||||||
from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
|
from sglang.srt.multimodal.mm_utils import materialize_multimodal_features
|
||||||
from sglang.srt.runtime_context import get_exec, get_parallel, get_server_args
|
from sglang.srt.runtime_context import (
|
||||||
|
configured_tp_size,
|
||||||
|
get_exec,
|
||||||
|
get_parallel,
|
||||||
|
get_server_args,
|
||||||
|
)
|
||||||
from sglang.srt.utils import is_blackwell_supported, is_hip, make_layers
|
from sglang.srt.utils import is_blackwell_supported, is_hip, make_layers
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
BumpAllocator,
|
BumpAllocator,
|
||||||
@@ -3082,7 +3087,10 @@ class KimiK3ForConditionalGeneration(nn.Module):
|
|||||||
|
|
||||||
def materialize_item_features(image_indices: List[int]) -> torch.Tensor:
|
def materialize_item_features(image_indices: List[int]) -> torch.Tensor:
|
||||||
"""Materialize only the images assigned to this vision-DP rank."""
|
"""Materialize only the images assigned to this vision-DP rank."""
|
||||||
ipc_consumer_count = max(get_parallel().tp_size, 1)
|
# Same source as MmItemMemoryPool.try_to_recycle(), which waits on
|
||||||
|
# configured_tp_size(): the live world size agrees once dist is up,
|
||||||
|
# but a refcount that disagrees with the waiter would strand items.
|
||||||
|
ipc_consumer_count = max(configured_tp_size(), 1)
|
||||||
device_index = device.index
|
device_index = device.index
|
||||||
if device.type == "cuda" and device_index is None:
|
if device.type == "cuda" and device_index is None:
|
||||||
device_index = torch.cuda.current_device()
|
device_index = torch.cuda.current_device()
|
||||||
|
|||||||
@@ -85,7 +85,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_forward,
|
get_forward,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_schedule,
|
get_schedule,
|
||||||
get_server_args,
|
process_model_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
# get_bool_env_var is defined in sglang.srt.utils.common, not sglang.srt.distributed.
|
# get_bool_env_var is defined in sglang.srt.utils.common, not sglang.srt.distributed.
|
||||||
@@ -431,10 +431,9 @@ class MiniMaxM2QKRMSNorm:
|
|||||||
|
|
||||||
props = torch.cuda.get_device_properties(device)
|
props = torch.cuda.get_device_properties(device)
|
||||||
# probe the maximum tokens for one prefill
|
# probe the maximum tokens for one prefill
|
||||||
server_args = get_server_args()
|
|
||||||
max_tokens = get_schedule().chunked_prefill_size
|
max_tokens = get_schedule().chunked_prefill_size
|
||||||
if max_tokens is None:
|
if max_tokens is None:
|
||||||
max_tokens = server_args.model_config.context_len
|
max_tokens = process_model_config().context_len
|
||||||
max_tokens = max(max_tokens, get_schedule().max_prefill_tokens)
|
max_tokens = max(max_tokens, get_schedule().max_prefill_tokens)
|
||||||
logger.info(f"[AR] Using CustomAllReduceV2 for MiniMaxM2 with {max_tokens = }")
|
logger.info(f"[AR] Using CustomAllReduceV2 for MiniMaxM2 with {max_tokens = }")
|
||||||
ALIGN = 512
|
ALIGN = 512
|
||||||
|
|||||||
@@ -61,8 +61,10 @@ from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha imp
|
|||||||
DeepseekMHAForwardMixin,
|
DeepseekMHAForwardMixin,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
|
attention_backends,
|
||||||
get_exec,
|
get_exec,
|
||||||
get_forward,
|
get_forward,
|
||||||
|
get_memory,
|
||||||
get_model,
|
get_model,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
@@ -609,18 +611,12 @@ class SarvamMoEMLAAttention(nn.Module):
|
|||||||
return k
|
return k
|
||||||
|
|
||||||
def _set_current_attention_backend(self, forward_batch: ForwardBatch) -> None:
|
def _set_current_attention_backend(self, forward_batch: ForwardBatch) -> None:
|
||||||
if self._server_args is None:
|
prefill_backend, decode_backend = attention_backends()
|
||||||
self._server_args = get_server_args()
|
self.current_attention_backend = (
|
||||||
if forward_batch.forward_mode.is_decode_or_idle():
|
decode_backend
|
||||||
self.current_attention_backend = (
|
if forward_batch.forward_mode.is_decode_or_idle()
|
||||||
self._server_args.decode_attention_backend
|
else prefill_backend
|
||||||
or self._server_args.attention_backend
|
)
|
||||||
)
|
|
||||||
else:
|
|
||||||
self.current_attention_backend = (
|
|
||||||
self._server_args.prefill_attention_backend
|
|
||||||
or self._server_args.attention_backend
|
|
||||||
)
|
|
||||||
|
|
||||||
def _maybe_fp8_bmm(
|
def _maybe_fp8_bmm(
|
||||||
self,
|
self,
|
||||||
@@ -684,7 +680,7 @@ class SarvamMoEMLAAttention(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
self._set_current_attention_backend(forward_batch)
|
self._set_current_attention_backend(forward_batch)
|
||||||
can_use_prefix_cache = not self._server_args.disable_radix_cache
|
can_use_prefix_cache = not get_memory().disable_radix_cache
|
||||||
do_prefix_merge = has_extend_prefix and can_use_prefix_cache
|
do_prefix_merge = has_extend_prefix and can_use_prefix_cache
|
||||||
|
|
||||||
if do_prefix_merge and forward_batch.num_prefix_chunks is None:
|
if do_prefix_merge and forward_batch.num_prefix_chunks is None:
|
||||||
|
|||||||
@@ -47,6 +47,7 @@ test-only ``override(**kw)``.
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
|
import functools
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
@@ -827,6 +828,10 @@ class RuntimeContext:
|
|||||||
server_args, "enable_torch_compile", False
|
server_args, "enable_torch_compile", False
|
||||||
)
|
)
|
||||||
self._server_args = server_args
|
self._server_args = server_args
|
||||||
|
# The adaptive draft-token bound memoizes on the config *path*, so a new
|
||||||
|
# publication that reuses the path must not inherit the bound computed
|
||||||
|
# from the file's previous contents.
|
||||||
|
_adaptive_draft_token_bound.cache_clear()
|
||||||
# Snapshot resolved config into the namespace bags (the single source of
|
# Snapshot resolved config into the namespace bags (the single source of
|
||||||
# truth for config reads). Driven by NS(...) metadata; a mock/partial
|
# truth for config reads). Driven by NS(...) metadata; a mock/partial
|
||||||
# config with no NS markers yields an empty tree (no bags projected).
|
# config with no NS markers yields an empty tree (no bags projected).
|
||||||
@@ -1376,6 +1381,7 @@ def reset_context() -> None:
|
|||||||
"""
|
"""
|
||||||
_CONTEXT._server_args = None
|
_CONTEXT._server_args = None
|
||||||
_CONTEXT._config_bags = None
|
_CONTEXT._config_bags = None
|
||||||
|
_adaptive_draft_token_bound.cache_clear()
|
||||||
_CONTEXT._overrides_log = []
|
_CONTEXT._overrides_log = []
|
||||||
_CONTEXT._publish_role = None
|
_CONTEXT._publish_role = None
|
||||||
_CONTEXT.parallel._config = None
|
_CONTEXT.parallel._config = None
|
||||||
@@ -1406,3 +1412,144 @@ def mamba_extra_buffer_lazy_enabled() -> bool:
|
|||||||
get_memory().disable_radix_cache is False
|
get_memory().disable_radix_cache is False
|
||||||
and get_exec().mamba.mamba_radix_cache_strategy == "extra_buffer_lazy"
|
and get_exec().mamba.mamba_radix_cache_strategy == "extra_buffer_lazy"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# --- Derived config accessors ------------------------------------------------
|
||||||
|
#
|
||||||
|
# A few values are computed from several config fields plus the HF config, so
|
||||||
|
# they are ``ServerArgs`` members rather than namespace leaves. Business code
|
||||||
|
# must not reach for the startup record to get them: these accessors are the
|
||||||
|
# named home, and this module — which owns the slot — is the only place that
|
||||||
|
# reads it. Each one keeps the member's exact semantics, including which model
|
||||||
|
# config it derives from (always the process's, i.e. the target's).
|
||||||
|
|
||||||
|
|
||||||
|
def mamba_cache_chunk_size() -> int:
|
||||||
|
"""The caching point granularity for mamba state: ``max(the model's mamba
|
||||||
|
chunk size, page_size)``. Cached on the config after the first call."""
|
||||||
|
return get_server_args().mamba_cache_chunk_size
|
||||||
|
|
||||||
|
|
||||||
|
def max_speculative_num_draft_tokens() -> int | None:
|
||||||
|
"""The largest draft-token count speculative decoding may use.
|
||||||
|
|
||||||
|
All three inputs are ``spec`` leaves, so this derives from the bags and
|
||||||
|
follows a post-publish override; ``ServerArgs.max_speculative_num_draft_tokens``
|
||||||
|
is the pre-publish equivalent. Adaptive spec resolves the count from its
|
||||||
|
candidate-step table instead of the flat field.
|
||||||
|
"""
|
||||||
|
spec = get_spec()
|
||||||
|
if spec.speculative_num_draft_tokens is None:
|
||||||
|
return None
|
||||||
|
if not spec.speculative_adaptive:
|
||||||
|
return spec.speculative_num_draft_tokens
|
||||||
|
# The adaptive branch parses a JSON config, and this is called per decode
|
||||||
|
# batch (`spec_prepare_for_decode`), so memoize on the inputs -- keyed, not
|
||||||
|
# cached once, so a post-publish override still recomputes.
|
||||||
|
return _adaptive_draft_token_bound(spec.speculative_adaptive_config)
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=8)
|
||||||
|
def _adaptive_draft_token_bound(cfg_path: str | None) -> int:
|
||||||
|
from sglang.srt.speculative.adaptive_spec_params import (
|
||||||
|
resolve_candidate_steps_from_config,
|
||||||
|
)
|
||||||
|
|
||||||
|
candidate_steps = resolve_candidate_steps_from_config(cfg_path=cfg_path)
|
||||||
|
# Adaptive spec requires topk=1 today, so each runtime state needs
|
||||||
|
# steps + 1 draft-token slots (mirrors the ServerArgs member).
|
||||||
|
return max(candidate_steps) + 1
|
||||||
|
|
||||||
|
|
||||||
|
def uses_mla_backend() -> bool:
|
||||||
|
"""Whether this process's model runs the MLA attention path."""
|
||||||
|
return get_server_args().use_mla_backend()
|
||||||
|
|
||||||
|
|
||||||
|
def attention_backends() -> tuple:
|
||||||
|
"""The configured ``(prefill, decode)`` backend pair, split fields falling
|
||||||
|
back to ``attention_backend``.
|
||||||
|
|
||||||
|
All three inputs are ``exec.kernel`` leaves, so this derives from the bags
|
||||||
|
and follows a post-publish override; ``ServerArgs.get_attention_backends``
|
||||||
|
is the pre-publish equivalent the resolution pipeline uses. A built runner
|
||||||
|
stamps its own resolved pair (``ModelRunner.prefill_attention_backend_str``);
|
||||||
|
read that when there is a runner in hand.
|
||||||
|
"""
|
||||||
|
from sglang.srt.arg_groups.overrides import attention_backends_of
|
||||||
|
|
||||||
|
# All three leaves live in the same bag, so the resolution pipeline's own
|
||||||
|
# helper applies directly -- one definition of the fallback rule.
|
||||||
|
return attention_backends_of(get_exec().kernel)
|
||||||
|
|
||||||
|
|
||||||
|
def process_model_config():
|
||||||
|
"""The process's ``ModelConfig`` (built once from the published config)."""
|
||||||
|
return get_server_args().get_model_config()
|
||||||
|
|
||||||
|
|
||||||
|
def cutedsl_moe_max_num_tokens() -> int:
|
||||||
|
"""The CuteDSL A2A per-rank token budget.
|
||||||
|
|
||||||
|
Every input is a published leaf (``spec``, ``schedule``, ``exec.graph``), so
|
||||||
|
this derives from the bags and follows a post-publish override;
|
||||||
|
``ServerArgs.cutedsl_moe_max_num_tokens`` is the pre-publish equivalent the
|
||||||
|
resolution pipeline uses. Max over the prefill bound, the piecewise-prefill
|
||||||
|
capture, and the decode/verify bound.
|
||||||
|
"""
|
||||||
|
from sglang.srt.model_executor.cuda_graph_config import Backend
|
||||||
|
|
||||||
|
spec = get_spec()
|
||||||
|
num_tokens_per_req = (
|
||||||
|
(spec.speculative_num_draft_tokens or 1) if spec.speculative_algorithm else 1
|
||||||
|
)
|
||||||
|
prefill_tokens = get_schedule().max_prefill_tokens
|
||||||
|
cg_config = get_exec().graph.cuda_graph_config
|
||||||
|
if cg_config is not None and cg_config.prefill.backend == Backend.TC_PIECEWISE:
|
||||||
|
prefill_tokens = max(prefill_tokens, cg_config.prefill.max_bs or 0)
|
||||||
|
decode_max_bs = (cg_config.decode.max_bs if cg_config is not None else 0) or 0
|
||||||
|
return max(prefill_tokens, decode_max_bs * num_tokens_per_req)
|
||||||
|
|
||||||
|
|
||||||
|
# --- Configured (not live) parallel sizes ------------------------------------
|
||||||
|
#
|
||||||
|
# ``get_parallel()`` shadows these names with the LIVE topology, which is the
|
||||||
|
# right answer almost everywhere. A handful of call sites need what was
|
||||||
|
# *configured* instead — before the groups exist, in a process that has none,
|
||||||
|
# or where the live value is deliberately aliased to another dimension. Each
|
||||||
|
# accessor below names that intent so no business call site has to reach for
|
||||||
|
# the startup record; the per-site reasons live in the read ratchet.
|
||||||
|
#
|
||||||
|
# They read the published leaf rather than the record: the bag is what
|
||||||
|
# ``override`` writes, and once the instance holds only the user's raw input
|
||||||
|
# the record would answer with what was *typed* instead of what resolution
|
||||||
|
# produced. Going through the bag directly is what gets past the live property
|
||||||
|
# that shadows these four names on ``get_parallel()``.
|
||||||
|
|
||||||
|
|
||||||
|
def _configured_parallel(name: str):
|
||||||
|
# The bag itself, not ParallelContext, whose live property shadows these
|
||||||
|
# four names. Read through the parallel slot the way the leaf accessor
|
||||||
|
# does — ``parallel`` is deliberately outside the per-role namespace table
|
||||||
|
# (every process reads topology config), so this must not route through
|
||||||
|
# ``config_bag()``'s role check, which would record or reject the read.
|
||||||
|
config = _CONTEXT.parallel._config
|
||||||
|
if config is None:
|
||||||
|
raise ValueError("config namespace 'parallel' not published")
|
||||||
|
return getattr(config, name)
|
||||||
|
|
||||||
|
|
||||||
|
def configured_tp_size() -> int:
|
||||||
|
return _configured_parallel("tp_size")
|
||||||
|
|
||||||
|
|
||||||
|
def configured_pp_size() -> int:
|
||||||
|
return _configured_parallel("pp_size")
|
||||||
|
|
||||||
|
|
||||||
|
def configured_moe_dp_size() -> int:
|
||||||
|
return _configured_parallel("moe_dp_size")
|
||||||
|
|
||||||
|
|
||||||
|
def configured_attn_cp_size() -> int:
|
||||||
|
return _configured_parallel("attn_cp_size")
|
||||||
|
|||||||
@@ -40,7 +40,12 @@ from sglang.srt.arg_groups.argparse_actions import (
|
|||||||
DeprecatedStoreTrueAction,
|
DeprecatedStoreTrueAction,
|
||||||
LoRAPathAction,
|
LoRAPathAction,
|
||||||
)
|
)
|
||||||
from sglang.srt.arg_groups.overrides import resolved_view
|
from sglang.srt.arg_groups.overrides import (
|
||||||
|
attention_backends_of,
|
||||||
|
mamba_extra_buffer_lazy_of,
|
||||||
|
mamba_extra_buffer_of,
|
||||||
|
resolved_view,
|
||||||
|
)
|
||||||
from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy
|
from sglang.srt.configs.embedding_model_spec import BCGPrefillPolicy
|
||||||
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch
|
from sglang.srt.configs.linear_attn_model_registry import get_linear_attn_spec_by_arch
|
||||||
from sglang.srt.connector import ConnectorType
|
from sglang.srt.connector import ConnectorType
|
||||||
@@ -8689,17 +8694,7 @@ class ServerArgs:
|
|||||||
return attention_backends_of(resolved_view(self))
|
return attention_backends_of(resolved_view(self))
|
||||||
|
|
||||||
def get_attention_backends(self):
|
def get_attention_backends(self):
|
||||||
prefill_attention_backend_str = (
|
return attention_backends_of(self)
|
||||||
self.prefill_attention_backend
|
|
||||||
if self.prefill_attention_backend
|
|
||||||
else self.attention_backend
|
|
||||||
)
|
|
||||||
decode_attention_backend_str = (
|
|
||||||
self.decode_attention_backend
|
|
||||||
if self.decode_attention_backend
|
|
||||||
else self.attention_backend
|
|
||||||
)
|
|
||||||
return prefill_attention_backend_str, decode_attention_backend_str
|
|
||||||
|
|
||||||
def use_mla_backend(self):
|
def use_mla_backend(self):
|
||||||
from sglang.srt.configs.model_config import AttentionArch
|
from sglang.srt.configs.model_config import AttentionArch
|
||||||
@@ -8715,16 +8710,10 @@ class ServerArgs:
|
|||||||
)
|
)
|
||||||
|
|
||||||
def enable_mamba_extra_buffer(self) -> bool:
|
def enable_mamba_extra_buffer(self) -> bool:
|
||||||
return (
|
return mamba_extra_buffer_of(self)
|
||||||
self.disable_radix_cache is False
|
|
||||||
and self.mamba_radix_cache_strategy in ("extra_buffer", "extra_buffer_lazy")
|
|
||||||
)
|
|
||||||
|
|
||||||
def enable_mamba_extra_buffer_lazy(self) -> bool:
|
def enable_mamba_extra_buffer_lazy(self) -> bool:
|
||||||
return (
|
return mamba_extra_buffer_lazy_of(self)
|
||||||
self.disable_radix_cache is False
|
|
||||||
and self.mamba_radix_cache_strategy == "extra_buffer_lazy"
|
|
||||||
)
|
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def max_speculative_num_draft_tokens(self) -> Optional[int]:
|
def max_speculative_num_draft_tokens(self) -> Optional[int]:
|
||||||
|
|||||||
@@ -48,10 +48,10 @@ from sglang.srt.mem_cache.allocation import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_exec,
|
get_exec,
|
||||||
get_server_args,
|
|
||||||
get_spec,
|
get_spec,
|
||||||
mamba_extra_buffer_enabled,
|
mamba_extra_buffer_enabled,
|
||||||
mamba_extra_buffer_lazy_enabled,
|
mamba_extra_buffer_lazy_enabled,
|
||||||
|
max_speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
is_cpu,
|
is_cpu,
|
||||||
@@ -1030,12 +1030,11 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None:
|
|||||||
"""eagle/ngram share a stateless free function; dflash keeps stateful
|
"""eagle/ngram share a stateless free function; dflash keeps stateful
|
||||||
prep on its draft input -- the dispatcher routes.
|
prep on its draft input -- the dispatcher routes.
|
||||||
"""
|
"""
|
||||||
server_args = get_server_args()
|
|
||||||
if mamba_extra_buffer_lazy_enabled():
|
if mamba_extra_buffer_lazy_enabled():
|
||||||
# Scheduler phase (outside forward isolation).
|
# Scheduler phase (outside forward isolation).
|
||||||
batch.mamba_lazy_spec_prepare(
|
batch.mamba_lazy_spec_prepare(
|
||||||
get_exec().mamba.mamba_track_interval,
|
get_exec().mamba.mamba_track_interval,
|
||||||
server_args.max_speculative_num_draft_tokens,
|
max_speculative_num_draft_tokens(),
|
||||||
)
|
)
|
||||||
if batch.spec_algorithm.is_dflash_family():
|
if batch.spec_algorithm.is_dflash_family():
|
||||||
batch.spec_info.prepare_for_decode(batch)
|
batch.spec_info.prepare_for_decode(batch)
|
||||||
|
|||||||
@@ -9,7 +9,9 @@ import numpy as np
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
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 (
|
||||||
|
configured_tp_size,
|
||||||
|
)
|
||||||
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
|
from sglang.srt.utils.stale_shm_cleanup import make_shm_name
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -130,7 +132,7 @@ class MmItemMemoryChunk:
|
|||||||
|
|
||||||
def try_to_recycle(self) -> bool:
|
def try_to_recycle(self) -> bool:
|
||||||
try:
|
try:
|
||||||
tp_num = get_server_args().tp_size
|
tp_num = configured_tp_size()
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.info(
|
logger.info(
|
||||||
"server_args has not been published yet, skip this turn's recycle"
|
"server_args has not been published yet, skip this turn's recycle"
|
||||||
|
|||||||
@@ -18,6 +18,7 @@ from sglang.srt.layers.attention.linear.kernels.gdn_flashinfer import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
|
from sglang.srt.layers.attention.linear.kernels.gdn_triton import TritonGDNKernel
|
||||||
from sglang.srt.layers.attention.linear.utils import LinearAttnKernelBackend
|
from sglang.srt.layers.attention.linear.utils import LinearAttnKernelBackend
|
||||||
|
from sglang.srt.runtime_context import get_context
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
@@ -151,12 +152,12 @@ class TestFlashInferGDNPrefillBackendPolicy(unittest.TestCase):
|
|||||||
track_ssm_h_dst=torch.empty(4),
|
track_ssm_h_dst=torch.empty(4),
|
||||||
)
|
)
|
||||||
|
|
||||||
with patch(
|
# The chunk size is a derived config member; seed its private cache on a
|
||||||
"sglang.srt.layers.attention.linear.kernels.gdn_flashinfer."
|
# published config rather than patching an import binding.
|
||||||
"get_server_args",
|
override = get_context().override_server_args(_mamba_cache_chunk_size=64)
|
||||||
return_value=SimpleNamespace(mamba_cache_chunk_size=64),
|
override.install()
|
||||||
):
|
self.addCleanup(override.restore)
|
||||||
maybe_build_flashinfer_checkpoint_plan(forward_batch, metadata, "cpu")
|
maybe_build_flashinfer_checkpoint_plan(forward_batch, metadata, "cpu")
|
||||||
|
|
||||||
torch.testing.assert_close(
|
torch.testing.assert_close(
|
||||||
metadata.state_checkpoint_cu_starts,
|
metadata.state_checkpoint_cu_starts,
|
||||||
|
|||||||
@@ -786,13 +786,6 @@ class TestShardConfig(unittest.TestCase):
|
|||||||
# `_collect_shard_config` is the exact failure mode that left
|
# `_collect_shard_config` is the exact failure mode that left
|
||||||
# moe_dense_tp_size / LM-head flags out of the cache key before.
|
# moe_dense_tp_size / LM-head flags out of the cache key before.
|
||||||
loader = object.__new__(PreshardedModelLoader)
|
loader = object.__new__(PreshardedModelLoader)
|
||||||
server_args = SimpleNamespace(
|
|
||||||
moe_dp_size=2,
|
|
||||||
enable_fp32_lm_head=True,
|
|
||||||
ep_num_redundant_experts=4,
|
|
||||||
enable_eplb=True,
|
|
||||||
init_expert_location="trivial",
|
|
||||||
)
|
|
||||||
model_config = SimpleNamespace(quantization="fp8", dtype=torch.bfloat16)
|
model_config = SimpleNamespace(quantization="fp8", dtype=torch.bfloat16)
|
||||||
required = {
|
required = {
|
||||||
"tp",
|
"tp",
|
||||||
@@ -819,8 +812,8 @@ class TestShardConfig(unittest.TestCase):
|
|||||||
enable_dp_lm_head=True,
|
enable_dp_lm_head=True,
|
||||||
)
|
)
|
||||||
with mock.patch(
|
with mock.patch(
|
||||||
"sglang.srt.model_loader.loader.get_server_args",
|
"sglang.srt.model_loader.loader.configured_moe_dp_size",
|
||||||
return_value=server_args,
|
return_value=2,
|
||||||
), mock.patch(
|
), mock.patch(
|
||||||
"sglang.srt.model_loader.loader.get_parallel",
|
"sglang.srt.model_loader.loader.get_parallel",
|
||||||
return_value=parallel,
|
return_value=parallel,
|
||||||
|
|||||||
@@ -39,7 +39,7 @@ from sglang.srt.multimodal.processors.kimi_k25 import (
|
|||||||
_resize_bicubic_if_needed,
|
_resize_bicubic_if_needed,
|
||||||
_resize_images_by_source_shape,
|
_resize_images_by_source_shape,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import get_parallel
|
from sglang.srt.runtime_context import get_context, get_parallel
|
||||||
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
from sglang.srt.utils.cuda_ipc_transport_utils import (
|
||||||
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
DEFER_CUDA_IPC_FEATURE_RECONSTRUCTION_KEY,
|
||||||
CudaIpcTensorTransportProxy,
|
CudaIpcTensorTransportProxy,
|
||||||
@@ -317,7 +317,11 @@ def test_dp_helper_supports_moonvit3d_packed_embeddings_on_tp1():
|
|||||||
tower = _MoonViT3dTower()
|
tower = _MoonViT3dTower()
|
||||||
pixel_values = torch.randn(4, 2)
|
pixel_values = torch.randn(4, 2)
|
||||||
|
|
||||||
with get_parallel().override(tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0):
|
# The IPC consumer count asks for the *configured* TP size (matching
|
||||||
|
# MmItemMemoryPool.try_to_recycle), so the double publishes one too.
|
||||||
|
with get_context().override_server_args(tp_size=1), get_parallel().override(
|
||||||
|
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
|
||||||
|
):
|
||||||
output = run_dp_sharded_mrope_vision_model(
|
output = run_dp_sharded_mrope_vision_model(
|
||||||
tower, pixel_values, [[1, 2, 2]], rope_type="rope_2d_packed"
|
tower, pixel_values, [[1, 2, 2]], rope_type="rope_2d_packed"
|
||||||
)
|
)
|
||||||
@@ -331,7 +335,11 @@ def test_dp_helper_can_lazily_load_kimi_features_on_tp1():
|
|||||||
pixel_values = torch.randn(4, 2)
|
pixel_values = torch.randn(4, 2)
|
||||||
loader = Mock(return_value=pixel_values)
|
loader = Mock(return_value=pixel_values)
|
||||||
|
|
||||||
with get_parallel().override(tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0):
|
# The IPC consumer count asks for the *configured* TP size (matching
|
||||||
|
# MmItemMemoryPool.try_to_recycle), so the double publishes one too.
|
||||||
|
with get_context().override_server_args(tp_size=1), get_parallel().override(
|
||||||
|
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
|
||||||
|
):
|
||||||
output = run_dp_sharded_mrope_vision_model(
|
output = run_dp_sharded_mrope_vision_model(
|
||||||
tower,
|
tower,
|
||||||
None,
|
None,
|
||||||
@@ -479,11 +487,10 @@ def test_kimi_non_dp_keeps_grid_thws_on_the_host():
|
|||||||
model.mm_projector = _IdentityProjector()
|
model.mm_projector = _IdentityProjector()
|
||||||
items = [_image_item(torch.randn(4, 2), [[1, 2, 2]])]
|
items = [_image_item(torch.randn(4, 2), [[1, 2, 2]])]
|
||||||
|
|
||||||
with get_parallel().override(
|
# The IPC consumer count asks for the *configured* TP size (matching
|
||||||
|
# MmItemMemoryPool.try_to_recycle), so the double publishes one too.
|
||||||
|
with get_context().override_server_args(tp_size=1), get_parallel().override(
|
||||||
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
|
tp_size=1, tp_rank=0, attn_tp_size=1, attn_tp_rank=0
|
||||||
), patch(
|
|
||||||
"sglang.srt.models.kimi_k25.get_server_args",
|
|
||||||
return_value=SimpleNamespace(tp_size=1),
|
|
||||||
):
|
):
|
||||||
model.get_image_feature(items)
|
model.get_image_feature(items)
|
||||||
|
|
||||||
|
|||||||
@@ -464,15 +464,17 @@ def test_kimi_k3_encoder_dp_defers_feature_materialization(monkeypatch):
|
|||||||
]
|
]
|
||||||
sharded_embeddings = torch.randn(2, 2)
|
sharded_embeddings = torch.randn(2, 2)
|
||||||
|
|
||||||
|
# The IPC consumer count asks for the *configured* TP size (matching
|
||||||
|
# MmItemMemoryPool.try_to_recycle), so publish it; the live topology the
|
||||||
|
# sharding helper reads is forced through the context's own override.
|
||||||
with mock_patch(
|
with mock_patch(
|
||||||
"sglang.srt.multimodal.mm_utils.run_dp_sharded_mrope_vision_model",
|
"sglang.srt.multimodal.mm_utils.run_dp_sharded_mrope_vision_model",
|
||||||
return_value=sharded_embeddings,
|
return_value=sharded_embeddings,
|
||||||
) as run_dp, mock_patch(
|
) as run_dp, get_context().override_server_args(tp_size=1), get_parallel().override(
|
||||||
"sglang.srt.models.kimi_k3.get_parallel",
|
tp_size=1, attn_tp_size=1
|
||||||
return_value=SimpleNamespace(tp_size=1, attn_tp_size=1),
|
|
||||||
):
|
):
|
||||||
output = model.get_image_feature(items)
|
output = model.get_image_feature(items)
|
||||||
# Exercise the loader while the runtime topology is patched.
|
# Exercise the loader while the runtime topology is forced.
|
||||||
loader_in_scope = run_dp.call_args.kwargs["load_local_pixel_values"]
|
loader_in_scope = run_dp.call_args.kwargs["load_local_pixel_values"]
|
||||||
local = loader_in_scope([1])
|
local = loader_in_scope([1])
|
||||||
both = loader_in_scope([0, 1])
|
both = loader_in_scope([0, 1])
|
||||||
@@ -545,12 +547,13 @@ def test_kimi_k3_preprocesses_only_dp_owner_images(monkeypatch):
|
|||||||
calls.append([int(image[0, 0, 0]) for image in images])
|
calls.append([int(image[0, 0, 0]) for image in images])
|
||||||
return torch.tensor([[float(calls[-1][0]), 0.0]]), torch.tensor([[1, 1, 1]])
|
return torch.tensor([[float(calls[-1][0]), 0.0]]), torch.tensor([[1, 1, 1]])
|
||||||
|
|
||||||
|
# Configured TP size (the IPC consumer count) comes from the published
|
||||||
|
# bags; the live topology is forced through the context's own override.
|
||||||
with mock_patch(
|
with mock_patch(
|
||||||
"sglang.srt.multimodal.mm_utils.run_dp_sharded_mrope_vision_model",
|
"sglang.srt.multimodal.mm_utils.run_dp_sharded_mrope_vision_model",
|
||||||
return_value=torch.zeros(1, 2),
|
return_value=torch.zeros(1, 2),
|
||||||
) as run_dp, mock_patch(
|
) as run_dp, get_context().override_server_args(tp_size=1), get_parallel().override(
|
||||||
"sglang.srt.models.kimi_k3.get_parallel",
|
tp_size=1, attn_tp_size=1
|
||||||
return_value=SimpleNamespace(tp_size=1, attn_tp_size=1),
|
|
||||||
), mock_patch(
|
), mock_patch(
|
||||||
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
|
"sglang.srt.multimodal.processors.kimi_k25._gpu_preprocess_images",
|
||||||
side_effect=fake_preprocess,
|
side_effect=fake_preprocess,
|
||||||
|
|||||||
@@ -197,9 +197,6 @@ class TestNgramMambaVerifyUpdate(CustomTestCase):
|
|||||||
with patch(
|
with patch(
|
||||||
"sglang.srt.speculative.spec_utils.mambaish_config",
|
"sglang.srt.speculative.spec_utils.mambaish_config",
|
||||||
return_value={"some": "config"},
|
return_value={"some": "config"},
|
||||||
), patch(
|
|
||||||
"sglang.srt.speculative.spec_utils.get_server_args",
|
|
||||||
return_value=MagicMock(mamba_track_interval=256),
|
|
||||||
), patch(
|
), patch(
|
||||||
"sglang.srt.speculative.spec_utils.get_exec",
|
"sglang.srt.speculative.spec_utils.get_exec",
|
||||||
return_value=MagicMock(mamba=MagicMock(mamba_track_interval=256)),
|
return_value=MagicMock(mamba=MagicMock(mamba_track_interval=256)),
|
||||||
|
|||||||
@@ -0,0 +1,191 @@
|
|||||||
|
"""No local may shadow a ``runtime_context`` accessor it also calls.
|
||||||
|
|
||||||
|
A mechanical sweep that rewrites ``self.server_args.mamba_cache_chunk_size``
|
||||||
|
into ``mamba_cache_chunk_size()`` turns
|
||||||
|
|
||||||
|
mamba_cache_chunk_size = self.server_args.mamba_cache_chunk_size
|
||||||
|
|
||||||
|
into ``mamba_cache_chunk_size = mamba_cache_chunk_size()``, which is a
|
||||||
|
self-referential local: the name is local for the whole function, so the call
|
||||||
|
raises ``UnboundLocalError`` the first time that line runs. Five of these
|
||||||
|
shipped in one sweep and only one had unit coverage — a mamba model on the
|
||||||
|
radix-cache-v2 path found it at request time.
|
||||||
|
|
||||||
|
This scans for the shape directly: a function-scope assignment whose target
|
||||||
|
name is an imported accessor.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import ast
|
||||||
|
import unittest
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import sglang
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_PACKAGE_ROOT = Path(next(iter(sglang.__path__)))
|
||||||
|
_CONTEXT_MODULE = "sglang.srt.runtime_context"
|
||||||
|
|
||||||
|
|
||||||
|
def _module_level_accessor_imports(tree: ast.AST) -> set[str]:
|
||||||
|
"""Accessors imported at module scope — visible in every function.
|
||||||
|
|
||||||
|
A *function-local* import is visible only inside its own scope, so it is
|
||||||
|
collected per function in the scan below: charging it file-wide would flag
|
||||||
|
an unrelated sibling function that binds the same name, where no shadowing
|
||||||
|
can occur.
|
||||||
|
"""
|
||||||
|
names: set[str] = set()
|
||||||
|
stack = list(tree.body)
|
||||||
|
while stack:
|
||||||
|
stmt = stack.pop()
|
||||||
|
if isinstance(
|
||||||
|
stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.Lambda, ast.ClassDef)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
if isinstance(stmt, ast.ImportFrom) and stmt.module == _CONTEXT_MODULE:
|
||||||
|
for alias in stmt.names:
|
||||||
|
names.add(alias.asname or alias.name)
|
||||||
|
stack.extend(ast.iter_child_nodes(stmt))
|
||||||
|
return names
|
||||||
|
|
||||||
|
|
||||||
|
def _bound_names(target):
|
||||||
|
"""Every name a binding target introduces, unpacking included.
|
||||||
|
|
||||||
|
``a, (b, c) = ...`` and ``for x, y in ...`` bind through Tuple/List/Starred
|
||||||
|
nodes, so a check that only accepts a bare ``ast.Name`` misses them.
|
||||||
|
"""
|
||||||
|
if isinstance(target, ast.Name):
|
||||||
|
yield target.id
|
||||||
|
elif isinstance(target, ast.Starred):
|
||||||
|
yield from _bound_names(target.value)
|
||||||
|
elif isinstance(target, (ast.Tuple, ast.List)):
|
||||||
|
for element in target.elts:
|
||||||
|
yield from _bound_names(element)
|
||||||
|
|
||||||
|
|
||||||
|
def _own_scope_statements(node) -> tuple:
|
||||||
|
"""This function's OWN scope: its statements, plus the (name, lineno) of
|
||||||
|
each nested ``def``/``class`` — the definition's *name* is a binding in
|
||||||
|
this scope (an earlier accessor call raises UnboundLocalError just like an
|
||||||
|
assignment), while its *body* is the nested scope's own and descending into
|
||||||
|
it would misattribute bindings."""
|
||||||
|
own_scope = []
|
||||||
|
nested_def_bindings = []
|
||||||
|
pending = list(node.body)
|
||||||
|
while pending:
|
||||||
|
stmt = pending.pop()
|
||||||
|
if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)):
|
||||||
|
nested_def_bindings.append((stmt.name, stmt.lineno))
|
||||||
|
continue
|
||||||
|
if isinstance(stmt, ast.Lambda):
|
||||||
|
continue
|
||||||
|
own_scope.append(stmt)
|
||||||
|
pending.extend(ast.iter_child_nodes(stmt))
|
||||||
|
return own_scope, nested_def_bindings
|
||||||
|
|
||||||
|
|
||||||
|
def _child_functions(body) -> list:
|
||||||
|
"""Function defs directly beneath this scope — descending through plain
|
||||||
|
statements and class bodies (a method closes over the enclosing function's
|
||||||
|
names, not the class's), but never into another function."""
|
||||||
|
funcs = []
|
||||||
|
pending = list(body)
|
||||||
|
while pending:
|
||||||
|
stmt = pending.pop()
|
||||||
|
if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||||
|
funcs.append(stmt)
|
||||||
|
continue
|
||||||
|
if isinstance(stmt, ast.Lambda):
|
||||||
|
continue
|
||||||
|
pending.extend(ast.iter_child_nodes(stmt))
|
||||||
|
return funcs
|
||||||
|
|
||||||
|
|
||||||
|
def _shadowing_assignments(tree: ast.AST, module_accessors: set[str]):
|
||||||
|
"""Function-local bindings whose name shadows an accessor visible in that
|
||||||
|
scope -- every statement form that binds a local, not just ``=``.
|
||||||
|
|
||||||
|
Python decides a name is local from *any* binding in the function, so a
|
||||||
|
loop variable, a ``with ... as``, a walrus, a comprehension target, or an
|
||||||
|
``except ... as`` all shadow the accessor for the whole function body,
|
||||||
|
exactly like an assignment does.
|
||||||
|
|
||||||
|
Visibility follows lexical scope: module-level imports reach every
|
||||||
|
function; a function-local import reaches its own scope and nested
|
||||||
|
functions (closure), but NOT an unrelated sibling — charging it file-wide
|
||||||
|
would flag bindings where no shadowing occurs. A function-scope *re-import*
|
||||||
|
of the accessor is itself fine: it binds the name to the same callable, so
|
||||||
|
calls after it behave identically (and the module is full of deliberate
|
||||||
|
local imports).
|
||||||
|
"""
|
||||||
|
stack = [(fn, module_accessors) for fn in _child_functions(tree.body)]
|
||||||
|
while stack:
|
||||||
|
node, inherited = stack.pop()
|
||||||
|
own_scope, nested_def_bindings = _own_scope_statements(node)
|
||||||
|
local_imports = {
|
||||||
|
alias.asname or alias.name
|
||||||
|
for stmt in own_scope
|
||||||
|
if isinstance(stmt, ast.ImportFrom) and stmt.module == _CONTEXT_MODULE
|
||||||
|
for alias in stmt.names
|
||||||
|
}
|
||||||
|
visible = inherited | local_imports
|
||||||
|
# ``def get_exec(): ...`` nested in the function binds the name in
|
||||||
|
# THIS scope, exactly like an assignment would.
|
||||||
|
for name, lineno in nested_def_bindings:
|
||||||
|
if name in visible:
|
||||||
|
yield node.name, name, lineno
|
||||||
|
for inner in own_scope:
|
||||||
|
targets = []
|
||||||
|
if isinstance(inner, ast.Assign):
|
||||||
|
targets = inner.targets
|
||||||
|
elif isinstance(inner, (ast.AnnAssign, ast.AugAssign)):
|
||||||
|
targets = [inner.target]
|
||||||
|
elif isinstance(inner, (ast.For, ast.AsyncFor, ast.comprehension)):
|
||||||
|
targets = [inner.target]
|
||||||
|
elif isinstance(inner, ast.NamedExpr):
|
||||||
|
targets = [inner.target]
|
||||||
|
elif isinstance(inner, (ast.With, ast.AsyncWith)):
|
||||||
|
targets = [i.optional_vars for i in inner.items if i.optional_vars]
|
||||||
|
elif isinstance(inner, ast.ExceptHandler) and inner.name:
|
||||||
|
targets = [ast.Name(id=inner.name, ctx=ast.Store())]
|
||||||
|
for target in targets:
|
||||||
|
for name in _bound_names(target):
|
||||||
|
if name in visible:
|
||||||
|
yield node.name, name, getattr(inner, "lineno", node.lineno)
|
||||||
|
for nested in _child_functions(node.body):
|
||||||
|
stack.append((nested, visible))
|
||||||
|
|
||||||
|
|
||||||
|
class TestNoAccessorShadowing(CustomTestCase):
|
||||||
|
def test_no_local_shadows_a_context_accessor(self):
|
||||||
|
offenders = []
|
||||||
|
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
|
||||||
|
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
|
||||||
|
if rel.startswith("srt/runtime_context.py"):
|
||||||
|
continue
|
||||||
|
source = path.read_text()
|
||||||
|
# A file that never names the module cannot import an accessor
|
||||||
|
# from it, at module scope or inside any function.
|
||||||
|
if _CONTEXT_MODULE not in source:
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
tree = ast.parse(source)
|
||||||
|
except SyntaxError:
|
||||||
|
continue
|
||||||
|
module_accessors = _module_level_accessor_imports(tree)
|
||||||
|
for func, name, lineno in _shadowing_assignments(tree, module_accessors):
|
||||||
|
offenders.append(f"{rel}:{lineno}: {func}() binds {name!r}")
|
||||||
|
self.assertFalse(
|
||||||
|
offenders,
|
||||||
|
"locals shadow a runtime_context accessor imported in the same "
|
||||||
|
"module; the name is local for the whole function, so any call to "
|
||||||
|
"the accessor there raises UnboundLocalError:\n" + "\n".join(offenders),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
@@ -6,22 +6,24 @@ startup record. Config decisions read the namespace accessors instead
|
|||||||
including post-publish overrides, and per-runner values come from the runner
|
including post-publish overrides, and per-runner values come from the runner
|
||||||
that owns them.
|
that owns them.
|
||||||
|
|
||||||
Two shapes count as a read: the direct ``get_server_args().field``, and the
|
Business code no longer reads the published record for a config value at all:
|
||||||
alias ``sa = get_server_args()`` followed by ``sa.field`` in the same function.
|
the baselines are zero for both shapes, over the whole package minus the two
|
||||||
A whole-object pass (``def f(server_args)``) is not a global read and is not
|
modules that own the slot.
|
||||||
counted — there the caller decided which instance to hand over.
|
|
||||||
|
|
||||||
What legitimately remains:
|
Where the remaining reads live (``runtime_context.py``, exempt by module):
|
||||||
|
|
||||||
- **Derived APIs.** ``@property`` and method members of ``ServerArgs``
|
- **Derived members.** ``@property`` / method members of ``ServerArgs``
|
||||||
(``mamba_cache_chunk_size``, ``get_model_config()``,
|
(``mamba_cache_chunk_size``, ``max_speculative_num_draft_tokens``,
|
||||||
``enable_mamba_extra_buffer()``, …) are computed from several fields plus the
|
``use_mla_backend()``, ``get_attention_backends()``, ``get_model_config()``,
|
||||||
HF config, so they are not namespace leaves and ``ServerArgs`` is their only
|
``cutedsl_moe_max_num_tokens()``) are computed from several fields plus the HF
|
||||||
home. Exempt by name below.
|
config, so they are not namespace leaves and ``ServerArgs`` is their only
|
||||||
|
home. ``runtime_context`` exposes each one as a named accessor
|
||||||
|
(``mamba_cache_chunk_size()`` …) and is the only module that reads the slot
|
||||||
|
for them.
|
||||||
- **Config-intent reads of live-shadowed sizes.** ``get_parallel()`` shadows
|
- **Config-intent reads of live-shadowed sizes.** ``get_parallel()`` shadows
|
||||||
``tp/pp/dcp/attn_cp/moe_dp_size`` with the live topology, so a config-intent
|
``tp/pp/dcp/attn_cp/moe_dp_size`` with the live topology, and a few call sites
|
||||||
read of one has nowhere else to go. Each exempt site needs an answer the live
|
need what was *configured*: the ``configured_*_size()`` accessors. Their
|
||||||
property cannot give:
|
reasons, per call site:
|
||||||
|
|
||||||
- ``dsa_indexer.pp_size`` gates ``pp_size > 1 and not get_pp_group()...``, and
|
- ``dsa_indexer.pp_size`` gates ``pp_size > 1 and not get_pp_group()...``, and
|
||||||
the short circuit is the point: with PP off the group is never touched, which
|
the short circuit is the point: with PP off the group is never touched, which
|
||||||
@@ -35,8 +37,23 @@ What legitimately remains:
|
|||||||
sizes are equal there and a live comparison is always false.
|
sizes are equal there and a live comparison is always false.
|
||||||
- ``model_loader/loader.py`` reports both: the same dict carries the live
|
- ``model_loader/loader.py`` reports both: the same dict carries the live
|
||||||
``moe_dp_size`` under ``"dp"``, so this entry is the configured intent.
|
``moe_dp_size`` under ``"dp"``, so this entry is the configured intent.
|
||||||
- The alias-form baseline is not zero yet. Lowering it is the next slice; the
|
|
||||||
failure message lists the sites whenever the count moves.
|
What the ratchet sees, syntactically: ``get_server_args().field``,
|
||||||
|
``sa = get_server_args()`` followed by ``sa.field`` (function-local, module-level,
|
||||||
|
or parked on an instance attribute -- ``self._sa = get_server_args()`` read from
|
||||||
|
another method of the same class), function-local copies of an alias to a
|
||||||
|
fixpoint (``cfg = sa`` then ``cfg.field``), and the dynamic form of each --
|
||||||
|
``getattr(<either>, "field")`` -- since a string-named read reaches the same
|
||||||
|
slot. What it cannot see is a name computed at runtime (``getattr(sa, name)``)
|
||||||
|
or indirection deeper than a local name copy (through a container, an
|
||||||
|
attribute of another object, a cross-scope copy); the census tool in the
|
||||||
|
context repo is what audits those.
|
||||||
|
|
||||||
|
A whole-object pass (``def f(server_args)``) is not a global read and is not
|
||||||
|
counted -- there the caller decided which instance to hand over. An optional
|
||||||
|
parameter that falls back to the global (``f(server_args=None)``) hides one,
|
||||||
|
so those fallbacks were removed; the ratchet cannot see them and the census
|
||||||
|
tool in the context repo is what audits that shape.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
@@ -54,48 +71,92 @@ from sglang.test.test_utils import CustomTestCase
|
|||||||
# scanned so a new one cannot appear there unnoticed.
|
# scanned so a new one cannot appear there unnoticed.
|
||||||
_PACKAGE_ROOT = Path(next(iter(sglang.__path__)))
|
_PACKAGE_ROOT = Path(next(iter(sglang.__path__)))
|
||||||
|
|
||||||
_DERIVED_MEMBERS = frozenset(
|
# The modules that own the slot: runtime_context publishes it and exposes the
|
||||||
{
|
# named accessors for the derived members, server_args/arg_groups ARE the
|
||||||
"cutedsl_moe_max_num_tokens",
|
# resolution pipeline.
|
||||||
"enable_mamba_extra_buffer",
|
_SLOT_OWNERS = ("srt/runtime_context.py", "srt/server_args.py", "srt/arg_groups/")
|
||||||
"enable_mamba_extra_buffer_lazy",
|
|
||||||
"get_attention_backends",
|
|
||||||
"get_model_config",
|
|
||||||
"mamba_cache_chunk_size",
|
|
||||||
"max_speculative_num_draft_tokens",
|
|
||||||
"model_config",
|
|
||||||
"use_mla_backend",
|
|
||||||
}
|
|
||||||
)
|
|
||||||
|
|
||||||
_CONFIG_INTENT_SIZES = frozenset(
|
# Every call site of a ``configured_*_size()`` accessor, with the reason the
|
||||||
{
|
# live topology cannot answer there. The test below asserts this map is exactly
|
||||||
("srt/layers/attention/dsa/dsa_indexer.py", "pp_size"),
|
# the set of call sites, so the reasons cannot drift away from the code.
|
||||||
("srt/layers/dp_attention.py", "attn_cp_size"),
|
_CONFIGURED_SIZE_CALL_SITES = {
|
||||||
("srt/layers/dp_attention.py", "moe_dp_size"),
|
("srt/layers/attention/dsa/dsa_indexer.py", "configured_pp_size"): (
|
||||||
("srt/model_loader/loader.py", "moe_dp_size"),
|
"gates `pp_size > 1 and not get_pp_group()...`; the short circuit is the "
|
||||||
("srt/utils/cuda_ipc_transport_utils.py", "tp_size"),
|
"point, since with PP off the group is never touched, which is what lets "
|
||||||
}
|
"the Indexer be constructed before distributed init"
|
||||||
)
|
),
|
||||||
|
("srt/layers/dp_attention.py", "configured_attn_cp_size"): (
|
||||||
|
"compared against the configured moe_dp_size below"
|
||||||
|
),
|
||||||
|
("srt/layers/dp_attention.py", "configured_moe_dp_size"): (
|
||||||
|
"the configuration this predicate detects (attn_cp_size > moe_dp_size) is "
|
||||||
|
"the one where initialize_model_parallel aliases _MOE_DP to _ATTN_CP, so "
|
||||||
|
"the live sizes are equal there and a live comparison is always false"
|
||||||
|
),
|
||||||
|
("srt/model_loader/loader.py", "configured_moe_dp_size"): (
|
||||||
|
"the same dict already carries the live moe_dp_size under 'dp'; this entry "
|
||||||
|
"is the configured intent"
|
||||||
|
),
|
||||||
|
("srt/utils/cuda_ipc_transport_utils.py", "configured_tp_size"): (
|
||||||
|
"runs in the tokenizer process, which has no parallel groups at all"
|
||||||
|
),
|
||||||
|
("srt/models/kimi_k25.py", "configured_tp_size"): (
|
||||||
|
"the IPC refcount has to name the same number the recycler waits on, and "
|
||||||
|
"that waiter (MmItemMemoryPool.try_to_recycle) reads configured_tp_size() "
|
||||||
|
"because it runs in the tokenizer process; a refcount taken from the live "
|
||||||
|
"attention subgroup would strand items in the bounded pool"
|
||||||
|
),
|
||||||
|
("srt/models/kimi_k3.py", "configured_tp_size"): (
|
||||||
|
"same as kimi_k25: the IPC refcount must agree with the recycler's waiter"
|
||||||
|
),
|
||||||
|
}
|
||||||
|
|
||||||
|
# A dynamic read whose name is set nowhere in the tree, so the predicate it
|
||||||
|
# feeds is inert (the ``getattr`` default decides it). Converting it would mean
|
||||||
|
# choosing what it should have named, which is the CP path's call, not this
|
||||||
|
# sweep's -- so it is listed here rather than silently counted or "fixed".
|
||||||
|
_INERT_DYNAMIC_READS = frozenset({("srt/layers/cp/base.py", "_is_dsa_model_arch")})
|
||||||
|
|
||||||
_DIRECT_BASELINE = 0
|
_DIRECT_BASELINE = 0
|
||||||
_ALIAS_BASELINE = 0
|
_ALIAS_BASELINE = 0
|
||||||
|
|
||||||
|
|
||||||
def _is_global_call(node) -> bool:
|
def _is_global_call(node) -> bool:
|
||||||
return (
|
"""``get_server_args()`` however it is spelled: bare, or module-qualified
|
||||||
isinstance(node, ast.Call)
|
(``ctx.get_server_args()``), which an ast.Name check alone would miss."""
|
||||||
and isinstance(node.func, ast.Name)
|
if not isinstance(node, ast.Call):
|
||||||
and node.func.id == "get_server_args"
|
return False
|
||||||
)
|
func = node.func
|
||||||
|
if isinstance(func, ast.Name):
|
||||||
|
return func.id == "get_server_args"
|
||||||
|
return isinstance(func, ast.Attribute) and func.attr == "get_server_args"
|
||||||
|
|
||||||
|
|
||||||
def _collect(rel: str, tree: ast.AST):
|
def _collect(rel: str, tree: ast.AST, inert: frozenset = frozenset()):
|
||||||
"""The (direct, alias) field reads in one module."""
|
"""The (direct, alias) field reads in one module.
|
||||||
|
|
||||||
|
``inert`` names the fields listed in ``_INERT_DYNAMIC_READS`` for this file;
|
||||||
|
they are dropped here, at the point the read is recognized, so the filter
|
||||||
|
matches on the field name rather than on the rendered message.
|
||||||
|
"""
|
||||||
direct, alias = [], []
|
direct, alias = [], []
|
||||||
|
|
||||||
def counted(attr: str) -> bool:
|
def counted(attr: str) -> bool:
|
||||||
return attr not in _DERIVED_MEMBERS and (rel, attr) not in _CONFIG_INTENT_SIZES
|
return attr not in inert
|
||||||
|
|
||||||
|
def _getattr_name(node):
|
||||||
|
"""``getattr(<record>, "field")`` names a field just as ``.field`` does;
|
||||||
|
matching only ast.Attribute would let a dynamic read walk past."""
|
||||||
|
if not (
|
||||||
|
isinstance(node, ast.Call)
|
||||||
|
and isinstance(node.func, ast.Name)
|
||||||
|
and node.func.id == "getattr"
|
||||||
|
and len(node.args) >= 2
|
||||||
|
and isinstance(node.args[1], ast.Constant)
|
||||||
|
and isinstance(node.args[1].value, str)
|
||||||
|
):
|
||||||
|
return None
|
||||||
|
return node.args[1].value
|
||||||
|
|
||||||
for node in ast.walk(tree):
|
for node in ast.walk(tree):
|
||||||
if (
|
if (
|
||||||
@@ -105,17 +166,54 @@ def _collect(rel: str, tree: ast.AST):
|
|||||||
):
|
):
|
||||||
direct.append(f"{rel}:{node.lineno}: get_server_args().{node.attr}")
|
direct.append(f"{rel}:{node.lineno}: get_server_args().{node.attr}")
|
||||||
|
|
||||||
|
name = _getattr_name(node)
|
||||||
|
if name is not None and _is_global_call(node.args[0]) and counted(name):
|
||||||
|
direct.append(f"{rel}:{node.lineno}: getattr(get_server_args(), {name!r})")
|
||||||
|
|
||||||
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||||
continue
|
continue
|
||||||
params = {a.arg for a in list(node.args.args) + list(node.args.kwonlyargs)}
|
params = {a.arg for a in list(node.args.args) + list(node.args.kwonlyargs)}
|
||||||
bound = {}
|
bound = {}
|
||||||
for inner in ast.walk(node):
|
for inner in ast.walk(node):
|
||||||
if isinstance(inner, ast.Assign) and _is_global_call(inner.value):
|
# ``sa = get_server_args()`` and its annotated form
|
||||||
for target in inner.targets:
|
# ``sa: ServerArgs = get_server_args()``.
|
||||||
if isinstance(target, ast.Name) and target.id not in params:
|
if isinstance(inner, (ast.Assign, ast.AnnAssign)) and _is_global_call(
|
||||||
bound.setdefault(target.id, inner.lineno)
|
getattr(inner, "value", None)
|
||||||
|
):
|
||||||
|
targets = (
|
||||||
|
inner.targets if isinstance(inner, ast.Assign) else [inner.target]
|
||||||
|
)
|
||||||
|
for target in targets:
|
||||||
|
if not isinstance(target, ast.Name):
|
||||||
|
continue
|
||||||
|
# A parameter reassigned from the global is the
|
||||||
|
# optional-injection shape (``f(server_args=None)`` then
|
||||||
|
# ``server_args = get_server_args()``): the reads that
|
||||||
|
# follow are global reads wearing a parameter's name, so
|
||||||
|
# they count from the bind on.
|
||||||
|
bound.setdefault(target.id, inner.lineno)
|
||||||
if not bound:
|
if not bound:
|
||||||
continue
|
continue
|
||||||
|
# A copy of an alias reaches the same record (``cfg = sa`` after
|
||||||
|
# ``sa = get_server_args()``), so follow Name-to-Name assignments to a
|
||||||
|
# fixpoint. Deeper indirection (through containers, attributes of
|
||||||
|
# other objects, cross-scope copies) stays census-tool territory.
|
||||||
|
changed = True
|
||||||
|
while changed:
|
||||||
|
changed = False
|
||||||
|
for inner in ast.walk(node):
|
||||||
|
if not isinstance(inner, (ast.Assign, ast.AnnAssign)):
|
||||||
|
continue
|
||||||
|
value = getattr(inner, "value", None)
|
||||||
|
if not (isinstance(value, ast.Name) and value.id in bound):
|
||||||
|
continue
|
||||||
|
targets = (
|
||||||
|
inner.targets if isinstance(inner, ast.Assign) else [inner.target]
|
||||||
|
)
|
||||||
|
for target in targets:
|
||||||
|
if isinstance(target, ast.Name) and target.id not in bound:
|
||||||
|
bound[target.id] = inner.lineno
|
||||||
|
changed = True
|
||||||
for inner in ast.walk(node):
|
for inner in ast.walk(node):
|
||||||
if (
|
if (
|
||||||
isinstance(inner, ast.Attribute)
|
isinstance(inner, ast.Attribute)
|
||||||
@@ -128,6 +226,170 @@ def _collect(rel: str, tree: ast.AST):
|
|||||||
f"{rel}:{inner.lineno}: {inner.value.id}.{inner.attr} "
|
f"{rel}:{inner.lineno}: {inner.value.id}.{inner.attr} "
|
||||||
f"(bound from get_server_args() at line {bound[inner.value.id]})"
|
f"(bound from get_server_args() at line {bound[inner.value.id]})"
|
||||||
)
|
)
|
||||||
|
name = _getattr_name(inner)
|
||||||
|
if (
|
||||||
|
name is not None
|
||||||
|
and isinstance(inner.args[0], ast.Name)
|
||||||
|
and inner.args[0].id in bound
|
||||||
|
and inner.lineno >= bound[inner.args[0].id]
|
||||||
|
and counted(name)
|
||||||
|
):
|
||||||
|
alias.append(
|
||||||
|
f"{rel}:{inner.lineno}: getattr({inner.args[0].id}, {name!r}) "
|
||||||
|
f"(bound from get_server_args() at line {bound[inner.args[0].id]})"
|
||||||
|
)
|
||||||
|
# A module-level alias is visible to every function in the file, so it needs
|
||||||
|
# its own pass -- the per-function scan above deliberately does not reach
|
||||||
|
# across scopes.
|
||||||
|
module_bound = {}
|
||||||
|
module_stack = list(tree.body)
|
||||||
|
while module_stack:
|
||||||
|
stmt = module_stack.pop()
|
||||||
|
# A module-level bind can sit inside an `if` / `try` / `with`, so the
|
||||||
|
# walk descends into those bodies -- but not into a nested function or
|
||||||
|
# class, whose binds are that scope's own.
|
||||||
|
if isinstance(
|
||||||
|
stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef, ast.Lambda)
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
module_stack.extend(ast.iter_child_nodes(stmt))
|
||||||
|
if isinstance(stmt, (ast.Assign, ast.AnnAssign)) and _is_global_call(
|
||||||
|
getattr(stmt, "value", None)
|
||||||
|
):
|
||||||
|
targets = stmt.targets if isinstance(stmt, ast.Assign) else [stmt.target]
|
||||||
|
for target in targets:
|
||||||
|
if isinstance(target, ast.Name):
|
||||||
|
module_bound.setdefault(target.id, stmt.lineno)
|
||||||
|
if module_bound:
|
||||||
|
# Shadowing is per lexical scope: a function with its own `sa` hides the
|
||||||
|
# module alias *inside that function only*. Aggregating the names
|
||||||
|
# file-wide would suppress every read in the module, including the
|
||||||
|
# top-level ones and the ones in functions that do resolve to the alias.
|
||||||
|
parents = {}
|
||||||
|
scope_binds = {}
|
||||||
|
stack = [tree]
|
||||||
|
while stack:
|
||||||
|
node = stack.pop()
|
||||||
|
enclosing = parents.get(id(node))
|
||||||
|
for child in ast.iter_child_nodes(node):
|
||||||
|
parents[id(child)] = (
|
||||||
|
node
|
||||||
|
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef))
|
||||||
|
else enclosing
|
||||||
|
)
|
||||||
|
stack.append(child)
|
||||||
|
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||||
|
names = {
|
||||||
|
a.arg for a in list(node.args.args) + list(node.args.kwonlyargs)
|
||||||
|
}
|
||||||
|
# Only this scope's own stores: a nested function's local `sa`
|
||||||
|
# shadows the alias inside *that* function, not in its parent.
|
||||||
|
pending = list(node.body)
|
||||||
|
while pending:
|
||||||
|
inner = pending.pop()
|
||||||
|
if isinstance(
|
||||||
|
inner,
|
||||||
|
(
|
||||||
|
ast.FunctionDef,
|
||||||
|
ast.AsyncFunctionDef,
|
||||||
|
ast.Lambda,
|
||||||
|
ast.ClassDef,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
if isinstance(inner, ast.Name) and isinstance(inner.ctx, ast.Store):
|
||||||
|
names.add(inner.id)
|
||||||
|
pending.extend(ast.iter_child_nodes(inner))
|
||||||
|
scope_binds[id(node)] = names
|
||||||
|
|
||||||
|
def _shadowed(node, name):
|
||||||
|
scope = parents.get(id(node))
|
||||||
|
while scope is not None:
|
||||||
|
if name in scope_binds.get(id(scope), ()):
|
||||||
|
return True
|
||||||
|
scope = parents.get(id(scope))
|
||||||
|
return False
|
||||||
|
|
||||||
|
for node in ast.walk(tree):
|
||||||
|
base = attr = None
|
||||||
|
if (
|
||||||
|
isinstance(node, ast.Attribute)
|
||||||
|
and isinstance(node.value, ast.Name)
|
||||||
|
and node.value.id in module_bound
|
||||||
|
):
|
||||||
|
base, attr = node.value.id, node.attr
|
||||||
|
shown = f"{base}.{attr}"
|
||||||
|
else:
|
||||||
|
attr_name = _getattr_name(node)
|
||||||
|
if (
|
||||||
|
attr_name is not None
|
||||||
|
and isinstance(node.args[0], ast.Name)
|
||||||
|
and node.args[0].id in module_bound
|
||||||
|
):
|
||||||
|
base, attr = node.args[0].id, attr_name
|
||||||
|
shown = f"getattr({base}, {attr!r})"
|
||||||
|
if base and not _shadowed(node, base) and counted(attr):
|
||||||
|
alias.append(
|
||||||
|
f"{rel}:{node.lineno}: {shown} "
|
||||||
|
f"(module-level bind from get_server_args() at line "
|
||||||
|
f"{module_bound[base]})"
|
||||||
|
)
|
||||||
|
|
||||||
|
# An alias parked on an instance attribute (``self._sa = get_server_args()``
|
||||||
|
# in one method, ``self._sa.field`` in another) reaches the same slot and
|
||||||
|
# crosses function scopes, so it is collected per class rather than per
|
||||||
|
# function.
|
||||||
|
for node in ast.walk(tree):
|
||||||
|
if not isinstance(node, ast.ClassDef):
|
||||||
|
continue
|
||||||
|
attr_bound = {}
|
||||||
|
for inner in ast.walk(node):
|
||||||
|
if isinstance(inner, (ast.Assign, ast.AnnAssign)) and _is_global_call(
|
||||||
|
getattr(inner, "value", None)
|
||||||
|
):
|
||||||
|
targets = (
|
||||||
|
inner.targets if isinstance(inner, ast.Assign) else [inner.target]
|
||||||
|
)
|
||||||
|
for target in targets:
|
||||||
|
if (
|
||||||
|
isinstance(target, ast.Attribute)
|
||||||
|
and isinstance(target.value, ast.Name)
|
||||||
|
and target.value.id in ("self", "cls")
|
||||||
|
):
|
||||||
|
attr_bound.setdefault(
|
||||||
|
(target.value.id, target.attr), inner.lineno
|
||||||
|
)
|
||||||
|
if not attr_bound:
|
||||||
|
continue
|
||||||
|
|
||||||
|
def _bound_attr(value):
|
||||||
|
"""``self._sa`` when that attribute was bound from the global."""
|
||||||
|
if (
|
||||||
|
isinstance(value, ast.Attribute)
|
||||||
|
and isinstance(value.value, ast.Name)
|
||||||
|
and (value.value.id, value.attr) in attr_bound
|
||||||
|
):
|
||||||
|
return (value.value.id, value.attr)
|
||||||
|
return None
|
||||||
|
|
||||||
|
for inner in ast.walk(node):
|
||||||
|
key = shown = None
|
||||||
|
if isinstance(inner, ast.Attribute):
|
||||||
|
key = _bound_attr(inner.value)
|
||||||
|
if key is not None and counted(inner.attr):
|
||||||
|
shown = f"{key[0]}.{key[1]}.{inner.attr}"
|
||||||
|
else:
|
||||||
|
name = _getattr_name(inner)
|
||||||
|
if name is not None:
|
||||||
|
key = _bound_attr(inner.args[0])
|
||||||
|
if key is not None and counted(name):
|
||||||
|
shown = f"getattr({key[0]}.{key[1]}, {name!r})"
|
||||||
|
if shown is not None:
|
||||||
|
alias.append(
|
||||||
|
f"{rel}:{inner.lineno}: {shown} "
|
||||||
|
f"(attribute bind from get_server_args() at line "
|
||||||
|
f"{attr_bound[key]})"
|
||||||
|
)
|
||||||
return direct, alias
|
return direct, alias
|
||||||
|
|
||||||
|
|
||||||
@@ -135,11 +397,14 @@ def _field_reads():
|
|||||||
direct, alias = [], []
|
direct, alias = [], []
|
||||||
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
|
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
|
||||||
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
|
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
|
||||||
|
if rel.startswith(_SLOT_OWNERS):
|
||||||
|
continue
|
||||||
try:
|
try:
|
||||||
tree = ast.parse(path.read_text())
|
tree = ast.parse(path.read_text())
|
||||||
except SyntaxError:
|
except SyntaxError:
|
||||||
continue
|
continue
|
||||||
module_direct, module_alias = _collect(rel, tree)
|
inert = frozenset(name for path_, name in _INERT_DYNAMIC_READS if path_ == rel)
|
||||||
|
module_direct, module_alias = _collect(rel, tree, inert)
|
||||||
direct += module_direct
|
direct += module_direct
|
||||||
alias += module_alias
|
alias += module_alias
|
||||||
return direct, alias
|
return direct, alias
|
||||||
@@ -167,5 +432,90 @@ class TestGlobalConfigReadRatchet(CustomTestCase):
|
|||||||
self._check("alias-form", alias, _ALIAS_BASELINE)
|
self._check("alias-form", alias, _ALIAS_BASELINE)
|
||||||
|
|
||||||
|
|
||||||
|
class TestConfiguredSizeCallSites(CustomTestCase):
|
||||||
|
"""The configured-vs-live exceptions are enumerated, with reasons.
|
||||||
|
|
||||||
|
``configured_*_size()`` answers what the user asked for where
|
||||||
|
``get_parallel()`` would answer what the process ended up with. Each such
|
||||||
|
exception is listed above with why the live property cannot serve it, and
|
||||||
|
this case fails if the code and that list disagree.
|
||||||
|
|
||||||
|
The unit is **(file, accessor)**, not the individual call: a second
|
||||||
|
`configured_pp_size()` in a file already registered for it collapses into
|
||||||
|
the same entry, so the reason has to cover the file's use of that accessor
|
||||||
|
rather than one line. A new file, or a new accessor in a listed file, is
|
||||||
|
what this catches -- in either call form (bare or module-qualified).
|
||||||
|
"""
|
||||||
|
|
||||||
|
def test_the_call_sites_match_the_documented_set(self):
|
||||||
|
found = set()
|
||||||
|
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
|
||||||
|
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
|
||||||
|
if rel.startswith(_SLOT_OWNERS):
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
tree = ast.parse(path.read_text())
|
||||||
|
except SyntaxError:
|
||||||
|
continue
|
||||||
|
for node in ast.walk(tree):
|
||||||
|
if not isinstance(node, ast.Call):
|
||||||
|
continue
|
||||||
|
func = node.func
|
||||||
|
name = (
|
||||||
|
func.id
|
||||||
|
if isinstance(func, ast.Name)
|
||||||
|
else (func.attr if isinstance(func, ast.Attribute) else None)
|
||||||
|
)
|
||||||
|
if name and name.startswith("configured_") and name.endswith("_size"):
|
||||||
|
found.add((rel, name))
|
||||||
|
documented = set(_CONFIGURED_SIZE_CALL_SITES)
|
||||||
|
self.assertEqual(
|
||||||
|
documented,
|
||||||
|
found,
|
||||||
|
"configured-size call sites drifted from their documented reasons.\n"
|
||||||
|
f" undocumented: {sorted(found - documented)}\n"
|
||||||
|
f" stale entries: {sorted(documented - found)}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestNoRenamedAccessorImports(CustomTestCase):
|
||||||
|
"""The scanners above match ``get_server_args`` and ``configured_*_size``
|
||||||
|
by their literal names, so an ``import ... as`` rename would walk a read
|
||||||
|
straight past both the zero baseline and the call-site registry. Renaming
|
||||||
|
these accessors buys nothing (the names are already short and unambiguous),
|
||||||
|
so it is banned outright — which is exactly what makes literal-name
|
||||||
|
matching sound."""
|
||||||
|
|
||||||
|
def test_the_scanned_accessors_are_never_import_renamed(self):
|
||||||
|
offenders = []
|
||||||
|
for path in sorted(_PACKAGE_ROOT.rglob("*.py")):
|
||||||
|
rel = path.relative_to(_PACKAGE_ROOT).as_posix()
|
||||||
|
try:
|
||||||
|
tree = ast.parse(path.read_text())
|
||||||
|
except SyntaxError:
|
||||||
|
continue
|
||||||
|
for node in ast.walk(tree):
|
||||||
|
if not isinstance(node, (ast.ImportFrom, ast.Import)):
|
||||||
|
continue
|
||||||
|
for imported in node.names:
|
||||||
|
if imported.asname is None or imported.asname == imported.name:
|
||||||
|
continue
|
||||||
|
base = imported.name.rsplit(".", 1)[-1]
|
||||||
|
if base == "get_server_args" or (
|
||||||
|
base.startswith("configured_") and base.endswith("_size")
|
||||||
|
):
|
||||||
|
offenders.append(
|
||||||
|
f"{rel}:{node.lineno}: {imported.name} as "
|
||||||
|
f"{imported.asname}"
|
||||||
|
)
|
||||||
|
self.assertFalse(
|
||||||
|
offenders,
|
||||||
|
"get_server_args / configured_*_size imported under another name; "
|
||||||
|
"the read ratchet and the configured-size registry match these "
|
||||||
|
"accessors by their literal names, so a rename silently escapes "
|
||||||
|
"both:\n" + "\n".join(offenders),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -5,7 +5,10 @@ from sglang.test.ci.ci_register import register_cpu_ci
|
|||||||
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
|
import json
|
||||||
import os
|
import os
|
||||||
|
import shutil
|
||||||
|
import tempfile
|
||||||
import unittest
|
import unittest
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
@@ -20,8 +23,10 @@ from sglang.srt.runtime_context import (
|
|||||||
get_flags,
|
get_flags,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
|
max_speculative_num_draft_tokens,
|
||||||
reset_context,
|
reset_context,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
from sglang.test.test_utils import CustomTestCase
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
_PS = "sglang.srt.distributed.parallel_state"
|
_PS = "sglang.srt.distributed.parallel_state"
|
||||||
@@ -378,6 +383,14 @@ class _FakeResolvedArgs:
|
|||||||
sampling_backend: A[
|
sampling_backend: A[
|
||||||
str | None, Arg(help="s", resolvable=True), NS("exec.kernel")
|
str | None, Arg(help="s", resolvable=True), NS("exec.kernel")
|
||||||
] = None
|
] = None
|
||||||
|
attention_backend: A[str | None, Arg(help="ab"), NS("exec.kernel")] = None
|
||||||
|
prefill_attention_backend: A[str | None, Arg(help="pab"), NS("exec.kernel")] = None
|
||||||
|
decode_attention_backend: A[str | None, Arg(help="dab"), NS("exec.kernel")] = None
|
||||||
|
disable_radix_cache: A[bool, Arg(help="drc"), NS("memory")] = False
|
||||||
|
mamba_radix_cache_strategy: A[str, Arg(help="mrcs"), NS("exec.mamba")] = "auto"
|
||||||
|
speculative_num_draft_tokens: A[int | None, Arg(help="d"), NS("spec")] = None
|
||||||
|
speculative_adaptive: A[bool, Arg(help="a"), NS("spec")] = False
|
||||||
|
speculative_adaptive_config: A[str | None, Arg(help="c"), NS("spec")] = None
|
||||||
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
_resolved_overrides: list = dataclasses.field(default_factory=list)
|
||||||
|
|
||||||
|
|
||||||
@@ -965,5 +978,180 @@ class TestPublishLifecycle(_IsolatedServerArgs):
|
|||||||
self.assertFalse(get_flags().capture.enable_torch_compile)
|
self.assertFalse(get_flags().capture.enable_torch_compile)
|
||||||
|
|
||||||
|
|
||||||
|
class TestDerivedPredicatesAgreeAcrossTiers(_IsolatedServerArgs):
|
||||||
|
"""One definition per predicate, checked rather than asserted in prose.
|
||||||
|
|
||||||
|
Each of these exists twice by construction -- once over a config-shaped
|
||||||
|
object (the resolution pipeline's `*_of` helper, which `ServerArgs`
|
||||||
|
delegates to) and once over the published bags. The pair must agree on
|
||||||
|
every input, or a decision made before publish differs from the same
|
||||||
|
decision made after it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
_STRATEGIES = ("auto", "no_buffer", "extra_buffer", "extra_buffer_lazy")
|
||||||
|
|
||||||
|
def test_mamba_extra_buffer_matches_the_member(self):
|
||||||
|
from sglang.srt.runtime_context import (
|
||||||
|
mamba_extra_buffer_enabled,
|
||||||
|
mamba_extra_buffer_lazy_enabled,
|
||||||
|
)
|
||||||
|
|
||||||
|
for disable_radix_cache in (False, True):
|
||||||
|
for strategy in self._STRATEGIES:
|
||||||
|
with self.subTest(radix=disable_radix_cache, strategy=strategy):
|
||||||
|
args = _FakeResolvedArgs(
|
||||||
|
disable_radix_cache=disable_radix_cache,
|
||||||
|
mamba_radix_cache_strategy=strategy,
|
||||||
|
)
|
||||||
|
get_context().set_server_args(args)
|
||||||
|
self.assertEqual(
|
||||||
|
ServerArgs.enable_mamba_extra_buffer(args),
|
||||||
|
mamba_extra_buffer_enabled(),
|
||||||
|
)
|
||||||
|
self.assertEqual(
|
||||||
|
ServerArgs.enable_mamba_extra_buffer_lazy(args),
|
||||||
|
mamba_extra_buffer_lazy_enabled(),
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_attention_backends_match_the_member(self):
|
||||||
|
from sglang.srt.runtime_context import attention_backends
|
||||||
|
|
||||||
|
backends = (None, "fa3", "triton")
|
||||||
|
for base in backends:
|
||||||
|
for prefill in backends:
|
||||||
|
for decode in backends:
|
||||||
|
with self.subTest(base=base, prefill=prefill, decode=decode):
|
||||||
|
args = _FakeResolvedArgs(
|
||||||
|
attention_backend=base,
|
||||||
|
prefill_attention_backend=prefill,
|
||||||
|
decode_attention_backend=decode,
|
||||||
|
)
|
||||||
|
get_context().set_server_args(args)
|
||||||
|
self.assertEqual(
|
||||||
|
ServerArgs.get_attention_backends(args),
|
||||||
|
attention_backends(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestAdaptiveDraftBoundLifecycle(_IsolatedServerArgs):
|
||||||
|
"""The adaptive draft-token bound is memoized on the config path, so the
|
||||||
|
memo has to end with the publication it was computed under.
|
||||||
|
|
||||||
|
Without that, a process that republishes with the same adaptive-config path
|
||||||
|
-- the file having been rewritten in between -- keeps the previous bound and
|
||||||
|
under-allocates the draft-token buffers sized from it.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _write_config(self, steps):
|
||||||
|
path = os.path.join(tempfile.mkdtemp(prefix="adaptive_cfg_"), "adaptive.json")
|
||||||
|
self.addCleanup(shutil.rmtree, os.path.dirname(path), ignore_errors=True)
|
||||||
|
with open(path, "w") as handle:
|
||||||
|
json.dump({"1": {"candidate_steps": steps}}, handle)
|
||||||
|
return path
|
||||||
|
|
||||||
|
def test_republishing_recomputes_the_bound(self):
|
||||||
|
path = self._write_config([2])
|
||||||
|
get_context().set_server_args(
|
||||||
|
_FakeResolvedArgs(
|
||||||
|
speculative_num_draft_tokens=3,
|
||||||
|
speculative_adaptive=True,
|
||||||
|
speculative_adaptive_config=path,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(max_speculative_num_draft_tokens(), 3)
|
||||||
|
|
||||||
|
with open(path, "w") as handle:
|
||||||
|
json.dump({"1": {"candidate_steps": [4]}}, handle)
|
||||||
|
# Same path, new contents: the memo must not survive the republish.
|
||||||
|
get_context().set_server_args(
|
||||||
|
_FakeResolvedArgs(
|
||||||
|
speculative_num_draft_tokens=3,
|
||||||
|
speculative_adaptive=True,
|
||||||
|
speculative_adaptive_config=path,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(max_speculative_num_draft_tokens(), 5)
|
||||||
|
|
||||||
|
def test_reset_clears_the_bound(self):
|
||||||
|
path = self._write_config([2])
|
||||||
|
get_context().set_server_args(
|
||||||
|
_FakeResolvedArgs(
|
||||||
|
speculative_num_draft_tokens=3,
|
||||||
|
speculative_adaptive=True,
|
||||||
|
speculative_adaptive_config=path,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(max_speculative_num_draft_tokens(), 3)
|
||||||
|
reset_context()
|
||||||
|
with open(path, "w") as handle:
|
||||||
|
json.dump({"1": {"candidate_steps": [6]}}, handle)
|
||||||
|
get_context().set_server_args(
|
||||||
|
_FakeResolvedArgs(
|
||||||
|
speculative_num_draft_tokens=3,
|
||||||
|
speculative_adaptive=True,
|
||||||
|
speculative_adaptive_config=path,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
self.assertEqual(max_speculative_num_draft_tokens(), 7)
|
||||||
|
|
||||||
|
|
||||||
|
class TestNamedAccessorsCallWhatTheyWrap(CustomTestCase):
|
||||||
|
"""A named accessor must *call* a member that is a method.
|
||||||
|
|
||||||
|
`return get_server_args().x` hands back a bound method when `x` is defined
|
||||||
|
with `def`; the failure then lands far away, in whatever arithmetic the
|
||||||
|
caller does with it. Checked statically so accessors that need a real model
|
||||||
|
config are covered too.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def test_accessors_that_wrap_methods_call_them(self):
|
||||||
|
import ast
|
||||||
|
import functools
|
||||||
|
import inspect
|
||||||
|
|
||||||
|
import sglang.srt.runtime_context as rc
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
tree = ast.parse(inspect.getsource(rc))
|
||||||
|
wrong = []
|
||||||
|
for node in tree.body:
|
||||||
|
if not isinstance(node, ast.FunctionDef):
|
||||||
|
continue
|
||||||
|
for inner in ast.walk(node):
|
||||||
|
if not (isinstance(inner, ast.Return) and inner.value is not None):
|
||||||
|
continue
|
||||||
|
value = inner.value
|
||||||
|
called = isinstance(value, ast.Call)
|
||||||
|
target = value.func if called else value
|
||||||
|
if not (
|
||||||
|
isinstance(target, ast.Attribute)
|
||||||
|
and isinstance(target.value, ast.Call)
|
||||||
|
and isinstance(target.value.func, ast.Name)
|
||||||
|
and target.value.func.id == "get_server_args"
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
member = getattr(ServerArgs, target.attr, None)
|
||||||
|
# A `property` / `functools.cached_property` member is already
|
||||||
|
# evaluated by the attribute access, so it is named here to keep
|
||||||
|
# the failure message from calling it "not a method" -- the fix
|
||||||
|
# for those is the opposite one.
|
||||||
|
kind = (
|
||||||
|
"a property"
|
||||||
|
if isinstance(member, (property, functools.cached_property))
|
||||||
|
else "not a method"
|
||||||
|
)
|
||||||
|
if inspect.isfunction(member) and not called:
|
||||||
|
wrong.append(
|
||||||
|
f"{node.name}(): returns ServerArgs.{target.attr} without "
|
||||||
|
"calling it, so callers get a bound method"
|
||||||
|
)
|
||||||
|
if not inspect.isfunction(member) and called:
|
||||||
|
wrong.append(
|
||||||
|
f"{node.name}(): calls ServerArgs.{target.attr}, which is "
|
||||||
|
f"{kind} -- the attribute access already produced the value"
|
||||||
|
)
|
||||||
|
self.assertEqual([], wrong, "\n".join(wrong))
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user