config: stop handing the record to code that does not read it (#36252)
This commit is contained in:
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user