config: stop handing the record to code that does not read it (#36252)

This commit is contained in:
Cheng Wan
2026-08-26 05:02:17 -07:00
committed by GitHub
parent d7b144f64e
commit ae5feb4b9c
60 changed files with 168 additions and 280 deletions
@@ -598,10 +598,10 @@ class CPUGraphRunner:
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_gathered_buffer = require_gathered_buffer(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_mlp_sync = require_mlp_sync(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_gathered_buffer = require_gathered_buffer()
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_mlp_sync = require_mlp_sync()
self.require_attn_tp_gather = require_attn_tp_gather()
self.enable_two_batch_overlap = (
model_runner.server_args.enable_two_batch_overlap
)
@@ -986,14 +986,14 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
pin_memory=is_pin_memory_available(batch.device),
).to(batch.device, non_blocking=True)
def adjust_num_token_non_padded_for_attn_tp(self, server_args) -> None:
def adjust_num_token_non_padded_for_attn_tp(self) -> None:
"""Make num_token_non_padded local to this attention-TP rank."""
from sglang.srt.utils.common import require_mlp_tp_gather
dp_rank = get_parallel().attn_dp_rank
assert self.global_num_tokens_cpu is not None
if require_mlp_tp_gather(server_args):
if require_mlp_tp_gather():
num_tokens_per_dp = self.global_num_tokens_cpu[dp_rank]
else:
num_tokens_per_dp = self.global_num_tokens_cpu[0]
@@ -223,7 +223,7 @@ from sglang.srt.utils.device_timer import device_timer_ctx
from sglang.srt.utils.nvtx_pytorch_hooks import PytHooks
from sglang.srt.utils.nvtx_utils import profile_range
from sglang.srt.utils.offloader import (
create_offloader_from_server_args,
create_offloader,
get_offloader,
set_offloader,
)
@@ -270,7 +270,6 @@ class ModelRunnerOutput:
def resolve_draft_attention_backend(
*,
draft_attention_backend: Optional[str],
server_args: ServerArgs,
is_draft_worker: bool,
) -> Optional[str]:
"""The attention backend a runner uses because it is a draft runner.
@@ -342,7 +341,6 @@ class ModelRunner:
self.device = get_device().device
self.draft_attention_backend = resolve_draft_attention_backend(
draft_attention_backend=draft_attention_backend,
server_args=server_args,
is_draft_worker=is_draft_worker,
)
# This runner's own load format, resolved before anything keys off it:
@@ -430,9 +428,7 @@ class ModelRunner:
self.shared_read_done_event: Optional[torch.cuda.Event] = None
# CPU offload
set_offloader(
create_offloader_from_server_args(server_args, dp_rank=self.ps.dp_rank)
)
set_offloader(create_offloader(dp_rank=self.ps.dp_rank))
self._weight_checker = WeightChecker(get_model=lambda: self.model, ps=self.ps)
@@ -591,7 +587,6 @@ class ModelRunner:
model=self.model,
model_config=self.model_config,
req_to_token_pool=self.req_to_token_pool,
server_args=self.server_args,
max_running_requests=self.max_running_requests,
device=self.device,
)
@@ -702,7 +697,6 @@ class ModelRunner:
)
set_global_expert_location_metadata(
compute_initial_expert_location_metadata(
server_args=self.server_args,
model_config=self.model_config,
moe_ep_rank=expert_rank,
)
@@ -727,7 +721,6 @@ class ModelRunner:
def maybe_init_eplb_manager(self):
self.eplb_manager = (
EPLBManager(
server_args=self.server_args,
model_config=self.model_config,
ps=self.ps,
get_model=lambda: self.model,
@@ -1076,9 +1069,7 @@ class ModelRunner:
self.pre_model_load_memory = result.pre_model_load_memory
def init_shared_mooncake_transfer_engine(self):
maybe_init_shared_mooncake_transfer_engine(
server_args=self.server_args, gpu_id=self.gpu_id
)
maybe_init_shared_mooncake_transfer_engine(gpu_id=self.gpu_id)
def load_model(self):
tic_total = time.perf_counter()
@@ -1091,9 +1082,7 @@ class ModelRunner:
if self.device != "cpu":
torch.set_num_threads(1)
if self.device == "cuda":
maybe_downgrade_dtype_for_legacy_gpu(
server_args=self.server_args, model_config=self.model_config
)
maybe_downgrade_dtype_for_legacy_gpu(model_config=self.model_config)
set_cuda_arch()
@@ -1113,7 +1102,6 @@ class ModelRunner:
# and derive the per-rank daemon socket. Idempotent across reloads.
maybe_enable_ipc_weight_cache(
load_config=self.load_config,
server_args=self.server_args,
tp_size=self.ps.tp_size,
pp_rank=self.ps.pp_rank,
tp_rank=self.ps.tp_rank,
@@ -1124,7 +1112,6 @@ class ModelRunner:
)
maybe_trigger_remote_instance_nccl_send_group(
server_args=self.server_args,
tp_rank=self.ps.tp_rank,
load_format=draft_load_format,
)
@@ -1195,11 +1182,12 @@ class ModelRunner:
f"mem usage={self.weight_load_mem_usage:.2f} GB."
)
report_online_quantization(model=self.model, server_args=self.server_args)
report_online_quantization(
model=self.model,
)
maybe_register_debug_tensor_dump_hook(
model=self.model,
server_args=self.server_args,
spec_algorithm=self.spec_algorithm,
is_draft_worker=self.is_draft_worker,
tp_size=self.ps.tp_size,
@@ -1281,7 +1269,6 @@ class ModelRunner:
)
if not cuda_graph_fully_disabled():
init_lora_cuda_graph_moe_buffers(
server_args=self.server_args,
model=self.model,
lora_manager=self.lora_manager,
dtype=self.dtype,
@@ -1471,13 +1458,11 @@ class ModelRunner:
if (
forward_batch.num_token_non_padded is not None
and forward_batch.global_num_tokens_gpu is not None
and require_gathered_buffer(self.server_args)
and require_gathered_buffer()
and not is_dsa_enable_prefill_cp()
and not is_mla_prefill_cp_enabled()
):
forward_batch.adjust_num_token_non_padded_for_attn_tp(
server_args=self.server_args,
)
forward_batch.adjust_num_token_non_padded_for_attn_tp()
# Hisparse coordinator — backends now read it from self.model_runner.
if self.hisparse_coordinator is not None:
@@ -1926,7 +1911,6 @@ class ModelRunner:
start=old_num_physical - num_local * initial_ep_size,
)
new_metadata = ExpertLocationMetadata.init_by_mapping(
self.server_args,
self.model_config,
physical_to_logical_map=expanded_p2l,
moe_ep_rank=self._elastic_global_rank(),
@@ -67,9 +67,7 @@ class LoadedModel(msgspec.Struct, frozen=True, kw_only=True):
startup_weight_load: Optional[Any] = None
def maybe_downgrade_dtype_for_legacy_gpu(
*, server_args: ServerArgs, model_config: ModelConfig
) -> None:
def maybe_downgrade_dtype_for_legacy_gpu(*, model_config: ModelConfig) -> None:
if torch.cuda.get_device_capability()[0] < 8:
logger.info(
"Compute capability below sm80. Use float16 due to lack of bfloat16 support."
@@ -85,7 +83,7 @@ def maybe_downgrade_dtype_for_legacy_gpu(
def maybe_trigger_remote_instance_nccl_send_group(
*, server_args: ServerArgs, tp_rank: int, load_format: Optional[str] = None
*, tp_rank: int, load_format: Optional[str] = None
) -> None:
"""``load_format`` is this runner's effective format: a draft loading under
``--speculative-draft-draft-load-format`` needs its own send group, and the
@@ -151,7 +149,7 @@ def resolve_sliding_window_size(model, model_config: ModelConfig) -> Optional[in
return sliding_window_size
def report_online_quantization(*, model, server_args: ServerArgs) -> None:
def report_online_quantization(*, model) -> None:
# TODO: Make sure all models have `quant_config` attribute, and all online quantization methods register which layers they actually quantize.
quantized_layers = getattr(
getattr(model, "quant_config", None), "quantized_layers", None
@@ -170,7 +168,6 @@ def report_online_quantization(*, model, server_args: ServerArgs) -> None:
def maybe_register_debug_tensor_dump_hook(
*,
model,
server_args: ServerArgs,
spec_algorithm: SpeculativeAlgorithm,
is_draft_worker: bool,
tp_size: int,
@@ -237,7 +234,6 @@ def build_load_config(
def maybe_enable_ipc_weight_cache(
*,
load_config: LoadConfig,
server_args: ServerArgs,
tp_size: int,
pp_rank: int,
tp_rank: int,
@@ -311,12 +307,11 @@ def load_model_with_memory_saver(
StartupWeightLoadManager,
)
startup_weight_load = StartupWeightLoadManager.create_from_server_args(
startup_weight_load = StartupWeightLoadManager.create_from_published_config(
loader=loader,
model_config=model_config,
load_config=load_config,
device_config=device_config,
server_args=server_args,
is_draft_worker=is_draft_worker,
)
model = startup_weight_load.prepare()
@@ -12,7 +12,6 @@ from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.managers.schedule_batch import ForwardMode
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.runtime_context import get_schedule
from sglang.srt.server_args import ServerArgs
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
@@ -33,7 +32,6 @@ class NgramEmbeddingManager:
model: torch.nn.Module,
model_config: ModelConfig,
req_to_token_pool: ReqToTokenPool,
server_args: ServerArgs,
max_running_requests: int,
device: str,
):
@@ -203,7 +203,6 @@ def _resolve_dflash_aux_hidden_state(
config.dflash_draft_num_layers = int(draft_num_layers)
config.dflash_target_layer_ids = target_layer_ids
config.dflash_draft_cell_size_per_token = _resolve_dflash_draft_cell_size(
server_args=server_args,
draft_model_config=draft_model_config,
draft_num_layers=int(draft_num_layers),
)
@@ -211,7 +210,6 @@ def _resolve_dflash_aux_hidden_state(
def _resolve_dflash_draft_cell_size(
*,
server_args: ServerArgs,
draft_model_config: ModelConfig,
draft_num_layers: int,
) -> int | None:
@@ -31,7 +31,6 @@ from sglang.srt.runtime_context import (
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -98,12 +97,16 @@ class StartupWeightLoadOptions:
prefetch_num_threads: int
@classmethod
def from_server_args(
def from_published_config(
cls,
*,
server_args: ServerArgs,
is_draft_worker: bool,
) -> StartupWeightLoadOptions:
"""Everything this needs is a published leaf; nothing comes off a record.
`is_draft_worker` is the exception and travels as an argument: it is
this runner's role, not the process's configuration.
"""
cuda_graph_config = get_exec().graph.cuda_graph_config
cuda_graph_enabled = any(
getattr(cuda_graph_config, phase).backend != Backend.DISABLED
@@ -267,29 +270,27 @@ class StartupWeightLoadManager:
self._prefetch_failure_reported = False
@classmethod
def create_from_server_args(
def create_from_published_config(
cls,
*,
loader,
model_config: ModelConfig,
load_config: LoadConfig,
device_config: DeviceConfig,
server_args: ServerArgs,
is_draft_worker: bool,
) -> StartupWeightLoadManager:
"""Build a manager straight from ``ServerArgs``.
"""Build a manager from the published configuration.
Callers on the model-loading path only decide *whether* to overlap; the
knowledge of which server arguments matter, and every support rule,
stays in this module.
knowledge of which config leaves matter, and every support rule, stays
in this module.
"""
return cls.create(
loader=loader,
model_config=model_config,
load_config=load_config,
device_config=device_config,
options=StartupWeightLoadOptions.from_server_args(
server_args=server_args,
options=StartupWeightLoadOptions.from_published_config(
is_draft_worker=is_draft_worker,
),
)
@@ -259,7 +259,6 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
calculate_mla_kv_cache_dim(
model_config=model_config,
kv_cache_dtype=kv_cache_dtype,
server_args=kvc.server_args,
)
* effective_num_layers
* kv_size
@@ -465,7 +464,6 @@ class HybridSWAPoolConfigurator(MemoryPoolConfigurator):
calculate_mla_kv_cache_dim(
model_config=model_config,
kv_cache_dtype=kv_cache_dtype,
server_args=kvc.server_args,
)
* kv_size
)
@@ -70,11 +70,10 @@ def get_batch_sizes_to_capture(
constraints and clamps to req_to_token_pool.size.
"""
server_args = model_runner.server_args
capture_bs = list(get_exec().graph.cuda_graph_config.decode.bs)
num_max_requests = model_runner.req_to_token_pool.size
mul_base = get_cuda_graph_batch_size_alignment(server_args)
mul_base = get_cuda_graph_batch_size_alignment()
# TBO splits each request's rows across two micro-batches, so the
# alignment constraint applies per request rather than per token row.
alignment_width = captured_req_width
@@ -82,7 +81,7 @@ def get_batch_sizes_to_capture(
alignment_width = 1
# pad `num_max_requests` to avoid being filtered out
num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests)
num_max_requests = get_cuda_graph_max_batch_size(num_max_requests)
if max(capture_bs) > num_max_requests:
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
# is very small. We add more values here to make sure we capture the maximum bs.
@@ -351,7 +351,7 @@ class BaseRunner(ABC):
dp_size=get_parallel().config.dp_size,
pp_size=get_parallel().config.pp_size,
is_encoder_decoder=mr.model_config.is_encoder_decoder,
require_mlp_tp_gather=require_mlp_tp_gather(mr.server_args),
require_mlp_tp_gather=require_mlp_tp_gather(),
seq_len_fill_value=mr.attn_backend.get_cuda_graph_seq_len_fill_value(),
encoder_len_fill_value=(
getattr(mr.model_config.hf_config, "max_source_positions", 0)
@@ -535,9 +535,9 @@ class BaseRunner(ABC):
)
# TP-gather requirements for global token metadata.
require_mlp_tp_gather_ = require_mlp_tp_gather(mr.server_args)
require_attn_tp_gather_ = require_attn_tp_gather(mr.server_args)
if require_gathered_buffer(mr.server_args):
require_mlp_tp_gather_ = require_mlp_tp_gather()
require_attn_tp_gather_ = require_attn_tp_gather()
if require_gathered_buffer():
assert require_mlp_tp_gather_ or require_attn_tp_gather_
if require_mlp_tp_gather_:
@@ -227,10 +227,10 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
self.enable_torch_compile = get_flags().capture.enable_torch_compile
self.disable_padding = model_runner.server_args.disable_cuda_graph_padding
self.is_encoder_decoder = model_runner.model_config.is_encoder_decoder
self.require_mlp_tp_gather = require_mlp_tp_gather(
model_runner.server_args
) and not self._forward_is_dp_local(model_runner)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_mlp_tp_gather = (
require_mlp_tp_gather() and not self._forward_is_dp_local(model_runner)
)
self.require_attn_tp_gather = require_attn_tp_gather()
# Composite predicates derive from the instance values so the dp-local
# draft exemption above stays consistent (require_gathered_buffer ==
# mlp_tp_gather or attn_tp_gather; require_mlp_sync adds dp attention).
@@ -597,7 +597,7 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
draft_is_deepseek_v4,
)
return not draft_is_deepseek_v4(server_args=model_runner.server_args)
return not draft_is_deepseek_v4()
def _ragged_capture_slots(self, num_tokens: int) -> int:
if envs.SGLANG_TEST_RAGGED_VERIFY_FORCE_UNIFORM_CAPTURE.get():
@@ -113,10 +113,10 @@ class EagerRunner(BaseRunner):
# (expand_for_topk_draft) before the eager fallback.
max_bs *= get_spec().speculative_eagle_topk
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
max_bs = get_eager_max_batch_size(sa, max_bs)
max_bs = get_eager_max_batch_size(max_bs)
prefill_ceiling = max(mr.max_total_num_tokens, max_prefill_buffer_tokens())
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
if require_mlp_sync(sa):
if require_mlp_sync():
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
@@ -332,7 +332,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
embed_dtype=self.model_runner.dtype,
enable_mamba_track=self.mamba_track_enabled,
enable_num_token_non_padded=enable_num_token_non_padded(),
require_gathered_buffer=require_gathered_buffer(model_runner.server_args),
require_gathered_buffer=require_gathered_buffer(),
enable_prefill_cp=(
is_dsa_enable_prefill_cp() or is_mla_prefill_cp_enabled()
),
@@ -349,8 +349,8 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
self.dsa_indexers = getattr(self.model_runner, "dsa_indexers", None)
self.dp_size = get_parallel().config.dp_size
self.require_mlp_tp_gather = require_mlp_tp_gather(model_runner.server_args)
self.require_attn_tp_gather = require_attn_tp_gather(model_runner.server_args)
self.require_mlp_tp_gather = require_mlp_tp_gather()
self.require_attn_tp_gather = require_attn_tp_gather()
# --- backend ---------------------------------------------------
# TcPiecewise resolves by running a compile pass that calls back into
@@ -602,7 +602,7 @@ class PrefillCudaGraphRunner(BaseCudaGraphRunner):
buf = self.buffer_registry.get_slot("num_token_non_padded").buffer
buf.fill_(num_tokens)
if require_gathered_buffer(self.model_runner.server_args):
if require_gathered_buffer():
local = compute_local_num_token_non_padded(
global_num_token_non_padded=buf,
num_tokens_per_dp=num_tokens,