Fix data parallel controller launch for num nodes > 2 (#12822)
This commit is contained in:
@@ -788,24 +788,26 @@ def _launch_subprocesses(
|
|||||||
|
|
||||||
scheduler_procs = []
|
scheduler_procs = []
|
||||||
if server_args.dp_size == 1:
|
if server_args.dp_size == 1:
|
||||||
|
# Launch tensor parallel scheduler processes
|
||||||
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
memory_saver_adapter = TorchMemorySaverAdapter.create(
|
||||||
enable=server_args.enable_memory_saver
|
enable=server_args.enable_memory_saver
|
||||||
)
|
)
|
||||||
scheduler_pipe_readers = []
|
scheduler_pipe_readers = []
|
||||||
|
|
||||||
nnodes_per_tp_group = max(server_args.nnodes // server_args.pp_size, 1)
|
pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1)
|
||||||
|
nnodes_per_pp_rank = max(server_args.nnodes // server_args.pp_size, 1)
|
||||||
|
pp_rank_range = range(
|
||||||
|
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank),
|
||||||
|
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1),
|
||||||
|
)
|
||||||
|
|
||||||
|
nnodes_per_tp_group = nnodes_per_pp_rank
|
||||||
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
||||||
tp_rank_range = range(
|
tp_rank_range = range(
|
||||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group),
|
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group),
|
||||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
||||||
)
|
)
|
||||||
|
|
||||||
pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1)
|
|
||||||
pp_rank_range = range(
|
|
||||||
pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group),
|
|
||||||
pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group + 1),
|
|
||||||
)
|
|
||||||
|
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
for tp_rank in tp_rank_range:
|
for tp_rank in tp_rank_range:
|
||||||
reader, writer = mp.Pipe(duplex=False)
|
reader, writer = mp.Pipe(duplex=False)
|
||||||
|
|||||||
@@ -49,8 +49,9 @@ from sglang.srt.tracing.trace import (
|
|||||||
trace_slice_end,
|
trace_slice_end,
|
||||||
trace_slice_start,
|
trace_slice_start,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils.common import (
|
||||||
bind_port,
|
bind_port,
|
||||||
|
configure_ipv6,
|
||||||
configure_logger,
|
configure_logger,
|
||||||
get_zmq_socket,
|
get_zmq_socket,
|
||||||
kill_itself_when_parent_died,
|
kill_itself_when_parent_died,
|
||||||
@@ -118,16 +119,16 @@ class DataParallelController:
|
|||||||
"""A controller that dispatches requests to multiple data parallel workers."""
|
"""A controller that dispatches requests to multiple data parallel workers."""
|
||||||
|
|
||||||
def __init__(self, server_args: ServerArgs, port_args: PortArgs) -> None:
|
def __init__(self, server_args: ServerArgs, port_args: PortArgs) -> None:
|
||||||
# for dp balance
|
|
||||||
self.global_balance_id = 0
|
|
||||||
|
|
||||||
# Parse args
|
# Parse args
|
||||||
self.max_total_num_tokens = None
|
|
||||||
self.server_args = server_args
|
self.server_args = server_args
|
||||||
self.port_args = port_args
|
self.port_args = port_args
|
||||||
self.load_balance_method = LoadBalanceMethod.from_str(
|
self.load_balance_method = LoadBalanceMethod.from_str(
|
||||||
server_args.load_balance_method
|
server_args.load_balance_method
|
||||||
)
|
)
|
||||||
|
self.run_scheduler_process = run_scheduler_process
|
||||||
|
|
||||||
|
# For DP balance
|
||||||
|
self.global_balance_id = 0
|
||||||
|
|
||||||
# Init inter-process communication
|
# Init inter-process communication
|
||||||
self.context = zmq.Context(1 + server_args.dp_size)
|
self.context = zmq.Context(1 + server_args.dp_size)
|
||||||
@@ -162,8 +163,6 @@ class DataParallelController:
|
|||||||
self.launch_dp_schedulers(server_args, port_args)
|
self.launch_dp_schedulers(server_args, port_args)
|
||||||
self.control_message_step = 1
|
self.control_message_step = 1
|
||||||
|
|
||||||
self.max_req_input_len = None
|
|
||||||
|
|
||||||
self.init_dispatcher()
|
self.init_dispatcher()
|
||||||
|
|
||||||
def send_to_all_workers(self, obj):
|
def send_to_all_workers(self, obj):
|
||||||
@@ -281,8 +280,12 @@ class DataParallelController:
|
|||||||
# Determine the endpoint for inter-node communication
|
# Determine the endpoint for inter-node communication
|
||||||
if server_args.dist_init_addr is None:
|
if server_args.dist_init_addr is None:
|
||||||
endpoint = f"tcp://127.0.0.1:{server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA}"
|
endpoint = f"tcp://127.0.0.1:{server_args.port + DP_ATTENTION_HANDSHAKE_PORT_DELTA}"
|
||||||
|
elif server_args.dist_init_addr.startswith("["): # ipv6 address
|
||||||
|
port, host = configure_ipv6(server_args.dist_init_addr)
|
||||||
|
endpoint = f"tcp://{host}:{int(port) + DP_ATTENTION_HANDSHAKE_PORT_DELTA}"
|
||||||
else:
|
else:
|
||||||
endpoint = f"tcp://{server_args.dist_init_addr}"
|
host, port = server_args.dist_init_addr.split(":")
|
||||||
|
endpoint = f"tcp://{host}:{int(port) + DP_ATTENTION_HANDSHAKE_PORT_DELTA}"
|
||||||
|
|
||||||
if server_args.node_rank == 0:
|
if server_args.node_rank == 0:
|
||||||
# Node 0: Broadcast worker ports to all other nodes
|
# Node 0: Broadcast worker ports to all other nodes
|
||||||
@@ -326,8 +329,8 @@ class DataParallelController:
|
|||||||
logger.debug(f"Connecting to node 0 to receive worker ports")
|
logger.debug(f"Connecting to node 0 to receive worker ports")
|
||||||
|
|
||||||
req_socket = get_zmq_socket(self.context, zmq.REQ, endpoint, False)
|
req_socket = get_zmq_socket(self.context, zmq.REQ, endpoint, False)
|
||||||
req_socket.setsockopt(zmq.RCVTIMEO, 60 * 1000) # 1 minute timeout
|
req_socket.setsockopt(zmq.RCVTIMEO, 600 * 1000) # 10 minute timeout
|
||||||
req_socket.setsockopt(zmq.SNDTIMEO, 60 * 1000)
|
req_socket.setsockopt(zmq.SNDTIMEO, 600 * 1000)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Send handshake with our node rank
|
# Send handshake with our node rank
|
||||||
@@ -381,19 +384,20 @@ class DataParallelController:
|
|||||||
|
|
||||||
scheduler_pipe_readers = []
|
scheduler_pipe_readers = []
|
||||||
|
|
||||||
nnodes_per_tp_group = max(server_args.nnodes // server_args.pp_size, 1)
|
pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1)
|
||||||
|
nnodes_per_pp_rank = max(server_args.nnodes // server_args.pp_size, 1)
|
||||||
|
pp_rank_range = range(
|
||||||
|
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank),
|
||||||
|
pp_size_per_node * (server_args.node_rank // nnodes_per_pp_rank + 1),
|
||||||
|
)
|
||||||
|
|
||||||
|
nnodes_per_tp_group = nnodes_per_pp_rank
|
||||||
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
tp_size_per_node = server_args.tp_size // nnodes_per_tp_group
|
||||||
tp_rank_range = range(
|
tp_rank_range = range(
|
||||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group),
|
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group),
|
||||||
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
tp_size_per_node * (server_args.node_rank % nnodes_per_tp_group + 1),
|
||||||
)
|
)
|
||||||
|
|
||||||
pp_size_per_node = max(server_args.pp_size // server_args.nnodes, 1)
|
|
||||||
pp_rank_range = range(
|
|
||||||
pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group),
|
|
||||||
pp_size_per_node * (server_args.node_rank // nnodes_per_tp_group + 1),
|
|
||||||
)
|
|
||||||
|
|
||||||
for pp_rank in pp_rank_range:
|
for pp_rank in pp_rank_range:
|
||||||
for tp_rank in tp_rank_range:
|
for tp_rank in tp_rank_range:
|
||||||
rank_port_args = port_args
|
rank_port_args = port_args
|
||||||
@@ -424,7 +428,7 @@ class DataParallelController:
|
|||||||
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
|
moe_ep_rank = tp_rank // (server_args.tp_size // server_args.ep_size)
|
||||||
with self.env_lock, maybe_reindex_device_id(gpu_id) as gpu_id:
|
with self.env_lock, maybe_reindex_device_id(gpu_id) as gpu_id:
|
||||||
proc = mp.Process(
|
proc = mp.Process(
|
||||||
target=run_scheduler_process,
|
target=self.run_scheduler_process,
|
||||||
args=(
|
args=(
|
||||||
server_args,
|
server_args,
|
||||||
rank_port_args,
|
rank_port_args,
|
||||||
@@ -504,8 +508,14 @@ def run_data_parallel_controller_process(
|
|||||||
server_args: ServerArgs,
|
server_args: ServerArgs,
|
||||||
port_args: PortArgs,
|
port_args: PortArgs,
|
||||||
pipe_writer,
|
pipe_writer,
|
||||||
|
data_parallel_controller_class=DataParallelController,
|
||||||
):
|
):
|
||||||
|
setproctitle.setproctitle("sglang::data_parallel_controller")
|
||||||
|
faulthandler.enable()
|
||||||
kill_itself_when_parent_died()
|
kill_itself_when_parent_died()
|
||||||
|
parent_process = psutil.Process().parent()
|
||||||
|
|
||||||
|
configure_logger(server_args)
|
||||||
if server_args.enable_trace:
|
if server_args.enable_trace:
|
||||||
process_tracing_init(server_args.otlp_traces_endpoint, "sglang")
|
process_tracing_init(server_args.otlp_traces_endpoint, "sglang")
|
||||||
thread_label = "DP Controller"
|
thread_label = "DP Controller"
|
||||||
@@ -514,13 +524,9 @@ def run_data_parallel_controller_process(
|
|||||||
elif server_args.disaggregation_mode == "decode":
|
elif server_args.disaggregation_mode == "decode":
|
||||||
thread_label = "Decode DP Controller"
|
thread_label = "Decode DP Controller"
|
||||||
trace_set_thread_info(thread_label)
|
trace_set_thread_info(thread_label)
|
||||||
setproctitle.setproctitle("sglang::data_parallel_controller")
|
|
||||||
faulthandler.enable()
|
|
||||||
configure_logger(server_args)
|
|
||||||
parent_process = psutil.Process().parent()
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
controller = DataParallelController(server_args, port_args)
|
controller = data_parallel_controller_class(server_args, port_args)
|
||||||
pipe_writer.send(
|
pipe_writer.send(
|
||||||
{
|
{
|
||||||
"status": "ready",
|
"status": "ready",
|
||||||
|
|||||||
@@ -2762,12 +2762,12 @@ def run_scheduler_process(
|
|||||||
dp_rank = int(os.environ["SGLANG_DP_RANK"])
|
dp_rank = int(os.environ["SGLANG_DP_RANK"])
|
||||||
if dp_rank is not None:
|
if dp_rank is not None:
|
||||||
prefix += f" DP{dp_rank}"
|
prefix += f" DP{dp_rank}"
|
||||||
|
if server_args.pp_size > 1:
|
||||||
|
prefix += f" PP{pp_rank}"
|
||||||
if server_args.tp_size > 1:
|
if server_args.tp_size > 1:
|
||||||
prefix += f" TP{tp_rank}"
|
prefix += f" TP{tp_rank}"
|
||||||
if server_args.ep_size > 1:
|
if server_args.ep_size > 1:
|
||||||
prefix += f" EP{moe_ep_rank}"
|
prefix += f" EP{moe_ep_rank}"
|
||||||
if server_args.pp_size > 1:
|
|
||||||
prefix += f" PP{pp_rank}"
|
|
||||||
|
|
||||||
# Config the process
|
# Config the process
|
||||||
setproctitle.setproctitle(f"sglang::scheduler{prefix.replace(' ', '_')}")
|
setproctitle.setproctitle(f"sglang::scheduler{prefix.replace(' ', '_')}")
|
||||||
|
|||||||
@@ -4026,7 +4026,7 @@ def prepare_server_args(argv: List[str]) -> ServerArgs:
|
|||||||
|
|
||||||
|
|
||||||
ZMQ_TCP_PORT_DELTA = 233
|
ZMQ_TCP_PORT_DELTA = 233
|
||||||
DP_ATTENTION_HANDSHAKE_PORT_DELTA = 5
|
DP_ATTENTION_HANDSHAKE_PORT_DELTA = 13
|
||||||
|
|
||||||
|
|
||||||
@dataclasses.dataclass
|
@dataclasses.dataclass
|
||||||
|
|||||||
Reference in New Issue
Block a user