[Ray] Add data parallel (DP) and DP attention support to RayEngine (#21887)
Co-authored-by: xyuzh <xyuzh@users.noreply.github.com>
This commit is contained in:
@@ -0,0 +1,253 @@
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
# ==============================================================================
|
||||
"""Ray-aware DataParallelController that launches SchedulerActors instead of mp.Process."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import List, Optional
|
||||
|
||||
import ray
|
||||
import zmq
|
||||
from ray.util.scheduling_strategies import PlacementGroupSchedulingStrategy
|
||||
|
||||
from sglang.srt.entrypoints.engine import (
|
||||
_calculate_rank_ranges,
|
||||
_compute_parallelism_ranks,
|
||||
)
|
||||
from sglang.srt.layers.dp_attention import compute_dp_attention_world_info
|
||||
from sglang.srt.managers.data_parallel_controller import DataParallelController
|
||||
from sglang.srt.ray.scheduler_actor import SchedulerActor
|
||||
from sglang.srt.server_args import PortArgs, ServerArgs
|
||||
from sglang.srt.utils.network import bind_port, get_zmq_socket, get_zmq_socket_on_host
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RayDataParallelController(DataParallelController):
|
||||
"""DataParallelController that uses Ray actors for scheduler processes.
|
||||
|
||||
Overrides the process-spawning methods to create SchedulerActor Ray actors
|
||||
instead of mp.Process. Runs in-process (not as a separate mp.Process) and
|
||||
reuses the parent's event_loop, dispatching, and ZMQ routing.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
server_args: ServerArgs,
|
||||
port_args: PortArgs,
|
||||
placement_group,
|
||||
bundle_for_node: List[int],
|
||||
rank0_node_ip: str,
|
||||
):
|
||||
# Set Ray-specific attributes BEFORE super().__init__() because the
|
||||
# parent constructor calls launch_dp_schedulers / launch_dp_attention_schedulers
|
||||
# which we override, and those methods need these attributes.
|
||||
self.pg = placement_group
|
||||
self.bundle_for_node = bundle_for_node
|
||||
self.rank0_node_ip = rank0_node_ip
|
||||
self.scheduler_actors: List = []
|
||||
self.event_loop_refs: List = []
|
||||
|
||||
# super().__init__ will call our overridden launch methods via MRO.
|
||||
# Pass run_scheduler_process_func=None since we don't spawn mp.Process.
|
||||
super().__init__(server_args, port_args, run_scheduler_process_func=None)
|
||||
|
||||
def launch_dp_schedulers(self, server_args: ServerArgs, port_args: PortArgs):
|
||||
"""Override: launch Ray scheduler actors per DP rank."""
|
||||
sockets = []
|
||||
dp_port_args_list = []
|
||||
|
||||
for dp_rank in range(server_args.dp_size):
|
||||
tmp_port_args = PortArgs.init_new(server_args)
|
||||
tmp_port_args.tokenizer_ipc_name = port_args.tokenizer_ipc_name
|
||||
tmp_port_args.detokenizer_ipc_name = port_args.detokenizer_ipc_name
|
||||
|
||||
# Hold NCCL port so the next DP rank gets a different one
|
||||
sockets.append(bind_port(tmp_port_args.nccl_port))
|
||||
dp_port_args_list.append(tmp_port_args)
|
||||
|
||||
# Create ZMQ PUSH socket for this DP rank (controller → scheduler)
|
||||
if server_args.node_rank == 0:
|
||||
self.workers[dp_rank] = get_zmq_socket(
|
||||
self.context,
|
||||
zmq.PUSH,
|
||||
tmp_port_args.scheduler_input_ipc_name,
|
||||
True,
|
||||
)
|
||||
|
||||
# Release held ports before creating actors
|
||||
for sock in sockets:
|
||||
sock.close()
|
||||
|
||||
# Create actors for each DP rank sequentially
|
||||
for dp_rank in range(server_args.dp_size):
|
||||
self._launch_ray_tp_group(
|
||||
server_args, dp_port_args_list[dp_rank], dp_rank
|
||||
)
|
||||
|
||||
def launch_dp_attention_schedulers(
|
||||
self, server_args: ServerArgs, port_args: PortArgs
|
||||
):
|
||||
"""Override: pre-allocate ports, skip broadcast, create Ray actors."""
|
||||
# Pre-allocate worker ports on the controller node, binding to the
|
||||
# rank-0 node IP instead of tcp://* to avoid exposing unauthenticated
|
||||
# ZMQ sockets (CVE-2026-3060).
|
||||
worker_ports = []
|
||||
for dp_rank in range(server_args.dp_size):
|
||||
worker_port, worker_socket = get_zmq_socket_on_host(
|
||||
self.context, zmq.PUSH, host=self.rank0_node_ip
|
||||
)
|
||||
worker_ports.append(worker_port)
|
||||
self.workers[dp_rank] = worker_socket
|
||||
logger.debug(f"Assigned port {worker_port} to worker {dp_rank}")
|
||||
|
||||
# Skip _broadcast_worker_ports — Ray creates all actors centrally,
|
||||
# so there's no need for the inter-node handshake protocol.
|
||||
self._launch_ray_tp_group(
|
||||
server_args, port_args, dp_rank=None, worker_ports=worker_ports
|
||||
)
|
||||
|
||||
def _launch_ray_tp_group(
|
||||
self,
|
||||
server_args: ServerArgs,
|
||||
port_args: PortArgs,
|
||||
dp_rank: Optional[int],
|
||||
worker_ports: Optional[List[int]] = None,
|
||||
):
|
||||
"""Create SchedulerActor Ray actors for one TP group (one DP rank).
|
||||
|
||||
For DP attention, dp_rank=None and worker_ports is provided; the dp_rank
|
||||
is derived from tp_rank via compute_dp_attention_world_info.
|
||||
|
||||
For regular DP, dp_rank is an integer and worker_ports is None.
|
||||
"""
|
||||
nnodes = server_args.nnodes
|
||||
batch_start_idx = len(self.scheduler_actors)
|
||||
|
||||
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
|
||||
)
|
||||
|
||||
for pp_rank in pp_range:
|
||||
for tp_rank in tp_range:
|
||||
rank_port_args = port_args
|
||||
actual_dp_rank = dp_rank
|
||||
|
||||
if server_args.enable_dp_attention:
|
||||
# DP attention: derive dp_rank from tp_rank
|
||||
_, _, actual_dp_rank = compute_dp_attention_world_info(
|
||||
server_args.enable_dp_attention,
|
||||
tp_rank,
|
||||
server_args.tp_size,
|
||||
server_args.dp_size,
|
||||
server_args.attn_cp_size,
|
||||
)
|
||||
rank_port_args = PortArgs.init_new(
|
||||
server_args, actual_dp_rank, worker_ports
|
||||
)
|
||||
# All DP ranks share the same NCCL port (reuse TP group)
|
||||
rank_port_args.nccl_port = port_args.nccl_port
|
||||
# The detokenizer and tokenizer bind using the
|
||||
# original port_args addresses (127.0.0.1 when
|
||||
# dist_init_addr is unset). Scheduler actors must
|
||||
# connect to the same addresses.
|
||||
rank_port_args.detokenizer_ipc_name = (
|
||||
port_args.detokenizer_ipc_name
|
||||
)
|
||||
rank_port_args.tokenizer_ipc_name = (
|
||||
port_args.tokenizer_ipc_name
|
||||
)
|
||||
|
||||
local_gpu_idx = (pp_rank % pp_per_node) * tp_per_node + (
|
||||
tp_rank % tp_per_node
|
||||
)
|
||||
|
||||
attn_cp_rank, moe_dp_rank, moe_ep_rank = (
|
||||
_compute_parallelism_ranks(server_args, tp_rank)
|
||||
)
|
||||
|
||||
# Each DP group needs a unique dist_init_addr for its own
|
||||
# torch.distributed process group. Use nccl_port which is
|
||||
# unique per DP group (regular DP) or shared (DP attention).
|
||||
dist_init_addr = (
|
||||
f"{self.rank0_node_ip}:{rank_port_args.nccl_port}"
|
||||
)
|
||||
|
||||
actor = SchedulerActor.options(
|
||||
num_cpus=0,
|
||||
num_gpus=1,
|
||||
name=(
|
||||
f"sglang_scheduler_rank0node={self.rank0_node_ip}"
|
||||
f"_dp{actual_dp_rank}_pp{pp_rank}_tp{tp_rank}"
|
||||
),
|
||||
scheduling_strategy=PlacementGroupSchedulingStrategy(
|
||||
placement_group=self.pg,
|
||||
placement_group_bundle_index=bundle_idx,
|
||||
),
|
||||
).remote(
|
||||
server_args=server_args,
|
||||
port_args=rank_port_args,
|
||||
gpu_id=local_gpu_idx,
|
||||
tp_rank=tp_rank,
|
||||
attn_cp_rank=attn_cp_rank,
|
||||
moe_dp_rank=moe_dp_rank,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
pp_rank=pp_rank,
|
||||
dp_rank=actual_dp_rank,
|
||||
dist_init_addr=dist_init_addr,
|
||||
)
|
||||
self.scheduler_actors.append(actor)
|
||||
|
||||
# Wait for all actors created in this call to initialize
|
||||
batch_actors = self.scheduler_actors[batch_start_idx:]
|
||||
try:
|
||||
scheduler_infos = ray.get(
|
||||
[actor.get_info.remote() for actor in batch_actors]
|
||||
)
|
||||
except ray.exceptions.RayActorError as e:
|
||||
for actor in self.scheduler_actors:
|
||||
try:
|
||||
ray.kill(actor)
|
||||
except Exception:
|
||||
logger.error(f"Failed to kill Ray scheduler actor: {actor}")
|
||||
raise RuntimeError(f"Scheduler actor failed to initialize: {e}")
|
||||
|
||||
# Store init info from the first actor (same across all actors)
|
||||
if scheduler_infos:
|
||||
self.max_total_num_tokens = scheduler_infos[0]["max_total_num_tokens"]
|
||||
self.max_req_input_len = scheduler_infos[0]["max_req_input_len"]
|
||||
|
||||
# Start event loops (non-blocking — runs until actor is killed)
|
||||
self.event_loop_refs.extend(
|
||||
[actor.run_event_loop.remote() for actor in batch_actors]
|
||||
)
|
||||
|
||||
# Override launch_tensor_parallel_group to be a no-op since we don't use it.
|
||||
# The parent's launch_dp_schedulers/launch_dp_attention_schedulers call this,
|
||||
# but our overrides call _launch_ray_tp_group instead.
|
||||
def launch_tensor_parallel_group(
|
||||
self,
|
||||
server_args: ServerArgs,
|
||||
port_args: PortArgs,
|
||||
base_gpu_id: int,
|
||||
dp_rank: Optional[int],
|
||||
worker_ports: Optional[List[int]] = None,
|
||||
):
|
||||
raise RuntimeError(
|
||||
"RayDataParallelController should not call launch_tensor_parallel_group. "
|
||||
"Use _launch_ray_tp_group instead."
|
||||
)
|
||||
+158
-69
@@ -17,6 +17,7 @@ from __future__ import annotations
|
||||
|
||||
import dataclasses
|
||||
import logging
|
||||
import threading
|
||||
from typing import Callable
|
||||
|
||||
import ray
|
||||
@@ -97,12 +98,6 @@ class RayEngine(Engine):
|
||||
Tuple of (RaySchedulerInitResult, None).
|
||||
scheduler_procs is None since Ray uses actors instead of mp.Process.
|
||||
"""
|
||||
if server_args.dp_size > 1:
|
||||
raise NotImplementedError(
|
||||
"Ray support for dp_size > 1 is not yet implemented. "
|
||||
"Set dp_size=1 or use_ray=False."
|
||||
)
|
||||
|
||||
pg = ray.util.get_current_placement_group()
|
||||
if pg is None:
|
||||
raise RuntimeError(
|
||||
@@ -110,77 +105,174 @@ class RayEngine(Engine):
|
||||
"Schedule the Engine actor onto a placement group"
|
||||
)
|
||||
|
||||
world_size = server_args.tp_size * server_args.pp_size
|
||||
nnodes = server_args.nnodes
|
||||
gpus_per_node = world_size // nnodes
|
||||
|
||||
logger.info(
|
||||
f"Ray cluster: {nnodes} nodes, "
|
||||
f"Use {gpus_per_node} GPUs/node, world_size={world_size}"
|
||||
)
|
||||
|
||||
# co-located with the Engine and rank0 scheduler at the same node
|
||||
engine_bundle, engine_ip = _find_engine_bundle(pg, nnodes)
|
||||
bundle_for_node = [engine_bundle] + [
|
||||
i for i in range(nnodes) if i != engine_bundle
|
||||
]
|
||||
|
||||
rank0_node_ip = engine_ip
|
||||
dist_init_addr = f"{rank0_node_ip}:{server_args.port + ZMQ_TCP_PORT_DELTA}"
|
||||
logger.info(f"dist_init_addr: {dist_init_addr}")
|
||||
|
||||
scheduler_actors = []
|
||||
if server_args.dp_size == 1:
|
||||
# Launch tensor parallel scheduler actors
|
||||
world_size = server_args.tp_size * server_args.pp_size
|
||||
gpus_per_node = world_size // nnodes
|
||||
|
||||
for node_idx in range(nnodes):
|
||||
bundle_idx = 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
|
||||
logger.info(
|
||||
f"Ray cluster: {nnodes} nodes, "
|
||||
f"Use {gpus_per_node} GPUs/node, world_size={world_size}"
|
||||
)
|
||||
for pp_rank in pp_range:
|
||||
for tp_rank in tp_range:
|
||||
local_gpu_idx = (pp_rank % pp_per_node) * tp_per_node + (
|
||||
tp_rank % tp_per_node
|
||||
)
|
||||
|
||||
attn_cp_rank, moe_dp_rank, moe_ep_rank = _compute_parallelism_ranks(
|
||||
server_args, tp_rank
|
||||
)
|
||||
|
||||
actor = SchedulerActor.options(
|
||||
num_cpus=0,
|
||||
num_gpus=1,
|
||||
name=f"sglang_scheduler_rank0node={rank0_node_ip}_pp{pp_rank}_tp{tp_rank}",
|
||||
scheduling_strategy=PlacementGroupSchedulingStrategy(
|
||||
placement_group=pg,
|
||||
placement_group_bundle_index=bundle_idx,
|
||||
),
|
||||
).remote(
|
||||
server_args=server_args,
|
||||
port_args=port_args,
|
||||
gpu_id=local_gpu_idx,
|
||||
tp_rank=tp_rank,
|
||||
attn_cp_rank=attn_cp_rank,
|
||||
moe_dp_rank=moe_dp_rank,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
pp_rank=pp_rank,
|
||||
dp_rank=0,
|
||||
dist_init_addr=dist_init_addr,
|
||||
)
|
||||
scheduler_actors.append(actor)
|
||||
|
||||
try:
|
||||
scheduler_infos = ray.get(
|
||||
[actor.get_info.remote() for actor in scheduler_actors]
|
||||
dist_init_addr = (
|
||||
f"{rank0_node_ip}:{server_args.port + ZMQ_TCP_PORT_DELTA}"
|
||||
)
|
||||
except ray.exceptions.RayActorError as e:
|
||||
for actor in scheduler_actors:
|
||||
logger.info(f"dist_init_addr: {dist_init_addr}")
|
||||
|
||||
scheduler_actors = []
|
||||
|
||||
for node_idx in range(nnodes):
|
||||
bundle_idx = 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,
|
||||
)
|
||||
)
|
||||
for pp_rank in pp_range:
|
||||
for tp_rank in tp_range:
|
||||
local_gpu_idx = (
|
||||
pp_rank % pp_per_node
|
||||
) * tp_per_node + (tp_rank % tp_per_node)
|
||||
|
||||
attn_cp_rank, moe_dp_rank, moe_ep_rank = (
|
||||
_compute_parallelism_ranks(server_args, tp_rank)
|
||||
)
|
||||
|
||||
actor = SchedulerActor.options(
|
||||
num_cpus=0,
|
||||
num_gpus=1,
|
||||
name=f"sglang_scheduler_rank0node={rank0_node_ip}_pp{pp_rank}_tp{tp_rank}",
|
||||
scheduling_strategy=PlacementGroupSchedulingStrategy(
|
||||
placement_group=pg,
|
||||
placement_group_bundle_index=bundle_idx,
|
||||
),
|
||||
).remote(
|
||||
server_args=server_args,
|
||||
port_args=port_args,
|
||||
gpu_id=local_gpu_idx,
|
||||
tp_rank=tp_rank,
|
||||
attn_cp_rank=attn_cp_rank,
|
||||
moe_dp_rank=moe_dp_rank,
|
||||
moe_ep_rank=moe_ep_rank,
|
||||
pp_rank=pp_rank,
|
||||
dp_rank=0,
|
||||
dist_init_addr=dist_init_addr,
|
||||
)
|
||||
scheduler_actors.append(actor)
|
||||
|
||||
try:
|
||||
scheduler_infos = ray.get(
|
||||
[actor.get_info.remote() for actor in scheduler_actors]
|
||||
)
|
||||
except ray.exceptions.RayActorError as e:
|
||||
for actor in scheduler_actors:
|
||||
try:
|
||||
ray.kill(actor)
|
||||
except Exception:
|
||||
logger.error(
|
||||
f"Failed to kill Ray scheduler actor: {actor}"
|
||||
)
|
||||
raise RuntimeError(
|
||||
f"Scheduler actor failed to initialize: {e}"
|
||||
)
|
||||
|
||||
event_loop_refs = [
|
||||
actor.run_event_loop.remote() for actor in scheduler_actors
|
||||
]
|
||||
|
||||
def wait_for_completion():
|
||||
try:
|
||||
ray.kill(actor)
|
||||
except Exception:
|
||||
logger.error(f"Failed to kill Ray scheduler actor: {actor}")
|
||||
raise RuntimeError(f"Scheduler actor failed to initialize: {e}")
|
||||
ray.get(event_loop_refs)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"Ray scheduler actor terminated with error: {e}"
|
||||
)
|
||||
|
||||
event_loop_refs = [actor.run_event_loop.remote() for actor in scheduler_actors]
|
||||
return (
|
||||
RaySchedulerInitResult(
|
||||
scheduler_infos=scheduler_infos,
|
||||
wait_for_completion=wait_for_completion,
|
||||
scheduler_actors=scheduler_actors,
|
||||
),
|
||||
None,
|
||||
)
|
||||
else:
|
||||
# Launch the data parallel controller
|
||||
return (
|
||||
cls._launch_dp_scheduler_processes(
|
||||
server_args, port_args, pg, bundle_for_node, rank0_node_ip
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _launch_dp_scheduler_processes(
|
||||
cls,
|
||||
server_args: ServerArgs,
|
||||
port_args: PortArgs,
|
||||
pg,
|
||||
bundle_for_node: list,
|
||||
rank0_node_ip: str,
|
||||
) -> RaySchedulerInitResult:
|
||||
"""Launch DP schedulers via RayDataParallelController."""
|
||||
from sglang.srt.ray.data_parallel_controller import (
|
||||
RayDataParallelController,
|
||||
)
|
||||
|
||||
if server_args.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
|
||||
else:
|
||||
total_gpus = server_args.dp_size * server_args.tp_size * server_args.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={server_args.dp_size}, "
|
||||
f"tp_size={server_args.tp_size}, pp_size={server_args.pp_size}, "
|
||||
f"enable_dp_attention={server_args.enable_dp_attention}"
|
||||
)
|
||||
|
||||
# Set dist_init_addr on server_args so PortArgs.init_new() can compute
|
||||
# TCP addresses correctly (required for DP attention path).
|
||||
dp_server_args = dataclasses.replace(
|
||||
server_args,
|
||||
dist_init_addr=f"{rank0_node_ip}:{server_args.port + ZMQ_TCP_PORT_DELTA}",
|
||||
)
|
||||
|
||||
# Create the DP controller in-process. This blocks until all actors
|
||||
# are initialized and their event loops have started.
|
||||
controller = RayDataParallelController(
|
||||
dp_server_args, port_args, pg, bundle_for_node, rank0_node_ip
|
||||
)
|
||||
|
||||
# Start the DP controller's event loop in a daemon thread.
|
||||
# It routes requests from the tokenizer to per-DP-rank schedulers.
|
||||
dp_thread = threading.Thread(
|
||||
target=controller.event_loop, daemon=True, name="dp_controller"
|
||||
)
|
||||
dp_thread.start()
|
||||
|
||||
scheduler_infos = [
|
||||
{
|
||||
"max_total_num_tokens": controller.max_total_num_tokens,
|
||||
"max_req_input_len": controller.max_req_input_len,
|
||||
}
|
||||
]
|
||||
|
||||
event_loop_refs = controller.event_loop_refs
|
||||
|
||||
def wait_for_completion():
|
||||
try:
|
||||
@@ -188,11 +280,8 @@ class RayEngine(Engine):
|
||||
except Exception as e:
|
||||
logger.error(f"Ray scheduler actor terminated with error: {e}")
|
||||
|
||||
return (
|
||||
RaySchedulerInitResult(
|
||||
scheduler_infos=scheduler_infos,
|
||||
wait_for_completion=wait_for_completion,
|
||||
scheduler_actors=scheduler_actors,
|
||||
),
|
||||
None,
|
||||
return RaySchedulerInitResult(
|
||||
scheduler_infos=scheduler_infos,
|
||||
wait_for_completion=wait_for_completion,
|
||||
scheduler_actors=controller.scheduler_actors,
|
||||
)
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
|
||||
Tests the Ray actor scheduler backend:
|
||||
- Offline inference via Engine(use_ray=True) inside a Ray actor on a placement group
|
||||
- Data parallel (DP) and DP attention support
|
||||
- Error paths in RayEngine._launch_scheduler_processes()
|
||||
- HTTP server launched via --use-ray flag
|
||||
|
||||
@@ -14,6 +15,8 @@ Usage:
|
||||
# 2-GPU tests
|
||||
python -m pytest test/manual/test_ray_engine.py::TestRayEngineOfflineTP2 -v -s
|
||||
python -m pytest test/manual/test_ray_engine.py::TestRayEngineOfflinePP2 -v -s
|
||||
python -m pytest test/manual/test_ray_engine.py::TestRayEngineOfflineDP2 -v -s
|
||||
python -m pytest test/manual/test_ray_engine.py::TestRayEngineOfflineDPAttention -v -s
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -29,6 +32,11 @@ from sglang.test.test_utils import DEFAULT_SMALL_MODEL_NAME_FOR_TEST
|
||||
# Allow overriding the model via env var for environments without gated access
|
||||
_MODEL = os.environ.get("SGLANG_TEST_MODEL", DEFAULT_SMALL_MODEL_NAME_FOR_TEST)
|
||||
|
||||
# DP attention requires a model whose num_kv_heads divides evenly across the
|
||||
# attention-TP dimension. Qwen2.5-0.5B (kv_heads=2, attn_heads=14) hits a
|
||||
# shape mismatch in the KV cache, so we use a larger model here.
|
||||
_DP_ATTN_MODEL = os.environ.get("SGLANG_TEST_DP_ATTN_MODEL", "Qwen/Qwen3-8B")
|
||||
|
||||
try:
|
||||
import ray
|
||||
from ray.runtime_env import RuntimeEnv
|
||||
@@ -64,7 +72,9 @@ _PROMPTS = [
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _create_engine_on_pg(tp_size, pp_size=1, model=_MODEL, extra_kwargs=None):
|
||||
def _create_engine_on_pg(
|
||||
tp_size, pp_size=1, dp_size=1, model=_MODEL, extra_kwargs=None
|
||||
):
|
||||
"""Create an EngineActor on a placement group and wait for it to be ready.
|
||||
|
||||
Returns (engine_actor, placement_group).
|
||||
@@ -88,7 +98,12 @@ def _create_engine_on_pg(tp_size, pp_size=1, model=_MODEL, extra_kwargs=None):
|
||||
self.engine.shutdown()
|
||||
self.engine = None
|
||||
|
||||
total_gpus = tp_size * pp_size
|
||||
enable_dp_attention = (extra_kwargs or {}).get("enable_dp_attention", False)
|
||||
if enable_dp_attention:
|
||||
# DP attention folds DP into TP — total GPUs = tp_size * pp_size
|
||||
total_gpus = tp_size * pp_size
|
||||
else:
|
||||
total_gpus = dp_size * tp_size * pp_size
|
||||
pg = placement_group(
|
||||
[{"CPU": 1, "GPU": total_gpus}],
|
||||
strategy="STRICT_PACK",
|
||||
@@ -99,6 +114,7 @@ def _create_engine_on_pg(tp_size, pp_size=1, model=_MODEL, extra_kwargs=None):
|
||||
model_path=model,
|
||||
tp_size=tp_size,
|
||||
pp_size=pp_size,
|
||||
dp_size=dp_size,
|
||||
)
|
||||
if extra_kwargs:
|
||||
kwargs.update(extra_kwargs)
|
||||
@@ -244,6 +260,87 @@ class TestRayEngineOfflinePP2(unittest.TestCase):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@unittest.skipUnless(_has_ray, "ray is not installed")
|
||||
@unittest.skipUnless(_NUM_GPUS >= 2, "requires at least 2 GPUs")
|
||||
class TestRayEngineOfflineDP2(unittest.TestCase):
|
||||
"""Test Ray engine with dp_size=2, tp_size=1."""
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
if not ray.is_initialized():
|
||||
ray.init(log_to_driver=True, runtime_env=_RAY_RUNTIME_ENV)
|
||||
cls.actor, cls.pg = _create_engine_on_pg(tp_size=1, dp_size=2)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
_cleanup(cls.actor, cls.pg)
|
||||
ray.shutdown()
|
||||
|
||||
def test_offline_generate_dp2(self):
|
||||
result = ray.get(
|
||||
self.actor.generate.remote("The capital of France is", _SAMPLING_PARAMS)
|
||||
)
|
||||
self.assertIn("text", result)
|
||||
self.assertGreater(len(result["text"]), 0)
|
||||
print(f"Generated (DP=2): {result['text'][:200]}")
|
||||
|
||||
def test_batch_generate_dp2(self):
|
||||
for prompt in _PROMPTS:
|
||||
result = ray.get(self.actor.generate.remote(prompt, _SAMPLING_PARAMS))
|
||||
self.assertIn("text", result)
|
||||
self.assertGreater(len(result["text"]), 0, f"Empty output for: {prompt}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests: Offline DP Attention (dp=2, tp=2)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@unittest.skipUnless(_has_ray, "ray is not installed")
|
||||
@unittest.skipUnless(_NUM_GPUS >= 2, "requires at least 2 GPUs")
|
||||
class TestRayEngineOfflineDPAttention(unittest.TestCase):
|
||||
"""Test Ray engine with dp_size=2, tp_size=2, enable_dp_attention=True."""
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
if not ray.is_initialized():
|
||||
ray.init(log_to_driver=True, runtime_env=_RAY_RUNTIME_ENV)
|
||||
cls.actor, cls.pg = _create_engine_on_pg(
|
||||
tp_size=2,
|
||||
dp_size=2,
|
||||
model=_DP_ATTN_MODEL,
|
||||
extra_kwargs={
|
||||
"enable_dp_attention": True,
|
||||
"disable_cuda_graph": True,
|
||||
"port": 31500,
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def tearDownClass(cls):
|
||||
_cleanup(cls.actor, cls.pg)
|
||||
ray.shutdown()
|
||||
|
||||
def test_offline_generate_dp_attention(self):
|
||||
result = ray.get(
|
||||
self.actor.generate.remote("The capital of France is", _SAMPLING_PARAMS)
|
||||
)
|
||||
self.assertIn("text", result)
|
||||
self.assertGreater(len(result["text"]), 0)
|
||||
print(f"Generated (DP-Attention): {result['text'][:200]}")
|
||||
|
||||
def test_batch_generate_dp_attention(self):
|
||||
for prompt in _PROMPTS:
|
||||
result = ray.get(self.actor.generate.remote(prompt, _SAMPLING_PARAMS))
|
||||
self.assertIn("text", result)
|
||||
self.assertGreater(len(result["text"]), 0, f"Empty output for: {prompt}")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests: Error paths
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@unittest.skipUnless(_has_ray, "ray is not installed")
|
||||
@unittest.skipUnless(_NUM_GPUS >= 1, "requires at least 1 GPU")
|
||||
class TestRayEngineErrors(unittest.TestCase):
|
||||
@@ -257,44 +354,6 @@ class TestRayEngineErrors(unittest.TestCase):
|
||||
def tearDownClass(cls):
|
||||
ray.shutdown()
|
||||
|
||||
def test_dp_greater_than_1_raises(self):
|
||||
"""RayEngine with dp_size > 1 should raise NotImplementedError."""
|
||||
|
||||
@ray.remote
|
||||
class _BadActor:
|
||||
def try_create(self):
|
||||
from sglang.srt.ray.engine import RayEngine
|
||||
|
||||
try:
|
||||
RayEngine(
|
||||
model_path=_MODEL,
|
||||
tp_size=1,
|
||||
dp_size=2,
|
||||
use_ray=True,
|
||||
)
|
||||
return None
|
||||
except (NotImplementedError, RuntimeError) as e:
|
||||
return str(e)
|
||||
|
||||
pg = placement_group([{"CPU": 1, "GPU": 1}], strategy="STRICT_PACK")
|
||||
ray.get(pg.ready())
|
||||
|
||||
actor = _BadActor.options(
|
||||
num_cpus=1,
|
||||
num_gpus=0,
|
||||
scheduling_strategy=PlacementGroupSchedulingStrategy(
|
||||
placement_group=pg,
|
||||
placement_group_bundle_index=0,
|
||||
),
|
||||
).remote()
|
||||
|
||||
try:
|
||||
error_msg = ray.get(actor.try_create.remote(), timeout=120)
|
||||
self.assertIsNotNone(error_msg, "Expected error but RayEngine created OK")
|
||||
self.assertIn("dp_size", error_msg.lower())
|
||||
finally:
|
||||
ray.util.remove_placement_group(pg)
|
||||
|
||||
def test_missing_placement_group_raises(self):
|
||||
"""RayEngine without a placement group should raise RuntimeError."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user