[feat][Ascend][Mindspore]: support model-impl of mindspore (#9234)
This commit is contained in:
118
python/sglang/srt/model_executor/mindspore_runner.py
Normal file
118
python/sglang/srt/model_executor/mindspore_runner.py
Normal file
@@ -0,0 +1,118 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the SGLang project
|
||||
"""ms_runner launch MindSpore distributed modules."""
|
||||
|
||||
import multiprocessing as mp
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import mindspore as ms
|
||||
import torch
|
||||
from mindspore._c_expression import GroupOptions
|
||||
from mindspore.communication import create_group
|
||||
|
||||
from sglang.srt.distributed.parallel_state import _groups
|
||||
|
||||
|
||||
class _Tmp:
|
||||
def __init__(self):
|
||||
self.sched_p = None
|
||||
|
||||
def set_sched_process(self, p):
|
||||
self.sched_p = p
|
||||
|
||||
def __del__(self):
|
||||
if self.sched_p:
|
||||
self.sched_p.kill()
|
||||
|
||||
|
||||
_tmp = _Tmp()
|
||||
|
||||
|
||||
def _get_host_and_ip(distributed_init_method):
|
||||
try:
|
||||
_, ip_str, port_str = distributed_init_method.split(":")
|
||||
ip = ip_str.split("/")[-1]
|
||||
port = int(port_str)
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
"Cannot get host and port information from %s, error: %s!"
|
||||
% (distributed_init_method, str(e))
|
||||
)
|
||||
|
||||
return ip, port
|
||||
|
||||
|
||||
def run_scheduler_init(rank, local_rank, world_size, master_addr, master_port):
|
||||
with open(str(Path() / "schedule.log"), "w") as scheduler_f:
|
||||
# For Python outputs.
|
||||
sys.stdout = scheduler_f
|
||||
sys.stderr = scheduler_f
|
||||
# For C++ outputs.
|
||||
os.dup2(scheduler_f.fileno(), 1)
|
||||
os.dup2(scheduler_f.fileno(), 2)
|
||||
os.environ["DEVICE_ID"] = str(local_rank)
|
||||
os.environ["MS_WORKER_NUM"] = str(world_size)
|
||||
os.environ["MS_ROLE"] = "MS_SCHED"
|
||||
os.environ["MS_NODE_ID"] = str(rank)
|
||||
os.environ["MS_SCHED_HOST"] = str(master_addr)
|
||||
os.environ["MS_SCHED_PORT"] = str(master_port)
|
||||
# This function is blocked until the whole cluster exits.
|
||||
ms.communication.init()
|
||||
|
||||
|
||||
def set_ms_parallel_env(rank, local_rank, world_size, init_method):
|
||||
master_addr, master_port = _get_host_and_ip(init_method)
|
||||
# change port avoiding port conflicts with torch
|
||||
master_port = master_port + 35 if master_port < 65500 else master_port - 35
|
||||
if not os.getenv("MS_ROLE"):
|
||||
if rank == 0:
|
||||
# Create a subprocess for scheduler of MindSpore, just for internal collaboration, not for collective communication
|
||||
sched_p = mp.Process(
|
||||
target=run_scheduler_init,
|
||||
args=(rank, local_rank, world_size, master_addr, master_port),
|
||||
)
|
||||
sched_p.start()
|
||||
global _tmp
|
||||
_tmp.set_sched_process(sched_p)
|
||||
|
||||
os.environ["DEVICE_ID"] = str(local_rank)
|
||||
os.environ["MS_WORKER_NUM"] = str(world_size)
|
||||
os.environ["MS_ROLE"] = "MS_WORKER"
|
||||
os.environ["MS_NODE_ID"] = str(rank)
|
||||
os.environ["MS_SCHED_HOST"] = str(master_addr)
|
||||
os.environ["MS_SCHED_PORT"] = str(master_port)
|
||||
|
||||
|
||||
def reuse_hccl_comm():
|
||||
for group_name, group in _groups.items():
|
||||
# Torch ProcessGroupHccl
|
||||
device_group = group().device_group
|
||||
hccl_comm_handle = device_group._get_backend(torch.device("npu")).get_hccl_comm(
|
||||
group().local_rank
|
||||
)
|
||||
print(
|
||||
f"MindSpore reuse torch group: {device_group}, group_name: {group_name}, local rank: {group().local_rank},"
|
||||
f"hccl communicator handle: {hex(hccl_comm_handle)}",
|
||||
flush=True,
|
||||
)
|
||||
# Create MS communication group by hccl comm handle to reuse Torch group.
|
||||
group_options = GroupOptions()
|
||||
group_options.hccl_config = {"hccl_comm": hccl_comm_handle}
|
||||
create_group(group_name, group().ranks, group_options)
|
||||
|
||||
|
||||
def init_ms_distributed(world_size, rank, local_rank, server_args, port):
|
||||
if server_args.dist_init_addr:
|
||||
dist_init_method = f"tcp://{server_args.dist_init_addr}"
|
||||
else:
|
||||
dist_init_method = f"tcp://{server_args.host}:{port}"
|
||||
set_ms_parallel_env(rank, local_rank, world_size, dist_init_method)
|
||||
|
||||
ms.set_context(infer_boost="on", jit_level="O0")
|
||||
ms.set_context(mode=ms.context.PYNATIVE_MODE)
|
||||
ms.set_device("Ascend", local_rank)
|
||||
ms.communication.init("hccl")
|
||||
# After distributed job is initialized, reuse hccl comms for MindSpore.
|
||||
reuse_hccl_comm()
|
||||
@@ -42,6 +42,7 @@ from sglang.srt.configs.load_config import LoadConfig, LoadFormat
|
||||
from sglang.srt.configs.model_config import (
|
||||
AttentionArch,
|
||||
ModelConfig,
|
||||
ModelImpl,
|
||||
get_nsa_index_head_dim,
|
||||
is_deepseek_nsa,
|
||||
)
|
||||
@@ -317,6 +318,8 @@ class ModelRunner:
|
||||
|
||||
if get_bool_env_var("SGLANG_DETECT_SLOW_RANK"):
|
||||
slow_rank_detector.execute()
|
||||
# Init mindspore running environment when model impl is "mindspore"
|
||||
self.init_mindspore_runner()
|
||||
|
||||
# Update deep gemm configure
|
||||
if deep_gemm_wrapper.ENABLE_JIT_DEEPGEMM:
|
||||
@@ -364,6 +367,20 @@ class ModelRunner:
|
||||
else:
|
||||
self.piecewise_cuda_graph_runner = None
|
||||
|
||||
def init_mindspore_runner(self):
|
||||
# Init the mindspore runner
|
||||
# for now, there is only some communication initialization work
|
||||
if self.server_args.model_impl.lower() == ModelImpl.MINDSPORE and _is_npu:
|
||||
from sglang.srt.model_executor.mindspore_runner import init_ms_distributed
|
||||
|
||||
init_ms_distributed(
|
||||
world_size=self.tp_size * self.pp_size,
|
||||
rank=self.tp_size * self.pp_rank + self.tp_rank,
|
||||
local_rank=self.gpu_id,
|
||||
server_args=self.server_args,
|
||||
port=self.dist_port,
|
||||
)
|
||||
|
||||
def initialize(self, min_per_gpu_memory: float):
|
||||
server_args = self.server_args
|
||||
|
||||
@@ -2018,6 +2035,9 @@ class ModelRunner:
|
||||
# TODO: Currently, cuda graph only captures decode steps, which only exists for generation models
|
||||
return
|
||||
|
||||
if self.server_args.model_impl.lower() == ModelImpl.MINDSPORE:
|
||||
return
|
||||
|
||||
if self.device != "cpu" and self.server_args.disable_cuda_graph:
|
||||
return
|
||||
|
||||
|
||||
Reference in New Issue
Block a user