config: the per-instance families read the bags (#35026)

This commit is contained in:
Cheng Wan
2026-08-17 16:17:53 -07:00
committed by GitHub
parent a97bc8db32
commit cba3c5d5ac
45 changed files with 909 additions and 530 deletions
@@ -22,7 +22,12 @@ from typing import Dict, List, NamedTuple, Optional, Tuple
import torch
from sglang.srt.parser.reasoning_parser import ReasoningParser
from sglang.srt.runtime_context import get_context, get_resources
from sglang.srt.runtime_context import (
get_context,
get_exec,
get_resources,
get_serving,
)
from sglang.srt.server_args import ServerArgs
logger = logging.getLogger(__name__)
@@ -315,7 +320,7 @@ def create_grammar_backend(
eos_token_ids: Optional[set] = None,
think_end_ids: Optional[List[int]] = None,
) -> Optional[BaseGrammarBackend]:
name = server_args.grammar_backend
name = get_exec().kernel.grammar_backend
# Custom grammar backend has the highest priority
if name in GRAMMAR_BACKEND_REGISTRY:
@@ -329,7 +334,7 @@ def create_grammar_backend(
grammar_backend = OutlinesGrammarBackend(
tokenizer,
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
whitespace_pattern=get_serving().constrained_json_whitespace_pattern,
)
elif name == "xgrammar":
from sglang.srt.constrained.xgrammar_backend import (
@@ -345,10 +350,10 @@ def create_grammar_backend(
tokenizer,
vocab_size=vocab_size,
model_eos_token_ids=eos_list,
any_whitespace=not server_args.constrained_json_disable_any_whitespace,
any_whitespace=not get_serving().constrained_json_disable_any_whitespace,
)
except TokenizerNotSupportedError as e:
if server_args.enable_strict_thinking:
if get_serving().enable_strict_thinking:
raise ValueError(
f"--enable-strict-thinking requires a grammar backend with "
f"token filtering support, but XGrammar failed to initialize: "
@@ -367,13 +372,13 @@ def create_grammar_backend(
grammar_backend = GuidanceBackend(
tokenizer=tokenizer,
any_whitespace=not server_args.constrained_json_disable_any_whitespace,
whitespace_pattern=server_args.constrained_json_whitespace_pattern,
any_whitespace=not get_serving().constrained_json_disable_any_whitespace,
whitespace_pattern=get_serving().constrained_json_whitespace_pattern,
n_vocab=vocab_size,
eos_token_ids=eos_token_ids,
)
elif name == "none":
if server_args.enable_strict_thinking:
if get_serving().enable_strict_thinking:
raise ValueError(
"--enable-strict-thinking requires a grammar backend that supports "
"token filtering, but grammar_backend='none' was specified. Use "
@@ -384,13 +389,13 @@ def create_grammar_backend(
else:
raise ValueError(f"Invalid grammar backend: {name}")
if server_args.reasoning_parser and think_end_ids:
if get_serving().reasoning_parser and think_end_ids:
from sglang.srt.constrained.reasoner_grammar_backend import (
ReasonerGrammarBackend,
)
reasoning_parser = ReasoningParser(
model_type=server_args.reasoning_parser,
model_type=get_serving().reasoning_parser,
stream_reasoning=False,
tokenizer=tokenizer,
)
@@ -399,7 +404,7 @@ def create_grammar_backend(
grammar_backend,
reasoning_parser,
tokenizer,
enable_strict_thinking=server_args.enable_strict_thinking,
enable_strict_thinking=get_serving().enable_strict_thinking,
)
return grammar_backend
@@ -34,6 +34,7 @@ from sglang.srt.managers.io_struct import GenerateReqInput, TokenizedGenerateReq
from sglang.srt.managers.multimodal_processor import get_mm_processor, import_processors
from sglang.srt.managers.schedule_batch import Modality, Req
from sglang.srt.multimodal.cache import media_preprocess_kwargs
from sglang.srt.runtime_context import get_disagg, get_exec, get_serving
from sglang.srt.server_args import ServerArgs
from sglang.srt.utils import ImageData
from sglang.srt.utils.common import safe_pickle_loads
@@ -1586,7 +1587,7 @@ class MMReceiverBase(ABC):
# context alive for the process instead of creating a temporary context
# whose destruction also closes its per-request socket.
self.scheduler_context = zmq.Context()
self.encoder_transfer_backend = server_args.encoder_transfer_backend
self.encoder_transfer_backend = get_disagg().encoder_transfer_backend
# When ``encode_urls`` is shared with an :class:`EncoderBootstrapServer`
# (tokenizer manager process), it grows / shrinks in place as encoders
# register or unregister; the receiver always sees the current set.
@@ -1642,8 +1643,8 @@ class MMReceiverBase(ABC):
self.embeddings_engine = init_mooncake_transfer_engine(
hostname=self.host,
ib_device=(
server_args.disaggregation_ib_device
or server_args.mooncake_ib_device
get_disagg().disaggregation_ib_device
or get_exec().moe.mooncake_ib_device
),
)
self.embeddings_buffer = dict()
@@ -1689,7 +1690,7 @@ class MMReceiverBase(ABC):
extra_kwargs["tokenizer_backend"] = server_args.tokenizer_backend
_processor = get_processor(
server_args.tokenizer_path,
get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
@@ -57,7 +57,7 @@ from sglang.srt.mem_cache.multimodal_cache import EmbeddingResult, MultiModalSta
from sglang.srt.model_executor.model_runner_components.load_model_utils import (
maybe_precompile_model_kernels_after_loading,
)
from sglang.srt.model_loader import get_model
from sglang.srt.model_loader import get_model as load_model
from sglang.srt.multimodal.cache import parse_content_hash, snapshot_media
from sglang.srt.multimodal.encoder_preprocessing import (
EncoderPreprocessOutput,
@@ -73,10 +73,15 @@ from sglang.srt.observability.trace import (
trace_set_thread_info,
)
from sglang.srt.runtime_context import (
configured_tp_size,
get_device,
get_disagg,
get_exec,
get_mm,
get_model,
get_observability,
get_parallel,
get_serving,
publish,
)
from sglang.srt.server_args import (
@@ -315,7 +320,7 @@ class MMEncoder:
logger.info(f"init MMEncoder {rank}/{server_args.tp_size}")
self.server_args = server_args
configure_media_url_security(
server_args.allowed_media_domains,
get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb,
)
self.rank = rank
@@ -329,7 +334,7 @@ class MMEncoder:
server_args,
)
self.load_config = LoadConfig(
load_format=server_args.load_format,
load_format=get_model().load_format,
download_dir=server_args.download_dir,
model_loader_extra_config=server_args.model_loader_extra_config,
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip,
@@ -340,7 +345,7 @@ class MMEncoder:
self.model_config.hf_config, "model_type", "unknown"
).lower()
self.device = server_args.device
self.device = get_device().device
self.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id
self.device_config = DeviceConfig(
@@ -354,7 +359,7 @@ class MMEncoder:
use_image_processor_gpu
and resolve_image_processor_backend(server_args) != "pil"
)
self._build_vision_config(server_args.mm_process_config)
self._build_vision_config(get_mm().mm_process_config)
self.model_audio_sr = self._resolve_audio_sr()
logger.info(f"Resolved model audio sample rate: {self.model_audio_sr} Hz")
@@ -368,7 +373,7 @@ class MMEncoder:
initialize_model_parallel(tensor_model_parallel_size=server_args.tp_size)
initialize_dp_attention(server_args, self.model_config)
self.model = get_model(
self.model = load_model(
model_config=self.model_config,
load_config=self.load_config,
device_config=self.device_config,
@@ -628,7 +633,7 @@ class MMEncoder:
)
try:
self.image_processor = AutoImageProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
get_serving().tokenizer_path or get_model().model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
**image_processor_kwargs,
@@ -639,7 +644,7 @@ class MMEncoder:
try:
self.video_processor = AutoVideoProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
get_serving().tokenizer_path or get_model().model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
@@ -650,7 +655,7 @@ class MMEncoder:
try:
# Note: AutoProcessor is used for audio processor
_audio_proc = AutoProcessor.from_pretrained(
server_args.tokenizer_path or server_args.model_path,
get_serving().tokenizer_path or get_model().model_path,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
)
@@ -2100,7 +2105,7 @@ class MMEncoder:
_zmq_xfer_start = time.perf_counter()
if (
self.server_args.encoder_transfer_backend == "zmq_to_scheduler"
get_disagg().encoder_transfer_backend == "zmq_to_scheduler"
and url is not None
):
lock = self.scheduler_send_locks.get(endpoint)
@@ -2146,7 +2151,7 @@ class MMEncoder:
if encoder_metrics_collector is not None:
encoder_metrics_collector.observe_transfer(
time.perf_counter() - _zmq_xfer_start,
backend=self.server_args.encoder_transfer_backend,
backend=get_disagg().encoder_transfer_backend,
)
return
@@ -3007,7 +3012,7 @@ async def _push_embedding_to_prefill(enc: MMEncoder, request: dict) -> None:
# No-op for mooncake (its /send is separate). embedding_port=None is
# rejected upfront, so ports is always a concrete list here.
req_id = request["req_id"]
backend = enc.server_args.encoder_transfer_backend
backend = get_disagg().encoder_transfer_backend
if backend == "zmq_to_tokenizer":
await enc.send(
@@ -3050,7 +3055,7 @@ async def _dp_worker_encode_and_send(
modality = Modality.from_str(request["modality"])
time_stats.modality = modality.name.lower()
time_stats.set_metrics_collector(encoder_metrics_collector)
backend = enc.server_args.encoder_transfer_backend
backend = get_disagg().encoder_transfer_backend
# URL state lives in main process module globals; workers don't see it.
if backend == "zmq_to_scheduler" and request.get("embedding_port") is None:
@@ -3666,14 +3671,14 @@ async def run_dp_worker(
)
global encoder_metrics_collector
if server_args.enable_metrics:
if get_observability().enable_metrics:
set_prometheus_multiproc_dir()
labels = {
"model_name": server_args.served_model_name,
"model_name": get_serving().served_model_name,
"dp_rank": str(dp_rank),
}
if server_args.extra_metric_labels:
labels.update(server_args.extra_metric_labels)
if get_observability().extra_metric_labels:
labels.update(get_observability().extra_metric_labels)
encoder_metrics_collector = EncoderMetricsCollector(labels)
enc.dp_rank = dp_rank
@@ -3964,14 +3969,14 @@ def launch_server(server_args: ServerArgs):
global encoder, encoder_metrics_collector
# Set up prometheus metrics.
if server_args.enable_metrics:
if get_observability().enable_metrics:
set_prometheus_multiproc_dir()
labels = {
"model_name": server_args.served_model_name,
"model_name": get_serving().served_model_name,
"dp_rank": "0",
}
if server_args.extra_metric_labels:
labels.update(server_args.extra_metric_labels)
if get_observability().extra_metric_labels:
labels.update(get_observability().extra_metric_labels)
encoder_metrics_collector = EncoderMetricsCollector(labels)
add_prometheus_middleware(app)
@@ -3979,21 +3984,21 @@ def launch_server(server_args: ServerArgs):
zmq_ctx = zmq.Context(10)
ipc_path_prefix = random_uuid()
port_args = PortArgs.init_new(server_args)
if server_args.dist_init_addr:
na = NetworkAddress.parse(server_args.dist_init_addr)
if get_parallel().dist_init_addr:
na = NetworkAddress.parse(get_parallel().dist_init_addr)
dist_init_method = na.to_tcp()
else:
dist_init_method = NetworkAddress(
server_args.host or "127.0.0.1", port_args.nccl_port
get_serving().host or "127.0.0.1", port_args.nccl_port
).to_tcp()
if server_args.enable_trace:
if get_observability().enable_trace:
process_tracing_init(
server_args.otlp_traces_endpoint,
get_observability().otlp_traces_endpoint,
"sglang",
trace_modules=server_args.trace_modules,
trace_modules=get_observability().trace_modules,
)
trace_set_thread_info("Encoder")
for rank in range(1, server_args.tp_size):
for rank in range(1, configured_tp_size()):
schedule_path = f"ipc:///tmp/{ipc_path_prefix}_schedule_{rank}"
send_sockets.append(
get_zmq_socket(zmq_ctx, zmq.PUSH, schedule_path, bind=False)
@@ -4006,13 +4011,13 @@ def launch_server(server_args: ServerArgs):
encoder = MMEncoder(server_args, dist_init_method=dist_init_method)
# Register this encoder's URL with prefill server(s) if configured.
if server_args.encoder_register_urls:
if get_disagg().encoder_register_urls:
import atexit
_register_encoder_url_with_bootstrap(server_args)
atexit.register(_unregister_encoder_url_from_bootstrap, server_args)
uvicorn.run(app, host=server_args.host, port=server_args.port)
uvicorn.run(app, host=get_serving().host, port=get_serving().port)
def _launch_server_dp(server_args: ServerArgs):
@@ -4083,7 +4088,7 @@ def _launch_server_dp(server_args: ServerArgs):
proc.start()
worker_processes.append(proc)
labels = {"model_name": server_args.served_model_name}
labels = {"model_name": get_serving().served_model_name}
if server_args.extra_metric_labels:
labels.update(server_args.extra_metric_labels)
dp_dispatcher = DPDispatcher(
@@ -4188,7 +4193,7 @@ async def handle_encode_request(request: dict):
# when multiple decoder TP ranks POST /encode
# with the same req_id, only the first triggers the VIT forward;
# subsequent callers wait and return the same metadata.
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
async with encoder._inflight_encode_lock:
if req_id in encoder._inflight_encode_events:
event = encoder._inflight_encode_events[req_id]
@@ -4274,7 +4279,7 @@ async def handle_encode_request(request: dict):
time_stats.set_mm_encode_end_time()
if error_msg:
if encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
if get_disagg().encoder_transfer_backend == "zmq_to_scheduler":
if request["embedding_port"] is None:
start_background_send(req_id)
else:
@@ -4285,7 +4290,7 @@ async def handle_encode_request(request: dict):
embedding_port=port,
)
# Signal waiters on failure for mooncake
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
encoder._inflight_encode_meta.pop(req_id, None)
evt = encoder._inflight_encode_events.pop(req_id, None)
if evt:
@@ -4299,7 +4304,7 @@ async def handle_encode_request(request: dict):
status_code=error_code,
content={"status": "error", "message": error_msg, "req_id": req_id},
)
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
# Store metadata for duplicate callers and signal them
encoder._inflight_encode_meta[req_id] = (
nbytes,
@@ -4323,7 +4328,7 @@ async def handle_encode_request(request: dict):
modality=modality_str, status="success"
)
return ORJSONResponse(content=request)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_scheduler":
elif get_disagg().encoder_transfer_backend == "zmq_to_scheduler":
logger.info(f"{request['embedding_port'] = }")
if request["embedding_port"] is None:
await encoder.send_with_url(
@@ -4347,7 +4352,7 @@ async def handle_encode_request(request: dict):
modality=modality_str, status="success"
)
return ORJSONResponse(content=None)
elif encoder.server_args.encoder_transfer_backend == "zmq_to_tokenizer":
elif get_disagg().encoder_transfer_backend == "zmq_to_tokenizer":
await encoder.send(
req_id=request["req_id"],
prefill_host=request["prefill_host"],
@@ -4369,7 +4374,7 @@ async def handle_encode_request(request: dict):
logger.error(f"Unexpected error in encoder logic for {req_id}: {error_msg}")
rid_to_err_msg[req_id] = error_msg
# Ensure inflight waiters are unblocked on unexpected errors
if encoder.server_args.encoder_transfer_backend == "mooncake":
if get_disagg().encoder_transfer_backend == "mooncake":
encoder._inflight_encode_meta.pop(req_id, None)
evt = encoder._inflight_encode_events.pop(req_id, None)
if evt:
+23 -16
View File
@@ -99,7 +99,14 @@ from sglang.srt.observability.trace import process_tracing_init, trace_set_threa
from sglang.srt.parser.template_detection import resolve_auto_parsers
from sglang.srt.parser.template_manager import TemplateManager
from sglang.srt.plugins import load_plugins
from sglang.srt.runtime_context import get_parallel, publish
from sglang.srt.runtime_context import (
configured_pp_size,
get_exec,
get_model,
get_parallel,
get_serving,
publish,
)
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import (
MultiprocessingSerializer,
@@ -158,9 +165,9 @@ def init_tokenizer_manager(
template_manager = TemplateManager()
template_manager.initialize_templates(
tokenizer_manager=tokenizer_manager,
model_path=server_args.model_path,
chat_template=server_args.chat_template,
completion_template=server_args.completion_template,
model_path=get_model().model_path,
chat_template=get_serving().chat_template,
completion_template=get_serving().completion_template,
)
# Resolve any remaining auto parsers using template manager's detection results
@@ -681,7 +688,7 @@ class Engine(EngineScoreMixin, EngineBase):
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
_calculate_rank_ranges(
server_args.nnodes,
server_args.pp_size,
configured_pp_size(),
tp_size,
server_args.node_rank,
)
@@ -702,7 +709,7 @@ class Engine(EngineScoreMixin, EngineBase):
daemon_procs = []
logger.info(
f"Launching {num_daemons} weight cache daemon(s) on node "
f"{server_args.node_rank} for model={server_args.model_path}, "
f"{server_args.node_rank} for model={get_model().model_path}, "
f"pp_ranks={pp_rank_range.start}..{pp_rank_range.stop - 1}, "
f"tp_ranks={tp_rank_range.start}..{tp_rank_range.stop - 1}, "
f"dist_init_method={dist_init_method}"
@@ -737,7 +744,7 @@ class Engine(EngineScoreMixin, EngineBase):
"-m",
"sglang.srt.weight_cache.daemon",
"--model-path",
server_args.model_path,
get_model().model_path,
"--gpu-id",
str(gpu_id),
"--tp-size",
@@ -745,7 +752,7 @@ class Engine(EngineScoreMixin, EngineBase):
"--tp-rank",
str(tp_rank),
"--pp-size",
str(server_args.pp_size),
str(configured_pp_size()),
"--pp-rank",
str(pp_rank),
"--dp-size",
@@ -753,14 +760,14 @@ class Engine(EngineScoreMixin, EngineBase):
"--ep-size",
str(get_parallel().ep_size),
"--load-format",
server_args.load_format,
get_model().load_format,
"--dtype",
server_args.dtype,
get_model().dtype,
"--dist-init-method",
dist_init_method,
]
if server_args.quantization:
cmd += ["--quantization", server_args.quantization]
if get_model().quantization:
cmd += ["--quantization", get_model().quantization]
if (
server_args.model_loader_extra_config
and server_args.model_loader_extra_config != "{}"
@@ -863,7 +870,7 @@ class Engine(EngineScoreMixin, EngineBase):
"""
scheduler_procs = []
use_dp_controller = (
get_parallel().dp_size > 1 or server_args.ep_join_mode == "scale"
get_parallel().dp_size > 1 or get_exec().moe.ep_join_mode == "scale"
)
if not use_dp_controller:
@@ -876,7 +883,7 @@ class Engine(EngineScoreMixin, EngineBase):
pp_rank_range, tp_rank_range, pp_size_per_node, tp_size_per_node = (
_calculate_rank_ranges(
server_args.nnodes,
server_args.pp_size,
configured_pp_size(),
server_args.tp_size,
server_args.node_rank,
)
@@ -983,7 +990,7 @@ class Engine(EngineScoreMixin, EngineBase):
processes: List[mp.Process] = []
names: List[str] = []
if server_args.detokenizer_worker_num <= 1:
if get_serving().detokenizer_worker_num <= 1:
proc = mp.Process(
target=run_detokenizer_process_func,
args=(server_args, port_args),
@@ -996,7 +1003,7 @@ class Engine(EngineScoreMixin, EngineBase):
router_ipc_name = port_args.detokenizer_ipc_name
worker_ipc_names: List[str] = []
try:
for i in range(server_args.detokenizer_worker_num):
for i in range(get_serving().detokenizer_worker_num):
worker_ipc = f"ipc://{tempfile.NamedTemporaryFile(delete=False).name}"
port_args.detokenizer_ipc_name = worker_ipc
proc = mp.Process(
+26 -23
View File
@@ -247,7 +247,7 @@ async def init_multi_tokenizer() -> ServerArgs:
template_manager = TemplateManager()
template_manager.initialize_templates(
tokenizer_manager=tokenizer_manager,
model_path=server_args.model_path,
model_path=get_model().model_path,
chat_template=server_args.chat_template,
completion_template=server_args.completion_template,
)
@@ -293,9 +293,9 @@ async def lifespan(fast_api_app: FastAPI):
"sglang",
trace_modules=server_args.trace_modules,
)
if server_args.disaggregation_mode == "prefill":
if get_disagg().disaggregation_mode == "prefill":
thread_label = "Prefill" + thread_label
elif server_args.disaggregation_mode == "decode":
elif get_disagg().disaggregation_mode == "decode":
thread_label = "Decode" + thread_label
trace_set_thread_info(thread_label)
@@ -380,7 +380,7 @@ async def lifespan(fast_api_app: FastAPI):
# Execute custom warmups
if server_args.warmups is not None:
await execute_warmups(
server_args.disaggregation_mode,
get_disagg().disaggregation_mode,
server_args.warmups.split(","),
_global_state.tokenizer_manager,
)
@@ -393,7 +393,7 @@ async def lifespan(fast_api_app: FastAPI):
try:
if (
getattr(fast_api_app, "is_single_tokenizer_mode", False)
and server_args.grpc_port is not None
and get_serving().grpc_port is not None
and not (server_args.smg_grpc_mode or server_args.grpc_mode)
):
grpc_handle = _start_native_grpc_server_for_runtime(
@@ -401,6 +401,7 @@ async def lifespan(fast_api_app: FastAPI):
tokenizer_manager=_global_state.tokenizer_manager,
template_manager=_global_state.template_manager,
scheduler_info=_global_state.scheduler_info,
grpc_port=get_serving().grpc_port,
)
if server_args.sidecar is not None:
from sglang.srt.entrypoints.sidecar import start_sidecar
@@ -480,7 +481,13 @@ v1_loads_router.route_class = ORJSONRoute
app.include_router(v1_loads_router)
from sglang.srt.entrypoints.elastic_ep import router as elastic_ep_router
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import (
get_disagg,
get_exec,
get_model,
get_parallel,
get_serving,
)
elastic_ep_router.route_class = ORJSONRoute
app.include_router(elastic_ep_router)
@@ -679,10 +686,7 @@ async def health_generate(request: Request) -> Response:
sampling_params=sampling_params,
log_metrics=False,
)
if (
_global_state.tokenizer_manager.server_args.disaggregation_mode
!= DisaggregationMode.NULL.value
):
if get_disagg().disaggregation_mode != DisaggregationMode.NULL.value:
gri.bootstrap_host = FAKE_BOOTSTRAP_HOST
gri.bootstrap_room = 0
else:
@@ -2224,7 +2228,7 @@ def _execute_server_warmup(server_args: ServerArgs):
json_data["input_ids"] = json_data["input_ids"][0]
elif (
is_vlm
and server_args.disaggregation_mode == "null"
and get_disagg().disaggregation_mode == "null"
and model_info["is_generation"]
):
served_model_name = ""
@@ -2234,9 +2238,9 @@ def _execute_server_warmup(server_args: ServerArgs):
# _global_state.tokenizer_manager is not initialized in the rust server,
# so we need to get the model name from the model_info
served_model_name = model_info.get(
"model_path", server_args.served_model_name
"model_path", get_serving().served_model_name
)
served_model_name = served_model_name or server_args.model_path
served_model_name = served_model_name or get_model().model_path
# TODO: ChatCompletionRequest does not have bootstrap info required by disaggregation mode, disable image-warmup for now
# Only use chat completions format for generation models, not embedding models
json_data = {
@@ -2280,7 +2284,7 @@ def _execute_server_warmup(server_args: ServerArgs):
# Send a warmup request
warmup_timeout = envs.SGLANG_WARMUP_TIMEOUT.get()
try:
if server_args.disaggregation_mode == "null":
if get_disagg().disaggregation_mode == "null":
res = requests.post(
url + request_name,
json=json_data,
@@ -2314,7 +2318,7 @@ def _execute_server_warmup(server_args: ServerArgs):
else:
logger.info(
"Disaggregation warmup failed (mode=%s), status codes: %s",
server_args.disaggregation_mode,
get_disagg().disaggregation_mode,
failed_status_codes,
)
# In rust-server mode there is no TokenizerManager (readiness is
@@ -2368,10 +2372,10 @@ def _wait_and_warmup(
logger.debug(
"[Elastic EP] Skipping server warmup for elastic joiner "
"(ep_join_mode=%s)",
server_args.ep_join_mode,
get_exec().moe.ep_join_mode,
)
if not server_args.skip_server_warmup and not skip_elastic_joiner_warmup:
if not get_serving().skip_server_warmup and not skip_elastic_joiner_warmup:
if not execute_warmup_func(server_args):
return
else:
@@ -2383,7 +2387,7 @@ def _wait_and_warmup(
logger.info("The server is fired up and ready to roll!")
if server_args.delete_ckpt_after_loading:
delete_directory(server_args.model_path)
delete_directory(get_model().model_path)
if server_args.debug_tensor_dump_input_file:
kill_process_tree(os.getpid())
@@ -2711,6 +2715,7 @@ def _start_native_grpc_server_for_runtime(
tokenizer_manager,
template_manager,
scheduler_info,
grpc_port,
):
from sglang.srt.entrypoints.grpc_bridge import RuntimeHandle
from sglang.srt.rust_extensions import load_rust_extension
@@ -2726,13 +2731,11 @@ def _start_native_grpc_server_for_runtime(
grpc_handle = grpc_native.start_server(
host=server_args.host,
port=server_args.grpc_port,
port=grpc_port,
runtime_handle=runtime_handle,
worker_threads=server_args.grpc_worker_threads,
)
logger.info(
f"Native gRPC server started on {server_args.host}:{server_args.grpc_port}"
)
logger.info(f"Native gRPC server started on {server_args.host}:{grpc_port}")
return grpc_handle
@@ -2790,7 +2793,7 @@ def launch_server(
# and /get_model_info endpoints are static (200 as soon as the server
# binds, before any forward pass), so without this the first real request
# pays the cold-start cost (observed as a >60s first generation).
if not server_args.skip_server_warmup:
if not get_serving().skip_server_warmup:
_execute_server_warmup(server_args)
logger.info("The server is fired up and ready to roll!")
if launch_callback is not None:
+2
View File
@@ -37,6 +37,8 @@ def _loopback_host(host: str) -> str:
def build_sidecar_endpoint(server_args) -> str:
"""Both halves of the endpoint come from the argument: this is a helper
over a config object, callable before anything is published."""
return NetworkAddress(
_loopback_host(server_args.host), server_args.grpc_port
).to_url()
@@ -1,21 +1,17 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Optional
from typing import Optional
from sglang.srt.kv_canary.token_oracle.oracle import HashOracle
from sglang.srt.kv_canary.token_oracle.oracle_manager import TokenOracleManager
from sglang.srt.kv_canary.token_oracle.sampler import install_oracle_sampler
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
from sglang.srt.runtime_context import get_exec
def install_token_oracle_from_env(
*, server_args: ServerArgs, vocab_size: int
) -> Optional[TokenOracleManager]:
def install_token_oracle_from_env(*, vocab_size: int) -> Optional[TokenOracleManager]:
# Must be called before create_sampler() so the factory is present when the
# Sampler is first constructed.
if server_args.sampling_backend != "token_oracle":
if get_exec().kernel.sampling_backend != "token_oracle":
return None
oracle = HashOracle(vocab_size=vocab_size)
@@ -39,7 +39,7 @@ from sglang.srt.managers.io_struct import (
)
from sglang.srt.managers.multi_tokenizer_mixin import MultiHttpWorkerDetokenizerMixin
from sglang.srt.observability.cpu_monitor import start_cpu_monitor_thread
from sglang.srt.runtime_context import publish
from sglang.srt.runtime_context import get_device, get_serving, publish
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils import configure_logger, freeze_gc, kill_itself_when_parent_died
from sglang.srt.utils.hf_transformers_utils import get_tokenizer
@@ -128,7 +128,7 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
self.vocab_size = None
else:
self.tokenizer = get_tokenizer(
server_args.tokenizer_path,
get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
@@ -142,11 +142,11 @@ class DetokenizerManager(MultiHttpWorkerDetokenizerMixin):
def init_running_status(self, server_args: ServerArgs):
self.decode_status = LimitedCapacityDict(capacity=DETOKENIZER_MAX_STATES)
self.disable_tokenizer_batch_decode = server_args.disable_tokenizer_batch_decode
self.is_tool_call_parser_gpt_oss = server_args.tool_call_parser == "gpt-oss"
self.is_tool_call_parser_gpt_oss = get_serving().tool_call_parser == "gpt-oss"
self.soft_watchdog = Watchdog.create(
debug_name="DetokenizerManager",
watchdog_timeout=server_args.soft_watchdog_timeout,
watchdog_timeout=get_device().soft_watchdog_timeout,
soft=True,
test_stuck_time=envs.SGLANG_TEST_STUCK_DETOKENIZER.get(),
)
+1 -3
View File
@@ -1252,9 +1252,7 @@ class Scheduler(
and not get_schedule().disable_priority_preemption
)
self.new_token_ratio_tracker = NewTokenRatioTracker.from_server_args(
self.server_args
)
self.new_token_ratio_tracker = NewTokenRatioTracker.from_config()
def init_soft_watchdog(self, server_args: ServerArgs):
if (x := server_args.soft_watchdog_timeout) is not None:
@@ -566,7 +566,7 @@ class SchedulerBatchResultProcessor:
return get_required_capture_hidden_mode(
max(
batch.return_hidden_states_mode,
get_server_return_hidden_states_mode(server_args),
get_server_return_hidden_states_mode(),
),
batch.spec_info,
)
@@ -4,7 +4,7 @@ from dataclasses import dataclass
from typing import TYPE_CHECKING, Sequence
from sglang.srt.environ import envs
from sglang.srt.server_args import ServerArgs
from sglang.srt.runtime_context import get_schedule
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import Req
@@ -18,10 +18,10 @@ class NewTokenRatioTracker:
current: float
@classmethod
def from_server_args(cls, server_args: ServerArgs) -> NewTokenRatioTracker:
def from_config(cls) -> NewTokenRatioTracker:
init = min(
envs.SGLANG_INIT_NEW_TOKEN_RATIO.get()
* server_args.schedule_conservativeness,
* get_schedule().schedule_conservativeness,
1.0,
)
min_ratio = min(
+46 -38
View File
@@ -121,7 +121,18 @@ from sglang.srt.observability.request_metrics_exporter import (
RequestMetricsExporterManager,
)
from sglang.srt.observability.trace import SpanAttributes, extract_trace_headers
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import (
get_device,
get_disagg,
get_exec,
get_lora,
get_memory,
get_mm,
get_model,
get_parallel,
get_serving,
get_spec,
)
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.srt.server_args import (
PortArgs,
@@ -160,12 +171,13 @@ _REQUEST_STATE_WAIT_TIMEOUT = envs.SGLANG_REQUEST_STATE_WAIT_TIMEOUT.get()
logger = logging.getLogger(__name__)
def _reject_missing_dispatched_encoder_embedding(server_args, request_obj, mm_inputs):
def _reject_missing_dispatched_encoder_embedding(request_obj, mm_inputs):
"""Do not silently turn a failed EPD request into local vision work."""
disagg = get_disagg()
if (
mm_inputs is None
and server_args.language_only
and server_args.encoder_transfer_backend == "zmq_to_tokenizer"
and disagg.language_only
and disagg.encoder_transfer_backend == "zmq_to_tokenizer"
and request_obj.need_wait_for_mm_inputs
):
raise fastapi.HTTPException(
@@ -405,11 +417,11 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.elastic_last_error = None
self.enable_metrics = server_args.enable_metrics
self.incremental_streaming_output = server_args.incremental_streaming_output
self.enable_lora = server_args.enable_lora
self.enable_lora = get_lora().enable_lora
self.enable_trace = server_args.enable_trace
self.allow_auto_truncate = server_args.allow_auto_truncate
self.skip_tokenizer_init = server_args.skip_tokenizer_init
self.preferred_sampling_params = server_args.preferred_sampling_params
self.preferred_sampling_params = get_serving().preferred_sampling_params
self.crash_dump_folder = server_args.crash_dump_folder
# Init model config
@@ -501,7 +513,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.tokenizer = None
else:
self.tokenizer = get_tokenizer(
server_args.tokenizer_path,
get_serving().tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
@@ -522,7 +534,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
self.async_dynamic_batch_tokenizer = None
def _validate_cuda_vmm_feature_transport_support(self) -> None:
if self.server_args.mm_feature_transport != "cuda_vmm":
if get_mm().mm_feature_transport != "cuda_vmm":
return
from sglang.srt.model_loader.utils import get_model_architecture
@@ -633,7 +645,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# The registry dynamically updates as adapters are loaded / unloaded during runtime. It
# serves as the source of truth for available adapters and maps user-friendly LoRA names
# to internally used unique LoRA IDs.
self.lora_registry = LoRARegistry(self.server_args.lora_paths)
self.lora_registry = LoRARegistry(get_lora().lora_paths)
# Lock to serialize LoRA update operations.
# Please note that, unlike `model_update_lock`, this does not block inference, allowing
# LoRA updates and inference to overlap.
@@ -642,15 +654,13 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# point to their latest LoRARef objects, so that they can be
# dynamically loaded if needed for inference
self.lora_ref_cache: Dict[str, LoRARef] = {}
if self.server_args.lora_paths is not None:
for lora_ref in self.server_args.lora_paths:
if get_lora().lora_paths is not None:
for lora_ref in get_lora().lora_paths:
self.lora_ref_cache[lora_ref.lora_name] = lora_ref
def init_disaggregation(self, *, start_pd_bootstrap_service: bool = True):
# PD Disaggregation
self.disaggregation_mode = DisaggregationMode(
self.server_args.disaggregation_mode
)
self.disaggregation_mode = DisaggregationMode(get_disagg().disaggregation_mode)
# Keep a reference so the bootstrap server is not garbage-collected.
self.bootstrap_server = (
start_disagg_service(self.server_args)
@@ -688,7 +698,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# Metrics
if self.enable_metrics:
engine_type = DisaggregationMode.to_engine_type(
self.server_args.disaggregation_mode
get_disagg().disaggregation_mode
)
labels = {
@@ -721,7 +731,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
configure_gc_warning(self.server_args.gc_warning_threshold_secs)
self.soft_watchdog = Watchdog.create(
debug_name="TokenizerManager",
watchdog_timeout=self.server_args.soft_watchdog_timeout,
watchdog_timeout=get_device().soft_watchdog_timeout,
soft=True,
test_stuck_time=envs.SGLANG_TEST_STUCK_TOKENIZER.get(),
)
@@ -997,7 +1007,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
isinstance(obj, EmbeddingReqInput) and obj.is_cross_encoder_request
)
if obj.input_embeds is not None:
if not self.server_args.disable_radix_cache:
if not get_memory().disable_radix_cache:
raise ValueError(
"input_embeds is provided while disable_radix_cache is False. "
"Please add `--disable-radix-cache` when you launch the server "
@@ -1062,7 +1072,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
if (
not self.server_args.language_only
or self.server_args.encoder_transfer_backend == "zmq_to_tokenizer"
or get_disagg().encoder_transfer_backend == "zmq_to_tokenizer"
):
if self.server_args.language_only:
mm_inputs = await self.mm_receiver.recv_mm_data(
@@ -1071,11 +1081,9 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
prompt=mm_processor_input,
need_wait_for_mm_inputs=obj.need_wait_for_mm_inputs,
)
_reject_missing_dispatched_encoder_embedding(
self.server_args, obj, mm_inputs
)
_reject_missing_dispatched_encoder_embedding(obj, mm_inputs)
if mm_inputs is None:
if self.server_args.language_only:
if get_disagg().language_only:
logger.warning(
"Encoder embedding not available, "
"falling back to local mm processing"
@@ -1089,7 +1097,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
)
elif (
self.server_args.language_only
and self.server_args.encoder_transfer_backend
and get_disagg().encoder_transfer_backend
in ["zmq_to_scheduler", "mooncake"]
and not obj.need_wait_for_mm_inputs
):
@@ -1256,12 +1264,12 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
requested_hidden_mode = get_request_return_hidden_states_mode(
obj.return_hidden_states
)
server_hidden_mode = get_server_return_hidden_states_mode(self.server_args)
server_hidden_mode = get_server_return_hidden_states_mode()
if requested_hidden_mode > server_hidden_mode:
if server_hidden_mode.need_capture():
raise ValueError(
"The requested return_hidden_states mode exceeds the "
f"server maximum `{self.server_args.return_hidden_states_mode}`. "
f"server maximum `{get_exec().features.return_hidden_states_mode}`. "
"Please launch with `--return-hidden-states-mode full` "
"to allow return_hidden_states=True."
)
@@ -1283,10 +1291,10 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
def _validate_mm_limits(
self, obj: Union[GenerateReqInput, EmbeddingReqInput]
) -> None:
if not self.server_args.limit_mm_data_per_request:
if not get_mm().limit_mm_data_per_request:
return
for modality, limit in self.server_args.limit_mm_data_per_request.items():
for modality, limit in get_mm().limit_mm_data_per_request.items():
data = getattr(obj, f"{modality}_data", None)
if data:
count = len(data) if isinstance(data, list) else 1
@@ -1399,7 +1407,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
bootstrap_room = obj.bootstrap_room
if (
bootstrap_room is None
and self.server_args.disaggregation_transfer_backend == "fake"
and get_disagg().disaggregation_transfer_backend == "fake"
):
bootstrap_room = self.fake_bootstrap_room_counter
self.fake_bootstrap_room_counter += 1
@@ -1581,7 +1589,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
- Batch tokenization does not support DP attention yet, and it will make everything goes to the first rank currently
"""
return batch_size > 0 and (
self.server_args.enable_tokenizer_batch_encode
get_serving().enable_tokenizer_batch_encode
or (
(not get_parallel().enable_dp_attention)
and (not self._batch_has_text(batch_size, requests))
@@ -2464,7 +2472,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
state.time_stats.set_finished_time()
meta_info["e2e_latency"] = state.time_stats.get_e2e_latency()
if self.server_args.speculative_algorithm:
if get_spec().speculative_algorithm:
self._calculate_spec_decoding_metrics(meta_info, recv_obj, i)
if self.enable_metrics:
scheduler_time_stats = (
@@ -2787,7 +2795,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
):
# Total number of proposed draft tokens per request.
num_proposed_drafts = recv_obj.spec_verify_ct[i] * (
self.server_args.speculative_num_draft_tokens - 1
get_spec().speculative_num_draft_tokens - 1
)
num_correct_drafts = recv_obj.spec_num_correct_drafts[i]
@@ -3484,7 +3492,7 @@ class TokenizerManager(TokenizerControlMixin, TokenizerManagerScoreMixin):
# This flag will be used in _tokenize_one_request to determine processing path
if should_dispatch:
obj.need_wait_for_mm_inputs = True
if self.server_args.encoder_transfer_backend in [
if get_disagg().encoder_transfer_backend in [
"zmq_to_scheduler",
"mooncake",
]:
@@ -3603,13 +3611,13 @@ async def print_exception_wrapper(func):
def get_processor_wrapper(server_args):
return get_processor(
server_args.tokenizer_path,
tokenizer_mode=server_args.tokenizer_mode,
trust_remote_code=server_args.trust_remote_code,
revision=server_args.revision,
get_serving().tokenizer_path,
tokenizer_mode=get_serving().tokenizer_mode,
trust_remote_code=get_model().trust_remote_code,
revision=get_model().revision,
image_processor_backend=resolve_image_processor_backend(server_args),
tokenizer_backend=server_args.tokenizer_backend,
model_name=server_args.model_path,
tokenizer_backend=get_serving().tokenizer_backend,
model_name=get_model().model_path,
)
@@ -567,7 +567,7 @@ class CPUGraphRunner:
self.return_hidden_states_mode = (
CaptureHiddenMode.NULL
if model_runner.is_draft_worker
else get_server_return_hidden_states_mode(model_runner.server_args)
else get_server_return_hidden_states_mode()
)
self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture()
# bs -> compiled fn (text-only / skip_cross_attention=True)
@@ -32,7 +32,7 @@ import warnings
from dataclasses import dataclass
from enum import IntEnum, auto
from functools import total_ordering
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional, Set, Tuple, Union
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Set, Tuple, Union
import torch
@@ -231,11 +231,12 @@ def register_attn_tp_sequence_sharded_predicate(
_attn_tp_sequence_sharded_predicate = predicate
def get_server_return_hidden_states_mode(server_args: Any) -> CaptureHiddenMode:
mode = getattr(server_args, "return_hidden_states_mode", None)
def get_server_return_hidden_states_mode() -> CaptureHiddenMode:
features = get_exec().features
mode = features.return_hidden_states_mode
if mode == "last":
return CaptureHiddenMode.LAST
if mode == "full" or getattr(server_args, "enable_return_hidden_states", False):
if mode == "full" or features.enable_return_hidden_states:
return CaptureHiddenMode.FULL
return CaptureHiddenMode.NULL
@@ -719,7 +720,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
if model_runner.is_draft_worker
else max(
batch.return_hidden_states_mode,
get_server_return_hidden_states_mode(model_runner.server_args),
get_server_return_hidden_states_mode(),
)
)
capture_hidden_mode = get_required_capture_hidden_mode(
@@ -726,7 +726,6 @@ class ModelRunner:
self._token_oracle_manager = None
return
self._token_oracle_manager = install_token_oracle_from_env(
server_args=self.server_args,
vocab_size=self.model_config.vocab_size,
)
@@ -261,8 +261,7 @@ def capture_prefill_graph(
if (
model_runner.spec_algorithm.is_eagle()
and not model_runner.is_draft_worker
and get_server_return_hidden_states_mode(model_runner.server_args)
< CaptureHiddenMode.FULL
and get_server_return_hidden_states_mode() < CaptureHiddenMode.FULL
and not check_cuda_graph_backend(Phase.PREFILL, Backend.BREAKABLE)
):
logger.info(
@@ -219,7 +219,7 @@ class BaseRunner(ABC):
self.return_hidden_states_mode = (
CaptureHiddenMode.NULL
if model_runner.is_draft_worker
else get_server_return_hidden_states_mode(model_runner.server_args)
else get_server_return_hidden_states_mode()
)
self.enable_return_hidden_states = self.return_hidden_states_mode.need_capture()
self.attn_tp_size = get_parallel().attn_tp_size
@@ -403,7 +403,7 @@ class BaseRunner(ABC):
capture_hidden_mode = (
CaptureHiddenMode.NULL
if mr.is_draft_worker
else get_server_return_hidden_states_mode(mr.server_args)
else get_server_return_hidden_states_mode()
)
num_tokens_per_req = 1
# A PD prefill target worker's pool has no SpeculativeState, so a
+7 -14
View File
@@ -67,7 +67,6 @@ from sglang.srt.runtime_context import (
get_memory,
get_model,
get_parallel,
get_server_args,
get_stream,
)
from sglang.srt.utils import (
@@ -159,12 +158,13 @@ for backend in CONCAT_ROPE_BACKENDS:
AttentionBackendRegistry.register(backend, _handle_concat_rope_backend)
def get_attn_forward_method(server_args, forward_batch) -> AttnForwardMethod:
def get_attn_forward_method(forward_batch) -> AttnForwardMethod:
prefill_backend, decode_backend = attention_backends()
is_decode = forward_batch.forward_mode.is_decode_or_idle()
if is_decode:
backend = server_args.decode_attention_backend or server_args.attention_backend
backend = decode_backend
else:
backend = server_args.prefill_attention_backend or server_args.attention_backend
backend = prefill_backend
if (
forward_batch.forward_mode.is_extend_without_speculative()
and backend == "fa3"
@@ -456,7 +456,6 @@ class SarvamMoEMLAAttention(nn.Module):
self.max_position_embeddings = max_position_embeddings
self.kv_cache_dtype = get_model().kv_cache_dtype
self._server_args = None
self.current_attention_backend = None
if self.q_lora_rank is None:
@@ -761,11 +760,9 @@ class SarvamMoEMLAAttention(nn.Module):
q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1)
if self._server_args is None:
self._server_args = get_server_args()
self._set_current_attention_backend(forward_batch)
forward_method = get_attn_forward_method(self._server_args, forward_batch)
forward_method = get_attn_forward_method(forward_batch)
if forward_method == AttnForwardMethod.MHA_PREFILL:
return self._run_mha_prefill(
@@ -875,10 +872,8 @@ class SarvamMoEMLAAttention(nn.Module):
q_nope, q_pe = q.split([self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1)
k_pe = latent_cache[..., self.kv_lora_rank :].unsqueeze(1)
if self._server_args is None:
self._server_args = get_server_args()
self._set_current_attention_backend(forward_batch)
forward_method = get_attn_forward_method(self._server_args, forward_batch)
forward_method = get_attn_forward_method(forward_batch)
if forward_method == AttnForwardMethod.MHA_PREFILL:
output = self._run_mha_prefill(
@@ -933,11 +928,9 @@ class SarvamMoEMLAAttention(nn.Module):
q_nope_out, k_nope, q_pe, k_pe, forward_batch, zero_allocator = inner_state
if self._server_args is None:
self._server_args = get_server_args()
self._set_current_attention_backend(forward_batch)
forward_method = get_attn_forward_method(self._server_args, forward_batch)
forward_method = get_attn_forward_method(forward_batch)
if forward_method == AttnForwardMethod.MLA_SEPARATE_ROPE:
attn_output = self.attn_mqa(
+14 -9
View File
@@ -9,7 +9,7 @@ import struct
from dataclasses import dataclass
from enum import Enum
from pathlib import Path
from typing import TYPE_CHECKING, Any, Mapping, Optional, Protocol, runtime_checkable
from typing import Any, Mapping, Optional, Protocol, runtime_checkable
from urllib.parse import unquote, urlparse
import numpy as np
@@ -17,8 +17,7 @@ import torch
import transformers
from PIL import Image
if TYPE_CHECKING:
from sglang.srt.server_args import ServerArgs
from sglang.srt.runtime_context import get_mm, get_model
CONTENT_HASH_PREFIX = "sha256:"
_SHA256_HEX_LENGTH = 64
@@ -379,11 +378,17 @@ def resolve_multimodal_item_hash(
def build_processor_fingerprint(
processor: Any,
hf_config: Any,
server_args: ServerArgs,
*,
extra: Optional[Mapping[str, Any]] = None,
) -> str:
"""Fingerprint preprocessing choices that can change processor output."""
"""Fingerprint preprocessing choices that can change processor output.
Every config value comes from the published bags, which is the only source
that answers the *effective* preprocessing config. Taking any of them from
a handed ``ServerArgs`` would let two callers with the same effective
config disagree on the digest -- and an omitted one silently fingerprint
the empty config, which is how incompatible artifacts get reused.
"""
processor_payload = (
processor.preprocess_fingerprint_payload()
if isinstance(processor, PreprocessFingerprintProvider)
@@ -395,10 +400,10 @@ def build_processor_fingerprint(
"processor_class": f"{type(processor).__module__}.{type(processor).__qualname__}",
"model_type": hf_payload.get("model_type"),
"architectures": hf_payload.get("architectures"),
"model_revision": server_args.revision,
"processor_revision": server_args.revision,
"disable_fast_image_processor": server_args.disable_fast_image_processor,
"mm_process_config": server_args.mm_process_config or {},
"model_revision": get_model().revision,
"processor_revision": get_model().revision,
"disable_fast_image_processor": get_mm().disable_fast_image_processor,
"mm_process_config": get_mm().mm_process_config or {},
"processor": processor_payload,
"extra": extra or {},
}
@@ -39,6 +39,7 @@ from sglang.srt.multimodal.transport.cuda_ipc import (
MmItemMemoryPool,
get_mm_feature_pool_size_per_worker,
)
from sglang.srt.runtime_context import get_mm
from sglang.srt.utils import (
CLIENT_MEDIA_EXCEPTIONS,
configure_media_url_security,
@@ -212,10 +213,10 @@ class BaseMultimodalProcessor(ABC):
self.server_args = server_args
self.transport_mode = transport_mode
configure_media_url_security(
server_args.allowed_media_domains,
get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb,
)
configured_mm_feature_transport = server_args.mm_feature_transport
configured_mm_feature_transport = get_mm().mm_feature_transport
self.mm_feature_transport = (
configured_mm_feature_transport
if configured_mm_feature_transport in ("cpu", "cuda_ipc", "cuda_vmm")
@@ -231,7 +232,7 @@ class BaseMultimodalProcessor(ABC):
self.disable_fast_image_processor = self.image_processor_backend == "pil"
self.skip_tokenizer_init = server_args.skip_tokenizer_init
mm_process_config = self.server_args.mm_process_config
mm_process_config = get_mm().mm_process_config
self.image_config = mm_process_config.get("image", {})
self.video_config = mm_process_config.get("video", {})
self.audio_config = mm_process_config.get("audio", {})
@@ -255,7 +256,7 @@ class BaseMultimodalProcessor(ABC):
# The fingerprint is needed only to build artifact keys. Avoid inspecting
# processor state when this processor will never retain artifacts.
self.processor_fingerprint = (
build_processor_fingerprint(self, hf_config, server_args)
build_processor_fingerprint(self, hf_config)
if self.mm_preprocess_cache.enabled
else None
)
@@ -39,6 +39,7 @@ from sglang.srt.multimodal.processors.mimo_audio import (
MiMoAudioPipeline,
)
from sglang.srt.multimodal.processors.qwen_vl import smart_nframes
from sglang.srt.runtime_context import get_device
from sglang.srt.utils import ImageData, VideoData
from sglang.srt.utils.common import download_remote_media
from sglang.utils import logger
@@ -1588,7 +1589,7 @@ class MiMoV2Processor(BaseMultimodalProcessor):
processor_config, "video_end_token_id"
)
self.use_image_processor_gpu = envs.SGLANG_ENCODER_IMAGE_PROCESSOR_USE_GPU.get()
device = server_args.device if self.use_image_processor_gpu else None
device = get_device().device if self.use_image_processor_gpu else None
self.mimo_processor = MiMoProcessor(
tokenizer=self._processor.tokenizer,
@@ -31,7 +31,11 @@ from sglang.srt.ray.engine import (
_get_bundle_node_ip,
_resolve_bundle_indices,
)
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import (
configured_attn_cp_size,
configured_pp_size,
get_parallel,
)
from sglang.srt.server_args import PortArgs, ServerArgs
from sglang.srt.utils.network import bind_port, get_zmq_socket, get_zmq_socket_on_host
@@ -144,7 +148,10 @@ class RayDataParallelController(DataParallelController):
for node_idx in range(nnodes):
bundle_idx = self.bundle_for_node[node_idx]
pp_range, tp_range, pp_per_node, tp_per_node = _calculate_rank_ranges(
nnodes, server_args.pp_size, server_args.tp_size, node_rank=node_idx
nnodes,
configured_pp_size(),
server_args.tp_size,
node_rank=node_idx,
)
for pp_rank in pp_range:
for tp_rank in tp_range:
@@ -161,7 +168,7 @@ class RayDataParallelController(DataParallelController):
tp_rank,
server_args.tp_size,
get_parallel().dp_size,
server_args.attn_cp_size,
configured_attn_cp_size(),
)
rank_port_args = PortArgs.init_new(
server_args, actual_dp_rank, worker_ports
@@ -202,7 +209,7 @@ class RayDataParallelController(DataParallelController):
world_size = _compute_world_size(server_args)
bundle_indices = _resolve_bundle_indices(self.pg, world_size)
ranks_per_tp_group = server_args.tp_size * server_args.pp_size
ranks_per_tp_group = server_args.tp_size * configured_pp_size()
if dp_rank is not None:
start_rank = dp_rank * ranks_per_tp_group
end_rank = start_rank + ranks_per_tp_group
@@ -232,7 +239,7 @@ class RayDataParallelController(DataParallelController):
tp_rank,
server_args.tp_size,
get_parallel().dp_size,
server_args.attn_cp_size,
configured_attn_cp_size(),
)
rank_port_args = PortArgs.init_new(
server_args, actual_dp_rank, worker_ports
+9 -9
View File
@@ -32,7 +32,7 @@ from sglang.srt.entrypoints.engine import (
)
from sglang.srt.environ import envs
from sglang.srt.ray.scheduler_actor import SchedulerActor
from sglang.srt.runtime_context import get_parallel
from sglang.srt.runtime_context import configured_pp_size, get_parallel
from sglang.srt.server_args import PortArgs, ServerArgs
logger = logging.getLogger(__name__)
@@ -109,8 +109,8 @@ def _compute_world_size(server_args: ServerArgs) -> int:
Normal: dp_size * tp_size * pp_size; DP attention: tp_size * pp_size.
"""
if get_parallel().enable_dp_attention:
return server_args.tp_size * server_args.pp_size
return get_parallel().dp_size * server_args.tp_size * server_args.pp_size
return server_args.tp_size * configured_pp_size()
return get_parallel().dp_size * server_args.tp_size * configured_pp_size()
def _resolve_bundle_indices(pg: PlacementGroup, world_size: int) -> List[int]:
@@ -269,10 +269,10 @@ class RayEngine(Engine):
)
if get_parallel().enable_dp_attention:
total_gpus = server_args.tp_size * server_args.pp_size
total_gpus = server_args.tp_size * configured_pp_size()
else:
total_gpus = (
get_parallel().dp_size * server_args.tp_size * server_args.pp_size
get_parallel().dp_size * server_args.tp_size * configured_pp_size()
)
nnodes = server_args.nnodes
@@ -332,7 +332,7 @@ class RayEngine(Engine):
pp_range, tp_range, pp_per_node, tp_per_node = (
_calculate_rank_ranges(
nnodes,
server_args.pp_size,
configured_pp_size(),
server_args.tp_size,
node_rank=node_idx,
)
@@ -449,16 +449,16 @@ class RayEngine(Engine):
if get_parallel().enable_dp_attention:
# DP attention folds DP into TP — total GPUs = tp_size * pp_size
total_gpus = server_args.tp_size * server_args.pp_size
total_gpus = server_args.tp_size * configured_pp_size()
else:
total_gpus = (
get_parallel().dp_size * server_args.tp_size * server_args.pp_size
get_parallel().dp_size * server_args.tp_size * configured_pp_size()
)
gpus_per_node = total_gpus // server_args.nnodes
logger.info(
f"Ray DP cluster: {server_args.nnodes} nodes, "
f"{gpus_per_node} GPUs/node, dp_size={get_parallel().dp_size}, "
f"tp_size={server_args.tp_size}, pp_size={server_args.pp_size}, "
f"tp_size={server_args.tp_size}, pp_size={configured_pp_size()}, "
f"enable_dp_attention={get_parallel().enable_dp_attention}"
)