Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion python/sglang/srt/constrained/grammar_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
from sglang.srt.constrained.reasoner_grammar_backend import ReasonerGrammarObject
from sglang.srt.distributed.communication_tags import P2PTag
from sglang.srt.environ import envs
from sglang.srt.runtime_context import get_serving
from sglang.srt.sampling.sampling_params import (
get_request_reasoning_end_token_ids,
)
Expand All @@ -31,7 +32,7 @@ def __init__(self, scheduler: Scheduler):
self.scheduler = scheduler
self.server_args = scheduler.server_args
self.grammar_queue: List[Req] = []
if not self.server_args.skip_tokenizer_init:
if not get_serving().skip_tokenizer_init:
self.grammar_backend = create_grammar_backend(
self.server_args,
scheduler.tokenizer,
Expand Down
14 changes: 5 additions & 9 deletions python/sglang/srt/disaggregation/decode.py
Original file line number Diff line number Diff line change
Expand Up @@ -468,11 +468,7 @@ def _swa_tail_len(self, seq_len: int) -> int:
return seq_len

page_size = self.token_to_kv_pool_allocator.page_size
if getattr(
self.scheduler.server_args,
"disaggregation_decode_enable_radix_cache",
False,
):
if get_disagg().disaggregation_decode_enable_radix_cache:
# Keep enough SWA before the page-aligned radix-cache insert
# boundary for the cached key to contain a complete window.
# `seq_len - 1` is the last committed position.
Expand Down Expand Up @@ -1194,7 +1190,7 @@ def pop_preallocated(
origin_input_len = self._rebootstrap_prefill_len(decode_req.req)
prefix_match: Optional[DecodePrefixMatch] = None
use_decode_radix_cache = (
self.scheduler.server_args.disaggregation_decode_enable_radix_cache
get_disagg().disaggregation_decode_enable_radix_cache
and not decode_req.is_rebootstrap
)
if use_decode_radix_cache:
Expand Down Expand Up @@ -1648,13 +1644,13 @@ def _allocatable_token_budgets(
available_size = logical_allocator.available_size()
elif self._uses_swa_tail_prealloc():
available_size = self.token_to_kv_pool_allocator.full_available_size()
if self.scheduler.server_args.disaggregation_decode_enable_radix_cache:
if get_disagg().disaggregation_decode_enable_radix_cache:
available_size += self._radix_full_evictable()
else:
available_size = self.token_to_kv_pool_allocator.available_size()
# Include evictable decode-radix cache entries in the budget -- they
# can be freed on demand before allocation.
if self.scheduler.server_args.disaggregation_decode_enable_radix_cache:
if get_disagg().disaggregation_decode_enable_radix_cache:
available_size += self._radix_full_evictable()
allocatable_tokens = available_size - max(
reserved_tokens, need_space_for_single_req
Expand Down Expand Up @@ -1801,7 +1797,7 @@ def _pre_alloc(

# Evict cached entries if the pool doesn't have enough free pages.
if (
self.scheduler.server_args.disaggregation_decode_enable_radix_cache
get_disagg().disaggregation_decode_enable_radix_cache
and self._radix_full_available() < required_alloc_tokens
):
num_to_evict = required_alloc_tokens - self._radix_full_available()
Expand Down
4 changes: 2 additions & 2 deletions python/sglang/srt/disaggregation/encoder/grpc_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -251,8 +251,8 @@ async def serve_grpc_encoder(server_args: ServerArgs):
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(
Expand Down
16 changes: 8 additions & 8 deletions python/sglang/srt/disaggregation/encoder/http_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -103,7 +103,7 @@ async def _lifespan(app: FastAPI):
app = FastAPI(lifespan=_lifespan)


def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
def _register_encoder_url_with_bootstrap():
"""Asynchronously register this encoder with each bootstrap URL.

Spawns a daemon thread that retries each URL independently with bounded
Expand All @@ -118,10 +118,10 @@ def _register_encoder_url_with_bootstrap(server_args: ServerArgs):
host = get_serving().host
if not host or host in ("0.0.0.0", "::"):
host = get_local_ip_auto(get_serving().host)
scheme = "https" if server_args.ssl_certfile else "http"
scheme = "https" if get_serving().ssl_certfile else "http"
encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme)
payload = {"url": encoder_url}
bootstrap_urls = list(server_args.encoder_register_urls)
bootstrap_urls = list(get_disagg().encoder_register_urls)
if not bootstrap_urls:
return

Expand Down Expand Up @@ -175,15 +175,15 @@ def _worker():
).start()


def _unregister_encoder_url_from_bootstrap(server_args: ServerArgs):
def _unregister_encoder_url_from_bootstrap():
host = get_serving().host
if not host or host in ("0.0.0.0", "::"):
host = get_local_ip_auto(get_serving().host)
scheme = "https" if server_args.ssl_certfile else "http"
scheme = "https" if get_serving().ssl_certfile else "http"
encoder_url = NetworkAddress(host, get_serving().port).to_url(scheme)
payload = {"url": encoder_url}

for bootstrap_url in server_args.encoder_register_urls:
for bootstrap_url in get_disagg().encoder_register_urls:
try:
resp = http_requests.delete(
f"{bootstrap_url}/unregister_encoder_url",
Expand Down Expand Up @@ -229,8 +229,8 @@ def launch_server(server_args: ServerArgs):
if get_disagg().encoder_register_urls:
import atexit

_register_encoder_url_with_bootstrap(server_args)
atexit.register(_unregister_encoder_url_from_bootstrap, server_args)
_register_encoder_url_with_bootstrap()
atexit.register(_unregister_encoder_url_from_bootstrap)

uvicorn.run(app, host=get_serving().host, port=get_serving().port)

Expand Down
2 changes: 1 addition & 1 deletion python/sglang/srt/disaggregation/encoder/receiver.py
Original file line number Diff line number Diff line change
Expand Up @@ -2000,7 +2000,7 @@ def _init_mm_processor(
import_processors("sglang.srt.multimodal.processors")

extra_kwargs = {}
if getattr(server_args, "tokenizer_backend", None) is not None:
if get_serving().tokenizer_backend is not None:
extra_kwargs["tokenizer_backend"] = get_serving().tokenizer_backend

_processor = get_processor(
Expand Down
7 changes: 4 additions & 3 deletions python/sglang/srt/disaggregation/encoder/runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,6 +52,7 @@
trace_set_thread_info,
)
from sglang.srt.runtime_context import (
get_device,
get_observability,
get_parallel,
get_serving,
Expand Down Expand Up @@ -1867,7 +1868,7 @@ def _kill_workers():
atexit.register(_kill_workers)

for dp_rank in range(dp_size):
gpu_id = server_args.base_gpu_id + dp_rank
gpu_id = get_device().base_gpu_id + dp_rank
# Pin the device parent-side around spawn (same convention as the
# scheduler launcher and DP controller) so the child inherits
# CUDA_VISIBLE_DEVICES from its first instruction, before any import
Expand All @@ -1890,8 +1891,8 @@ def _kill_workers():
worker_processes.append(process)

labels = {"model_name": get_serving().served_model_name}
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)
return DPDispatcher(
dp_size,
dispatch_sockets,
Expand Down
12 changes: 6 additions & 6 deletions python/sglang/srt/disaggregation/encoder/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -548,7 +548,7 @@ def __init__(
self.server_args = server_args
configure_media_url_security(
get_mm().allowed_media_domains,
server_args.media_url_max_file_size_mb,
get_mm().media_url_max_file_size_mb,
)
self.transfer_backend = get_disagg().encoder_transfer_backend
self.use_mooncake = self.transfer_backend == "mooncake"
Expand All @@ -563,18 +563,18 @@ def __init__(
)
self.load_config = LoadConfig(
load_format=get_model().load_format,
download_dir=server_args.download_dir,
download_dir=get_model().download_dir,
model_loader_extra_config=get_model().model_loader_extra_config,
remote_instance_weight_loader_seed_instance_ip=server_args.remote_instance_weight_loader_seed_instance_ip,
remote_instance_weight_loader_seed_instance_service_port=server_args.remote_instance_weight_loader_seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=server_args.remote_instance_weight_loader_send_weights_group_ports,
remote_instance_weight_loader_seed_instance_ip=get_model().remote_instance_weight_loader_seed_instance_ip,
remote_instance_weight_loader_seed_instance_service_port=get_model().remote_instance_weight_loader_seed_instance_service_port,
remote_instance_weight_loader_send_weights_group_ports=get_model().remote_instance_weight_loader_send_weights_group_ports,
)
self.model_type = getattr(
self.model_config.hf_config, "model_type", "unknown"
).lower()

self.device = get_device().device
self.gpu_id = server_args.base_gpu_id + rank if gpu_id is None else gpu_id
self.gpu_id = get_device().base_gpu_id + rank if gpu_id is None else gpu_id

self.device_config = DeviceConfig(
device=self.device,
Expand Down
19 changes: 8 additions & 11 deletions python/sglang/srt/distributed/bootstrap.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,14 +81,14 @@ def init_torch_distributed(
tic = time.perf_counter()
logger.info("Init torch distributed begin.")

backend = _resolve_backend(device=device, server_args=server_args)
backend = _resolve_backend(device=device)

before_avail_memory = get_available_gpu_memory(device, ps.gpu_id)
if not get_parallel().enable_p2p_check:
monkey_patch_p2p_access_check()

dist_init_method = _resolve_dist_init_method(dist_port=dist_port)
_set_all_reduce_flags(server_args=server_args)
_set_all_reduce_flags()

if not is_draft_worker:
if device == "cpu":
Expand Down Expand Up @@ -173,9 +173,9 @@ def init_torch_distributed(
)


def _resolve_backend(*, device: str, server_args: ServerArgs) -> str:
def _resolve_backend(*, device: str) -> str:
backend = get_default_distributed_backend(device)
if device == "cuda" and server_args.elastic_ep_backend == "mooncake":
if device == "cuda" and get_exec().moe.elastic_ep_backend == "mooncake":
backend = "mooncake"
return backend

Expand All @@ -198,9 +198,9 @@ def _resolve_dist_init_method(*, dist_port: int) -> str:
return dist_init_method


def _set_all_reduce_flags(*, server_args: ServerArgs) -> None:
def _set_all_reduce_flags() -> None:
set_custom_all_reduce(not get_exec().comm.disable_custom_all_reduce)
set_mscclpp_all_reduce(server_args.enable_mscclpp)
set_mscclpp_all_reduce(get_exec().comm.enable_mscclpp)
set_torch_symm_mem_all_reduce(get_exec().comm.enable_torch_symm_mem)
set_flashinfer_allreduce_only(
get_exec().comm.flashinfer_allreduce_fusion_backend is not None
Expand Down Expand Up @@ -295,7 +295,7 @@ def _init_parallel_groups(
duplicate_tp_group=get_disagg().enable_pdmux,
duplicate_attn_cp_group=(
is_hip()
and server_args.enable_two_batch_overlap
and get_exec().overlap.enable_two_batch_overlap
and get_parallel().enable_dsa_prefill_context_parallel
),
enable_symm_mem=get_exec().comm.enable_symm_mem,
Expand All @@ -308,10 +308,7 @@ def _init_parallel_groups(
server_args=server_args,
model_config=model_config,
)
initialize_layernorm_sp(
server_args=server_args,
model_config=model_config,
)
initialize_layernorm_sp(model_config=model_config)
if is_npu():
register_sgl_tp_rank(gpu_id)

Expand Down
5 changes: 3 additions & 2 deletions python/sglang/srt/entrypoints/elastic_ep.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
from fastapi import APIRouter, Request
from fastapi.responses import ORJSONResponse

from sglang.srt.runtime_context import get_exec
from sglang.srt.utils.auth import AuthLevel, auth_level

router = APIRouter()
Expand Down Expand Up @@ -43,7 +44,7 @@ async def scale_elastic_ep(raw_request: Request):
from sglang.srt.entrypoints.http_server import _global_state
from sglang.srt.managers.io_struct import ScaleElasticEPReqInput

if _global_state.tokenizer_manager.server_args.elastic_ep_backend is None:
if get_exec().moe.elastic_ep_backend is None:
return ORJSONResponse(
{"error": "elastic EP is not enabled (set --elastic-ep-backend)"},
status_code=HTTPStatus.NOT_FOUND,
Expand Down Expand Up @@ -78,7 +79,7 @@ async def is_scaling_elastic_ep(raw_request: Request):
"""Return the tokenizer's mirrored Elastic EP scale state."""
from sglang.srt.entrypoints.http_server import _global_state

if _global_state.tokenizer_manager.server_args.elastic_ep_backend is None:
if get_exec().moe.elastic_ep_backend is None:
return ORJSONResponse(
{"error": "elastic EP is not enabled (set --elastic-ep-backend)"},
status_code=HTTPStatus.NOT_FOUND,
Expand Down
Loading
Loading