[NPU]mindspore model support moe (#15363)
This commit is contained in:
@@ -2,6 +2,7 @@
|
|||||||
# SPDX-FileCopyrightText: Copyright contributors to the SGLang project
|
# SPDX-FileCopyrightText: Copyright contributors to the SGLang project
|
||||||
"""ms_runner launch MindSpore distributed modules."""
|
"""ms_runner launch MindSpore distributed modules."""
|
||||||
|
|
||||||
|
import logging
|
||||||
import multiprocessing as mp
|
import multiprocessing as mp
|
||||||
import os
|
import os
|
||||||
import sys
|
import sys
|
||||||
@@ -14,6 +15,8 @@ from mindspore.communication import create_group
|
|||||||
|
|
||||||
from sglang.srt.distributed.parallel_state import _groups
|
from sglang.srt.distributed.parallel_state import _groups
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
class _Tmp:
|
class _Tmp:
|
||||||
def __init__(self):
|
def __init__(self):
|
||||||
@@ -92,10 +95,9 @@ def reuse_hccl_comm():
|
|||||||
hccl_comm_handle = device_group._get_backend(torch.device("npu")).get_hccl_comm(
|
hccl_comm_handle = device_group._get_backend(torch.device("npu")).get_hccl_comm(
|
||||||
group().local_rank
|
group().local_rank
|
||||||
)
|
)
|
||||||
print(
|
logger.info(
|
||||||
f"MindSpore reuse torch group: {device_group}, group_name: {group_name}, local rank: {group().local_rank},"
|
f"MindSpore reuse torch group: {device_group}, group_name: {group_name}, local rank: {group().local_rank},"
|
||||||
f"hccl communicator handle: {hex(hccl_comm_handle)}",
|
f"hccl communicator handle: {hex(hccl_comm_handle)}",
|
||||||
flush=True,
|
|
||||||
)
|
)
|
||||||
# Create MS communication group by hccl comm handle to reuse Torch group.
|
# Create MS communication group by hccl comm handle to reuse Torch group.
|
||||||
group_options = GroupOptions()
|
group_options = GroupOptions()
|
||||||
|
|||||||
@@ -28,6 +28,19 @@ if _is_npu:
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _get_arch_from_config(config):
|
||||||
|
mindspore_models = import_model_classes("sgl_mindspore.models")
|
||||||
|
architectures = getattr(config, "architectures", [])
|
||||||
|
if isinstance(architectures, str):
|
||||||
|
architectures = [architectures]
|
||||||
|
if not architectures:
|
||||||
|
raise ValueError("No model architectures are specified")
|
||||||
|
for arch in architectures:
|
||||||
|
if arch in mindspore_models:
|
||||||
|
return mindspore_models[arch]
|
||||||
|
raise ValueError(f"Unsupported arch {architectures}")
|
||||||
|
|
||||||
|
|
||||||
def tensor_torch2ms(x: torch.Tensor):
|
def tensor_torch2ms(x: torch.Tensor):
|
||||||
if x is None or not isinstance(x, torch.Tensor):
|
if x is None or not isinstance(x, torch.Tensor):
|
||||||
return x
|
return x
|
||||||
@@ -178,28 +191,14 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
arch = self.get_arch(self.config)
|
arch = self.get_arch(self.config)
|
||||||
self.model = arch(config=config, quant_config=quant_config)
|
self.model = arch(config=config, quant_config=quant_config)
|
||||||
|
|
||||||
self.casual_mask = LowerTriangularMask(
|
self.causal_mask = LowerTriangularMask(
|
||||||
self.config.param_dtype, self.config.max_position_embeddings
|
self.config.param_dtype, self.config.max_position_embeddings
|
||||||
)
|
)
|
||||||
self.key_cache = []
|
self.key_cache = []
|
||||||
self.value_cache = []
|
self.value_cache = []
|
||||||
|
|
||||||
def get_arch(self, config):
|
def get_arch(self, config):
|
||||||
# Get all implemented models
|
return _get_arch_from_config(config)
|
||||||
mindspore_models = import_model_classes("sgl_mindspore.models")
|
|
||||||
|
|
||||||
# Get arch from config
|
|
||||||
architectures = config.architectures
|
|
||||||
if isinstance(architectures, str):
|
|
||||||
architectures = [architectures]
|
|
||||||
if not architectures:
|
|
||||||
logger.warning("No model architectures are specified")
|
|
||||||
|
|
||||||
for arch in architectures:
|
|
||||||
if arch in mindspore_models:
|
|
||||||
return mindspore_models[arch]
|
|
||||||
if arch is None:
|
|
||||||
raise ValueError(f"Unsupported arch {architectures}")
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def use_mla(self):
|
def use_mla(self):
|
||||||
@@ -273,7 +272,7 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
model_inputs["position_ids"] = tensor_torch2ms(positions)
|
model_inputs["position_ids"] = tensor_torch2ms(positions)
|
||||||
model_inputs["q_seq_lens"] = ms.Tensor(q_seq_lens, dtype=ms.int32)
|
model_inputs["q_seq_lens"] = ms.Tensor(q_seq_lens, dtype=ms.int32)
|
||||||
model_inputs["attention_mask"] = self.casual_mask.gen_attention_mask(
|
model_inputs["attention_mask"] = self.causal_mask.gen_attention_mask(
|
||||||
is_prefill, model_inputs["position_ids"], q_seq_lens, batch_valid_length
|
is_prefill, model_inputs["position_ids"], q_seq_lens, batch_valid_length
|
||||||
).contiguous()
|
).contiguous()
|
||||||
model_inputs["out_cache_loc"] = tensor_torch2ms(forward_batch.out_cache_loc).to(
|
model_inputs["out_cache_loc"] = tensor_torch2ms(forward_batch.out_cache_loc).to(
|
||||||
@@ -303,5 +302,16 @@ class MindSporeForCausalLM(torch.nn.Module):
|
|||||||
logits_result = LogitsProcessorOutput(next_token_logits=tensor_ms2torch(logits))
|
logits_result = LogitsProcessorOutput(next_token_logits=tensor_ms2torch(logits))
|
||||||
return logits_result
|
return logits_result
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def get_model_config_for_expert_location(cls, config):
|
||||||
|
try:
|
||||||
|
arch_cls = _get_arch_from_config(config)
|
||||||
|
method = getattr(arch_cls, "get_model_config_for_expert_location", None)
|
||||||
|
if method is None:
|
||||||
|
return None
|
||||||
|
return method(config)
|
||||||
|
except Exception:
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
EntryClass = [MindSporeForCausalLM]
|
EntryClass = [MindSporeForCausalLM]
|
||||||
|
|||||||
Reference in New Issue
Block a user