[diffusion] feat: data-parallel serving (--dp-size) (#33725)

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
This commit is contained in:
Mick
2026-08-06 11:50:13 +08:00
committed by GitHub
co-authored by Claude Fable 5
parent 407a65d3cb
commit 604d3561b0
8 changed files with 417 additions and 87 deletions
+1
View File
@@ -83,6 +83,7 @@ Use `sglang generate --help` and `sglang serve --help` for the full argument lis
- `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths; `speed` also enables `--enable-torch-compile` unless the model-specific deployment config opts out or you explicitly disable it. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes. - `--performance-mode {manual|auto|speed|memory}` / `--mode`: preset for latency/throughput and memory defaults. `auto` is the default and keeps safe offload defaults, using FSDP only for validated DiT-offload replacement paths; `speed` also enables `--enable-torch-compile` unless the model-specific deployment config opts out or you explicitly disable it. Use `manual` to keep performance-related server args under explicit user control. Explicit offload, FSDP, and parallelism flags take precedence in all modes.
- `--tp-size {N}`: tensor parallelism size, mainly for encoders - `--tp-size {N}`: tensor parallelism size, mainly for encoders
- `--sp-degree {N}`: sequence parallelism size - `--sp-degree {N}`: sequence parallelism size
- `--dp-size {N}` (alias `--data-parallel-size`): number of data-parallel replicas. Each replica is a full copy of the engine on `num_gpus / N` GPUs with its own ingress; generation requests round-robin across replicas, realtime sessions stick to the replica holding their state, and control operations (weights, LoRA, memory occupation, shutdown) apply to every replica. Combines with the other parallelism axes (`num_gpus = dp × cfg × tp × sp`); monolithic serving only.
- `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls - `--ulysses-degree {N}` and `--ring-degree {N}`: USP parallelism controls
- `--enable-cfg-parallel {true|false}`: enable or explicitly disable CFG parallelism - `--enable-cfg-parallel {true|false}`: enable or explicitly disable CFG parallelism
- `--encoder-parallel {auto|fold|dp|replicate}`: how the text/image encoders use the GPUs the DiT replica leaves idle during encoding. `auto` (the default for both `generate` and `serve`) TP-folds an encoder wide enough to pay for the per-layer all-reduce, selects DP for a server batch when it can engage, and otherwise replicates; `fold` forces the shard whenever the dims allow it; `dp` splits a batched encode across ranks and needs `--batching-max-size > 1` to engage; `replicate` encodes redundantly on every rank. `fold` and `replicate` are bitwise-identical to single-GPU encoding. See [Encoder Parallelism](/docs/sglang-diffusion/encoder_parallel). - `--encoder-parallel {auto|fold|dp|replicate}`: how the text/image encoders use the GPUs the DiT replica leaves idle during encoding. `auto` (the default for both `generate` and `serve`) TP-folds an encoder wide enough to pay for the per-layer all-reduce, selects DP for a server batch when it can engage, and otherwise replicates; `fold` forces the shard whenever the dims allow it; `dp` splits a batched encode across ranks and needs `--batching-max-size > 1` to engage; `replicate` encodes redundantly on every rank. `fold` and `replicate` are bitwise-identical to single-GPU encoding. See [Encoder Parallelism](/docs/sglang-diffusion/encoder_parallel).
@@ -137,6 +137,10 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
): ):
self.local_rank = local_rank self.local_rank = local_rank
self.rank = rank self.rank = rank
# the rank that materializes output and replies to the client: the
# first rank of this DP replica, which is global rank 0 only at dp=1
gpus_per_replica = max(1, server_args.num_gpus // (server_args.dp_size or 1))
self.is_output_rank = rank % gpus_per_replica == 0
self.master_port = master_port self.master_port = master_port
# FIXME: should we use tcp as distribute init method? # FIXME: should we use tcp as distribute init method?
self.server_args = server_args self.server_args = server_args
@@ -460,7 +464,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
output_batch = None output_batch = None
forward_failed = False forward_failed = False
try: try:
if self.rank == 0 and not current_platform.is_cpu(): if self.is_output_rank and not current_platform.is_cpu():
torch.get_device_module().reset_peak_memory_stats() torch.get_device_module().reset_peak_memory_stats()
start_time = ( start_time = (
@@ -474,7 +478,11 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
request_metrics = [ request_metrics = [
item.metrics for item in log_reqs if item.metrics is not None item.metrics for item in log_reqs if item.metrics is not None
] ]
if self.rank == 0 and request_metrics and not current_platform.is_cpu(): if (
self.is_output_rank
and request_metrics
and not current_platform.is_cpu()
):
baseline_snapshot = capture_memory_snapshot() baseline_snapshot = capture_memory_snapshot()
for metrics in request_metrics: for metrics in request_metrics:
metrics.record_memory_snapshot("before_forward", baseline_snapshot) metrics.record_memory_snapshot("before_forward", baseline_snapshot)
@@ -501,13 +509,13 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
self._record_output_peak_memory(output_batch) self._record_output_peak_memory(output_batch)
output_metrics = self._iter_output_metrics(output_batch) output_metrics = self._iter_output_metrics(output_batch)
if self.rank == 0 and output_metrics and not current_platform.is_cpu(): if self.is_output_rank and output_metrics and not current_platform.is_cpu():
peak_snapshot = capture_memory_snapshot() peak_snapshot = capture_memory_snapshot()
for metrics in output_metrics: for metrics in output_metrics:
metrics.record_memory_snapshot("after_forward", peak_snapshot) metrics.record_memory_snapshot("after_forward", peak_snapshot)
if ( if (
self.rank == 0 self.is_output_rank
and not req.suppress_logs and not req.suppress_logs
and not current_platform.is_cpu() and not current_platform.is_cpu()
and logger.isEnabledFor(logging.DEBUG) and logger.isEnabledFor(logging.DEBUG)
@@ -581,7 +589,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
def _materialize_raw_frame_transport( def _materialize_raw_frame_transport(
self, output_batch: OutputBatch, req: Req self, output_batch: OutputBatch, req: Req
) -> None: ) -> None:
if self.rank != 0: if not self.is_output_rank:
return return
if output_batch.output is not None: if output_batch.output is not None:
output_batch.raw_frame_content_type = RAW_RGB_CONTENT_TYPE output_batch.raw_frame_content_type = RAW_RGB_CONTENT_TYPE
@@ -603,7 +611,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
output_batch: OutputBatch, output_batch: OutputBatch,
save_output_paths: Callable[[OutputBatch], None], save_output_paths: Callable[[OutputBatch], None],
) -> None: ) -> None:
if self.rank == 0: if self.is_output_rank:
save_output_paths(output_batch) save_output_paths(output_batch)
output_batch.output = None output_batch.output = None
output_batch.audio = None output_batch.audio = None
@@ -614,7 +622,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
) -> None: ) -> None:
"""materialize the output from tensor to numpy frames for faster serialization""" """materialize the output from tensor to numpy frames for faster serialization"""
if ( if (
self.rank != 0 not self.is_output_rank
or output_batch.output is None or output_batch.output is None
or not getattr(req, "return_frames", False) or not getattr(req, "return_frames", False)
): ):
@@ -678,7 +686,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
return np.asarray(materialized.frames) return np.asarray(materialized.frames)
def _record_output_peak_memory(self, output_batch: OutputBatch) -> None: def _record_output_peak_memory(self, output_batch: OutputBatch) -> None:
if self.rank != 0 or current_platform.is_cpu(): if not self.is_output_rank or current_platform.is_cpu():
return return
peak_reserved_bytes = torch.get_device_module().max_memory_reserved() peak_reserved_bytes = torch.get_device_module().max_memory_reserved()
output_batch.peak_memory_mb = peak_reserved_bytes / (1024**2) output_batch.peak_memory_mb = peak_reserved_bytes / (1024**2)
@@ -691,7 +699,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
def _save_output_paths(self, req: Req, output_batch: OutputBatch) -> None: def _save_output_paths(self, req: Req, output_batch: OutputBatch) -> None:
"""save outputs to files""" """save outputs to files"""
if self.rank != 0 or output_batch.output is None: if not self.is_output_rank or output_batch.output is None:
return return
dynamic_output_paths = None dynamic_output_paths = None
@@ -742,7 +750,7 @@ class GPUWorker(GPUWorkerPostTrainingMixin):
reqs: list[Req], reqs: list[Req],
output_batch: OutputBatch, output_batch: OutputBatch,
) -> None: ) -> None:
if self.rank != 0 or output_batch.output is None: if not self.is_output_rank or output_batch.output is None:
return return
if len(output_batch.output) != len(reqs): if len(output_batch.output) != len(reqs):
raise RuntimeError( raise RuntimeError(
@@ -104,15 +104,23 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
set_global_server_args(server_args=server_args) set_global_server_args(server_args=server_args)
# Inter-process Communication # Each DP replica is a contiguous rank block (dp is the outermost
# layout axis); its first rank binds the replica's ingress, and the
# sp/cfg/tp broadcast relay in recv_reqs -- replica-internal by
# construction -- fans requests out within the replica only.
gpus_per_replica = max(1, server_args.num_gpus // server_args.dp_size)
self.dp_replica = gpu_id // gpus_per_replica
self.context = zmq.Context(io_threads=2) self.context = zmq.Context(io_threads=2)
endpoint = server_args.scheduler_endpoint if gpu_id % gpus_per_replica == 0:
if gpu_id == 0: endpoint = server_args.scheduler_endpoint_for(self.dp_replica)
# router allocates identify (envelope) for each connection # router allocates identify (envelope) for each connection
self.receiver, actual_endpoint = get_zmq_socket( self.receiver, actual_endpoint = get_zmq_socket(
self.context, zmq.ROUTER, endpoint, True self.context, zmq.ROUTER, endpoint, True
) )
logger.info(f"Scheduler bind at endpoint: {actual_endpoint}") logger.info(
f"Scheduler (dp replica {self.dp_replica}) bind at endpoint: "
f"{actual_endpoint}"
)
else: else:
self.receiver = None self.receiver = None
from sglang.multimodal_gen.runtime.platforms import current_platform from sglang.multimodal_gen.runtime.platforms import current_platform
@@ -1172,9 +1180,8 @@ class Scheduler(SchedulerWarmupMixin, SchedulerPostTrainingMixin, SchedulerDisag
self._disagg_event_loop() self._disagg_event_loop()
return return
logger.debug( if self.receiver is not None:
f"Rank 0 scheduler listening on tcp://*:{self.server_args.scheduler_port}" logger.debug("Driver scheduler of dp replica %d listening", self.dp_replica)
)
while self._running: while self._running:
# Update queue depth for metrics # Update queue depth for metrics
@@ -1,11 +1,29 @@
import itertools
import pickle import pickle
import time import time
import zlib
from typing import Any, Optional from typing import Any, Optional
import zmq import zmq
import zmq.asyncio import zmq.asyncio
from sglang.multimodal_gen.runtime.entrypoints.post_training.io_struct import (
GetWeightsChecksumReqInput,
ReleaseMemoryOccupationReqInput,
ResumeMemoryOccupationReqInput,
UpdateWeightFromDiskReqInput,
UpdateWeightFromTensorCheckerReqInput,
UpdateWeightFromTensorReqInput,
)
from sglang.multimodal_gen.runtime.entrypoints.utils import (
ListLorasReq,
MergeLoraWeightsReq,
SetLoraReq,
ShutdownReq,
UnmergeLoraWeightsReq,
)
from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs from sglang.multimodal_gen.runtime.ipc_array import materialize_file_refs
from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.server_args import ServerArgs from sglang.multimodal_gen.runtime.server_args import ServerArgs
from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger from sglang.multimodal_gen.runtime.utils.logging_utils import init_logger
@@ -15,6 +33,22 @@ from sglang.multimodal_gen.runtime.utils.request_logger import (
logger = init_logger(__name__) logger = init_logger(__name__)
# Control ops mutate replica state (weights, LoRA, memory, shutdown), so with
# DP they must reach every replica rather than one.
_CONTROL_REQ_TYPES = (
SetLoraReq,
MergeLoraWeightsReq,
UnmergeLoraWeightsReq,
ListLorasReq,
ShutdownReq,
UpdateWeightFromDiskReqInput,
UpdateWeightFromTensorReqInput,
UpdateWeightFromTensorCheckerReqInput,
GetWeightsChecksumReqInput,
ReleaseMemoryOccupationReqInput,
ResumeMemoryOccupationReqInput,
)
async def run_zeromq_broker(server_args: ServerArgs): async def run_zeromq_broker(server_args: ServerArgs):
""" """
@@ -49,6 +83,34 @@ async def run_zeromq_broker(server_args: ServerArgs):
pass pass
def _session_key(batch: Any) -> str | None:
"""Realtime sessions hold GPU state on one replica, so every request of a
session must land on the same one."""
reqs = batch if isinstance(batch, list) else [batch]
for req in reqs:
if isinstance(req, Req) and req.realtime_session_id:
return req.realtime_session_id
return None
def _select_replica(batch: Any, dp_size: int, counter: "itertools.count") -> int:
if dp_size <= 1:
return 0
session = _session_key(batch)
if session is not None:
return zlib.crc32(session.encode()) % dp_size
return next(counter) % dp_size
def _merge_fanout_results(results: list[Any]) -> Any:
"""One reply for a control op sent to every replica: the first error wins,
because "succeeded" must mean succeeded everywhere."""
for result in results:
if isinstance(result, OutputBatch) and result.error:
return result
return results[0]
class SchedulerClient: class SchedulerClient:
""" """
A synchronous, singleton client for communicating with the Scheduler service. A synchronous, singleton client for communicating with the Scheduler service.
@@ -57,9 +119,9 @@ class SchedulerClient:
def __init__(self): def __init__(self):
self.context = None self.context = None
self.scheduler_socket = None
self.server_args = None self.server_args = None
self.request_logger: Optional[DiffusionRequestLogger] = None self.request_logger: Optional[DiffusionRequestLogger] = None
self._replica_counter = itertools.count()
def initialize(self, server_args: ServerArgs): def initialize(self, server_args: ServerArgs):
if self.context is not None and not self.context.closed: if self.context is not None and not self.context.closed:
@@ -69,39 +131,38 @@ class SchedulerClient:
self.server_args = server_args self.server_args = server_args
self.request_logger = DiffusionRequestLogger.from_server_args(server_args) self.request_logger = DiffusionRequestLogger.from_server_args(server_args)
self.context = zmq.Context() self.context = zmq.Context()
self.scheduler_socket = self.context.socket(zmq.REQ)
# Set socket options for the main communication socket
self.scheduler_socket.setsockopt(zmq.LINGER, 0)
# 100 minute timeout for generation
self.scheduler_socket.setsockopt(zmq.RCVTIMEO, 6000000)
scheduler_endpoint = self.server_args.scheduler_endpoint
self.scheduler_socket.connect(scheduler_endpoint)
logger.debug(
f"SchedulerClient connected to backend scheduler at {scheduler_endpoint}"
)
def forward(self, batch: Any, timeout_ms: int | None = None) -> Any: def forward(self, batch: Any, timeout_ms: int | None = None) -> Any:
"""Sends a batch or request to the scheduler and waits for the response.""" """Sends a batch or request to the scheduler and waits for the response."""
self.request_logger.log_received_request(batch) return self._forward_routed(batch, timeout_ms)
previous_timeout_ms = None
if timeout_ms is not None: def _forward_one(self, endpoint: str, batch: Any, timeout_ms: int | None) -> Any:
previous_timeout_ms = self.scheduler_socket.getsockopt(zmq.RCVTIMEO) socket = self.context.socket(zmq.REQ)
self.scheduler_socket.setsockopt(zmq.RCVTIMEO, timeout_ms) socket.setsockopt(zmq.LINGER, 0)
socket.setsockopt(zmq.RCVTIMEO, timeout_ms if timeout_ms else 6000000)
try: try:
self.scheduler_socket.send_pyobj(batch) socket.connect(endpoint)
output_batch = self.scheduler_socket.recv_pyobj() socket.send_pyobj(batch)
output_batch = socket.recv_pyobj()
_materialize_output_batch_file_refs(output_batch) _materialize_output_batch_file_refs(output_batch)
self.request_logger.log_finished_request(batch, output_batch)
return output_batch return output_batch
except zmq.error.Again: except zmq.error.Again:
logger.error("Timeout waiting for response from scheduler.") logger.error("Timeout waiting for response from %s.", endpoint)
raise TimeoutError("Scheduler did not respond in time.") raise TimeoutError("Scheduler did not respond in time.")
finally: finally:
if previous_timeout_ms is not None and self.scheduler_socket is not None: socket.close()
self.scheduler_socket.setsockopt(zmq.RCVTIMEO, previous_timeout_ms)
def _forward_routed(self, batch: Any, timeout_ms: int | None) -> Any:
self.request_logger.log_received_request(batch)
endpoints = self.server_args.scheduler_endpoints
if isinstance(batch, _CONTROL_REQ_TYPES):
results = [self._forward_one(ep, batch, timeout_ms) for ep in endpoints]
output_batch = _merge_fanout_results(results)
else:
replica = _select_replica(batch, len(endpoints), self._replica_counter)
output_batch = self._forward_one(endpoints[replica], batch, timeout_ms)
self.request_logger.log_finished_request(batch, output_batch)
return output_batch
def ping(self) -> bool: def ping(self) -> bool:
""" """
@@ -111,27 +172,22 @@ class SchedulerClient:
logger.error("Cannot ping: client is not initialized.") logger.error("Cannot ping: client is not initialized.")
return False return False
for endpoint in self.server_args.scheduler_endpoints:
ping_socket = self.context.socket(zmq.REQ) ping_socket = self.context.socket(zmq.REQ)
ping_socket.setsockopt(zmq.LINGER, 0) ping_socket.setsockopt(zmq.LINGER, 0)
ping_socket.setsockopt(zmq.RCVTIMEO, 2000) # 2-second timeout for pings ping_socket.setsockopt(zmq.RCVTIMEO, 2000) # 2-second timeout for pings
endpoint = self.server_args.scheduler_endpoint
try: try:
ping_socket.connect(endpoint) ping_socket.connect(endpoint)
ping_socket.send_pyobj({"method": "ping"}) ping_socket.send_pyobj({"method": "ping"})
ping_socket.recv_pyobj() ping_socket.recv_pyobj()
return True
except zmq.error.Again: except zmq.error.Again:
return False return False
finally: finally:
ping_socket.close() ping_socket.close()
return True
def close(self): def close(self):
"""Closes the socket and terminates the context.""" """Terminates the context."""
if self.scheduler_socket:
self.scheduler_socket.close()
self.scheduler_socket = None
if self.context: if self.context:
self.context.term() self.context.term()
self.context = None self.context = None
@@ -150,6 +206,7 @@ class AsyncSchedulerClient:
self.context = None self.context = None
self.server_args = None self.server_args = None
self.request_logger: Optional[DiffusionRequestLogger] = None self.request_logger: Optional[DiffusionRequestLogger] = None
self._replica_counter = itertools.count()
def initialize(self, server_args: ServerArgs): def initialize(self, server_args: ServerArgs):
if self.context is not None and not self.context.closed: if self.context is not None and not self.context.closed:
@@ -171,24 +228,33 @@ class AsyncSchedulerClient:
"AsyncSchedulerClient is not initialized. Call initialize() first." "AsyncSchedulerClient is not initialized. Call initialize() first."
) )
# Create a temporary REQ socket for this request to allow concurrency endpoints = self.server_args.scheduler_endpoints
if isinstance(batch, _CONTROL_REQ_TYPES):
# replica state (weights, LoRA, memory) must change everywhere
results = [await self._forward_one(ep, batch) for ep in endpoints]
output_batch = _merge_fanout_results(results)
else:
replica = _select_replica(batch, len(endpoints), self._replica_counter)
output_batch = await self._forward_one(endpoints[replica], batch)
self.request_logger.log_finished_request(batch, output_batch)
return output_batch
async def _forward_one(self, endpoint: str, batch: Any) -> Any:
# a temporary REQ socket per request keeps concurrent requests from
# interleaving on one socket's strict send/recv alternation
socket = self.context.socket(zmq.REQ) socket = self.context.socket(zmq.REQ)
socket.setsockopt(zmq.LINGER, 0) socket.setsockopt(zmq.LINGER, 0)
# 100 minute timeout # 100 minute timeout
socket.setsockopt(zmq.RCVTIMEO, 6000000) socket.setsockopt(zmq.RCVTIMEO, 6000000)
endpoint = self.server_args.scheduler_endpoint
socket.connect(endpoint) socket.connect(endpoint)
try: try:
await socket.send(pickle.dumps(batch)) await socket.send(pickle.dumps(batch))
payload = await socket.recv() payload = await socket.recv()
output_batch = pickle.loads(payload) output_batch = pickle.loads(payload)
_materialize_output_batch_file_refs(output_batch) _materialize_output_batch_file_refs(output_batch)
self.request_logger.log_finished_request(batch, output_batch)
return output_batch return output_batch
except zmq.error.Again: except zmq.error.Again:
logger.error("Timeout waiting for response from scheduler.") logger.error("Timeout waiting for response from %s.", endpoint)
raise TimeoutError("Scheduler did not respond in time.") raise TimeoutError("Scheduler did not respond in time.")
finally: finally:
socket.close() socket.close()
@@ -201,21 +267,19 @@ class AsyncSchedulerClient:
logger.error("Cannot ping: client is not initialized.") logger.error("Cannot ping: client is not initialized.")
return False return False
for endpoint in self.server_args.scheduler_endpoints:
ping_socket = self.context.socket(zmq.REQ) ping_socket = self.context.socket(zmq.REQ)
ping_socket.setsockopt(zmq.LINGER, 0) ping_socket.setsockopt(zmq.LINGER, 0)
ping_socket.setsockopt(zmq.RCVTIMEO, 2000) ping_socket.setsockopt(zmq.RCVTIMEO, 2000)
endpoint = self.server_args.scheduler_endpoint
try: try:
ping_socket.connect(endpoint) ping_socket.connect(endpoint)
await ping_socket.send(pickle.dumps({"method": "ping"})) await ping_socket.send(pickle.dumps({"method": "ping"}))
await ping_socket.recv() await ping_socket.recv()
return True
except zmq.error.Again: except zmq.error.Again:
return False return False
finally: finally:
ping_socket.close() ping_socket.close()
return True
def close(self): def close(self):
"""Closes the socket and terminates the context.""" """Closes the socket and terminates the context."""
@@ -222,7 +222,6 @@ class ServerArgs(DisaggServerArgsMixin):
# number of data parallelism groups # number of data parallelism groups
dp_size: int = 1 dp_size: int = 1
# number of gpu in a dp group # number of gpu in a dp group
dp_degree: int = 1
# cfg parallel (None = auto-decide based on num_gpus) # cfg parallel (None = auto-decide based on num_gpus)
enable_cfg_parallel: Optional[bool] = None enable_cfg_parallel: Optional[bool] = None
# number of GPUs in each CFG parallel group (None = auto, 1 = disabled, N > 1 = enabled) # number of GPUs in each CFG parallel group (None = auto, 1 = disabled, N > 1 = enabled)
@@ -356,6 +355,8 @@ class ServerArgs(DisaggServerArgsMixin):
webui_port: int | None = 12312 webui_port: int | None = 12312
scheduler_port: int = 5555 scheduler_port: int = 5555
# settled ingress ports, one per DP replica; None until ports are settled
scheduler_ports: list[int] | None = None
batching_mode: str = "dynamic" batching_mode: str = "dynamic"
batching_max_size: int = 1 batching_max_size: int = 1
batching_delay_ms: float = 0.0 batching_delay_ms: float = 0.0
@@ -974,7 +975,10 @@ class ServerArgs(DisaggServerArgsMixin):
requested_ports = [] requested_ports = []
if needs_http: if needs_http:
requested_ports.append((self.port, "HTTP")) requested_ports.append((self.port, "HTTP"))
requested_ports.append((self.scheduler_port, "Scheduler")) for replica in range(self.dp_size or 1):
requested_ports.append(
(self.scheduler_port + replica, f"Scheduler[{replica}]")
)
if self.master_port is not None: if self.master_port is not None:
requested_ports.append((self.master_port, "Master")) requested_ports.append((self.master_port, "Master"))
seen_ports: dict[int, str] = {} seen_ports: dict[int, str] = {}
@@ -998,6 +1002,13 @@ class ServerArgs(DisaggServerArgsMixin):
initial_scheduler_port, avoid=settled_ports initial_scheduler_port, avoid=settled_ports
) )
settled_ports.add(self.scheduler_port) settled_ports.add(self.scheduler_port)
self.scheduler_ports = [self.scheduler_port]
for _ in range((self.dp_size or 1) - 1):
port = self.settle_port(
self.scheduler_ports[-1] + 1, avoid=settled_ports
)
settled_ports.add(port)
self.scheduler_ports.append(port)
if self.master_port is not None: if self.master_port is not None:
self.master_port = self.settle_port( self.master_port = self.settle_port(
self.master_port, 37, avoid=settled_ports self.master_port, 37, avoid=settled_ports
@@ -2047,10 +2058,22 @@ class ServerArgs(DisaggServerArgsMixin):
Internal endpoint for scheduler. Internal endpoint for scheduler.
Prefers the configured host but normalizes localhost -> 127.0.0.1 to avoid ZMQ issues. Prefers the configured host but normalizes localhost -> 127.0.0.1 to avoid ZMQ issues.
""" """
return self.scheduler_endpoint_for(0)
def scheduler_endpoint_for(self, replica: int) -> str:
"""Ingress endpoint of one DP replica's driver rank."""
scheduler_host = self.host scheduler_host = self.host
if scheduler_host is None or scheduler_host == "localhost": if scheduler_host is None or scheduler_host == "localhost":
scheduler_host = "127.0.0.1" scheduler_host = "127.0.0.1"
return f"tcp://{scheduler_host}:{self.scheduler_port}" if self.scheduler_ports is not None:
port = self.scheduler_ports[replica]
else:
port = self.scheduler_port + replica
return f"tcp://{scheduler_host}:{port}"
@property
def scheduler_endpoints(self) -> list[str]:
return [self.scheduler_endpoint_for(r) for r in range(self.dp_size or 1)]
def settle_port( def settle_port(
self, self,
@@ -2459,8 +2482,11 @@ class ServerArgs(DisaggServerArgsMixin):
if self.dp_size < 1: if self.dp_size < 1:
raise ValueError("--dp-size must be a natural number") raise ValueError("--dp-size must be a natural number")
if self.dp_size > 1: if self.dp_size > 1 and self.disagg_role != RoleType.MONOLITHIC:
raise ValueError("DP is not yet supported") raise ValueError(
"--dp-size > 1 is only supported for monolithic serving; "
"disaggregated roles scale by adding role instances instead"
)
num_gpus_per_group = self.dp_size * self.tp_size num_gpus_per_group = self.dp_size * self.tp_size
if self.enable_cfg_parallel: if self.enable_cfg_parallel:
@@ -1144,6 +1144,7 @@ STANDALONE_FILES = {
"../single_test_file/test_disagg_server.py", "../single_test_file/test_disagg_server.py",
"../single_test_file/test_ar_models.py", "../single_test_file/test_ar_models.py",
"../single_test_file/test_ipc_a2a_2_gpu.py", "../single_test_file/test_ipc_a2a_2_gpu.py",
"../single_test_file/test_dp_serving_2_gpu.py",
], ],
} }
@@ -1178,6 +1179,8 @@ STANDALONE_FILE_EST_TIMES = {
"../single_test_file/test_ar_models.py": 600.0, "../single_test_file/test_ar_models.py": 600.0,
# no model load; the cost is the one-time JIT build of the sync kernels # no model load; the cost is the one-time JIT build of the sync kernels
"../single_test_file/test_ipc_a2a_2_gpu.py": 240.0, "../single_test_file/test_ipc_a2a_2_gpu.py": 240.0,
# zimage server startup dominates; six short requests after warmup
"../single_test_file/test_dp_serving_2_gpu.py": 900.0,
}, },
} }
@@ -0,0 +1,150 @@
"""Data-parallel serving must route, replicate, and agree.
Launches one server with --dp-size 2 (one GPU per replica) and checks the three
properties that make DP real rather than a parsed flag: both replica drivers
bind their own ingress, both replicas serve traffic (round-robin means two
sequential requests land on different replicas), and a fixed seed produces the
same image bytes from either replica.
pytest -v python/sglang/multimodal_gen/test/single_test_file/test_dp_serving_2_gpu.py
"""
from __future__ import annotations
import base64
import hashlib
import json
import os
import signal
import subprocess
import sys
import time
import unittest
import urllib.request
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
import torch
from sglang.multimodal_gen.runtime.platforms import current_platform
from sglang.test.test_utils import CustomTestCase
_MODEL = "Tongyi-MAI/Z-Image-Turbo"
_PORT = 30811
_STARTUP_TIMEOUT_S = 1200
def _post_generation(prompt: str, seed: int) -> dict:
payload = json.dumps(
{
"prompt": prompt,
"size": "512x512",
"seed": seed,
"num_inference_steps": 20,
"response_format": "b64_json",
}
).encode()
req = urllib.request.Request(
f"http://127.0.0.1:{_PORT}/v1/images/generations",
data=payload,
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(req, timeout=600) as resp:
return json.loads(resp.read())
def _image_md5(response: dict) -> str:
b64 = response["data"][0]["b64_json"]
return hashlib.md5(base64.b64decode(b64)).hexdigest()
class TestDpServingTwoGpu(CustomTestCase):
def test_two_replicas_serve_and_agree(self):
if not current_platform.is_cuda():
self.skipTest("DP e2e is exercised on CUDA")
if torch.cuda.device_count() < 2:
self.skipTest("needs 2 GPUs")
log_path = Path("/tmp/dp_serving_test.log")
fh = open(log_path, "w")
proc = subprocess.Popen(
[
sys.executable,
"-m",
"sglang.multimodal_gen.runtime.entrypoints.cli.main",
"serve",
"--model-path",
_MODEL,
"--num-gpus",
"2",
"--dp-size",
"2",
"--enable-cfg-parallel",
"false",
"--port",
str(_PORT),
],
stdout=fh,
stderr=subprocess.STDOUT,
preexec_fn=os.setsid,
env=os.environ.copy(),
)
try:
deadline = time.monotonic() + _STARTUP_TIMEOUT_S
while time.monotonic() < deadline:
try:
with urllib.request.urlopen(
f"http://127.0.0.1:{_PORT}/health", timeout=5
) as resp:
if resp.status == 200:
break
except Exception:
time.sleep(5)
else:
self.fail(f"server not healthy in {_STARTUP_TIMEOUT_S}s")
# warm each replica once so first-request costs stay out of timing
for _ in range(2):
_post_generation("warm pear", seed=7)
# same seed through two round-robined replicas -> identical bytes
first = _post_generation("a pear on a table", seed=42)
second = _post_generation("a pear on a table", seed=42)
self.assertEqual(_image_md5(first), _image_md5(second))
log = log_path.read_text(errors="ignore")
self.assertIn("dp replica 0) bind", log)
self.assertIn("dp replica 1) bind", log)
# distribution, not just agreement: a concurrent pair must run in
# about one request's wall time. Two requests serialized on a
# single replica would take ~2x the single-request time, so the
# 1.6x bound separates the behaviors with margin for jitter.
t0 = time.monotonic()
single = _post_generation("a pear on a table", seed=42)
single_s = time.monotonic() - t0
self.assertTrue(single["data"])
t0 = time.monotonic()
with ThreadPoolExecutor(max_workers=2) as pool:
futures = [
pool.submit(_post_generation, "a pear on a table", 42)
for _ in range(2)
]
results = [f.result() for f in futures]
pair_s = time.monotonic() - t0
for r in results:
self.assertTrue(r["data"])
self.assertLess(
pair_s,
1.6 * single_s,
f"concurrent pair took {pair_s:.2f}s vs single {single_s:.2f}s; "
"requests are serializing on one replica",
)
finally:
os.killpg(os.getpgid(proc.pid), signal.SIGKILL)
fh.close()
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,71 @@
"""Routing invariants for DP serving.
Generation requests go to exactly one replica, control ops reach every replica,
and a realtime session always lands on the same replica it started on.
"""
import itertools
from sglang.multimodal_gen.runtime.entrypoints.utils import SetLoraReq, ShutdownReq
from sglang.multimodal_gen.runtime.pipelines_core import Req
from sglang.multimodal_gen.runtime.pipelines_core.schedule_batch import OutputBatch
from sglang.multimodal_gen.runtime.scheduler_client import (
_CONTROL_REQ_TYPES,
_merge_fanout_results,
_select_replica,
)
from sglang.multimodal_gen.runtime.server_args import ServerArgs
def _req(**kwargs) -> Req:
return Req(prompt="a pear", **kwargs)
def test_round_robin_covers_every_replica():
counter = itertools.count()
picks = [_select_replica([_req()], 3, counter) for _ in range(6)]
assert picks == [0, 1, 2, 0, 1, 2]
def test_dp1_always_selects_replica_zero():
counter = itertools.count()
assert {_select_replica([_req()], 1, counter) for _ in range(5)} == {0}
def test_session_requests_stick_to_one_replica():
counter = itertools.count()
session_req = _req(realtime_session_id="session-abc")
picks = {_select_replica([session_req], 4, counter) for _ in range(8)}
assert len(picks) == 1
# and the counter was never consumed by session traffic
assert _select_replica([_req()], 4, counter) == 0
def test_control_reqs_are_recognized():
assert isinstance(SetLoraReq(lora_nickname="x", lora_path="y"), _CONTROL_REQ_TYPES)
assert isinstance(ShutdownReq(), _CONTROL_REQ_TYPES)
assert not isinstance([_req()], _CONTROL_REQ_TYPES)
assert not isinstance(_req(), _CONTROL_REQ_TYPES)
def test_fanout_merge_surfaces_the_failing_replica():
ok = OutputBatch(output=None)
bad = OutputBatch(error="replica 1 exploded")
assert _merge_fanout_results([ok, bad]) is bad
assert _merge_fanout_results([ok, OutputBatch(output=None)]) is ok
def test_scheduler_endpoints_one_per_replica():
args = ServerArgs.__new__(ServerArgs)
args.host = "localhost"
args.dp_size = 3
args.scheduler_port = 6000
args.scheduler_ports = None
assert ServerArgs.scheduler_endpoints.fget(args) == [
"tcp://127.0.0.1:6000",
"tcp://127.0.0.1:6001",
"tcp://127.0.0.1:6002",
]
# settled ports need not be consecutive; the endpoint list follows them
args.scheduler_ports = [6000, 7005, 7100]
assert ServerArgs.scheduler_endpoint_for(args, 2) == "tcp://127.0.0.1:7100"