config: the per-instance families read the bags (#35026)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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(),
|
||||
)
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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}"
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user