[minimax m3][npu]Adaptation of Minimax M3(w8a8) for NPU platforms [1/2] (#32941)

Signed-off-by: Devashish Lal <devcode@fb.com>
Signed-off-by: Alexandre Milesi <milesial@users.noreply.github.com>
Signed-off-by: Faradawn Yang <73060648+faradawn@users.noreply.github.com>
Signed-off-by: Ryan Stewart <rystewart@nvidia.com>
Co-authored-by: ClownBin <chaobin1993@126.com>
Co-authored-by: huangzhenyu <q_m_p@qq.com>
Co-authored-by: clown <17490516+ClownBin@users.noreply.github.com>
Co-authored-by: badmer <374330057@qq.com.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
Co-authored-by: YAMY <74099316+YAMY1234@users.noreply.github.com>
Co-authored-by: Shangming Cai <csmthu@gmail.com>
Co-authored-by: Mick <mickjagger19@icloud.com>
Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
Co-authored-by: Brayden Zhong <brayden.zhong@radixark.ai>
Co-authored-by: Jimmy Shong <jimmysh341@gmail.com>
Co-authored-by: Zijie Xia <zijie.xia@radixark.ai>
Co-authored-by: Thomas Wang <thomawan@amd.com>
Co-authored-by: siyu <liusy58@linux.alibaba.com>
Co-authored-by: Yuang Chen <cya539102@antgroup.com>
Co-authored-by: Yuang Chen <1131578721@qq.com>
Co-authored-by: 黄孝君 <dingfangsu23@gmail.com>
Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com>
Co-authored-by: Xiaoyu Zhang <1182563586@qq.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
Co-authored-by: TobyMint <130973409+TobyMint@users.noreply.github.com>
Co-authored-by: TobyMint <tobymint@users.noreply.github.com>
Co-authored-by: Cheng Wan <54331508+ch-wan@users.noreply.github.com>
Co-authored-by: Tan Trinh <84185999+tanth47@users.noreply.github.com>
Co-authored-by: Lifan Shen <draftbks@gmail.com>
Co-authored-by: Justin Tong <justintong0323@outlook.com>
Co-authored-by: Qiaolin Yu <liin1211@outlook.com>
Co-authored-by: AMD-yanfeiwang <yanfei.wang@amd.com>
Co-authored-by: QIN2DIM <62018067+QIN2DIM@users.noreply.github.com>
Co-authored-by: Zhiyao Jiang <jessicajiang324@gmail.com>
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
Co-authored-by: DevashishLal-CB <devashish@rivosinc.com>
Co-authored-by: Devashish Lal <devcode@fb.com>
Co-authored-by: Alex Nails <alex.nails@radixark.ai>
Co-authored-by: Mohammad Miadh Angkad <176301910+mmangkad@users.noreply.github.com>
Co-authored-by: Baizhou Zhang <sobereddiezhang@gmail.com>
Co-authored-by: Michael Gschwind <mkgschwind+private@gmail.com>
Co-authored-by: weireweire <weiliangl@nvidia.com>
Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
Co-authored-by: Khoa Pham <khoa.pham@radixark.ai>
Co-authored-by: milesial <milesial@users.noreply.github.com>
Co-authored-by: elvischenv <219235043+elvischenv@users.noreply.github.com>
Co-authored-by: cctry <csycfl@gmail.com>
Co-authored-by: Jae B. <jlee5814@gmail.com>
Co-authored-by: Michael <13900043+michaelzhang-ai@users.noreply.github.com>
Co-authored-by: forrestl <16055533+forrestl111@users.noreply.github.com>
Co-authored-by: EchO <117733745+CyberSecurityErial@users.noreply.github.com>
Co-authored-by: Tanmay patil <tanmaypatil3151@gmail.com>
Co-authored-by: ybyang <10629930+whybeyoung@users.noreply.github.com>
Co-authored-by: Hsiu-Chun, Hung <160560375+Emmanuel0612@users.noreply.github.com>
Co-authored-by: Hung <Emmanuel0612@users.noreply.github.com>
Co-authored-by: HaiShaw <hixiao@gmail.com>
Co-authored-by: Bingxu Chen <bingxche@amd.com>
Co-authored-by: YC Yen-Ching Tseng <yctseng@amd.com>
Co-authored-by: Cherry_ming <136634645@qq.com>
Co-authored-by: Even Zhou <even.y.zhou@outlook.com>
Co-authored-by: sglang-npu-bot <sglangnpu@163.com>
Co-authored-by: Zhiqiang Xie <xiezhq@stanford.edu>
Co-authored-by: Tingwei Huang <huangtingwei9988@gmail.com>
Co-authored-by: Kaixi <kaiximatteoc@nvidia.com>
Co-authored-by: github-actions[bot] <github-actions[bot]@users.noreply.github.com>
Co-authored-by: Faradawn Yang <73060648+faradawn@users.noreply.github.com>
Co-authored-by: Ryan Stewart <rystewart@nvidia.com>
Co-authored-by: gjsheu <gjsheu@163.com>
Co-authored-by: Jinyan Yi <yjy20010615@gmail.com>
Co-authored-by: Ke Bao <ispobaoke@gmail.com>
Co-authored-by: huangtingwei <141888744+huangtingwei9988@users.noreply.github.com>
Co-authored-by: Hanming Lu <hanminglu@meta.com>
Co-authored-by: Jeremy Zhang <jeremy.zhang866@gmail.com>
Co-authored-by: Dmitrii Sergeev <dmi.sergeev@gmail.com>
Co-authored-by: Hao Zhang <zhisbug@users.noreply.github.com>
Co-authored-by: zhisbug <1654062+zhisbug@users.noreply.github.com>
Co-authored-by: Douglas Yang <dyang@college.harvard.edu>
Co-authored-by: gongwei1027 <gongwei833x@gmail.com>
Co-authored-by: ilyasher-harmonic <ilya.sherstyuk@harmonic.fun>
Co-authored-by: Hanming Lu <69857889+hanming-lu@users.noreply.github.com>
Co-authored-by: sglang-bot <sglangbot@gmail.com>
Co-authored-by: sglang-bot <232288953+sglang-bot@users.noreply.github.com>
Co-authored-by: Jimmy Shong <69131491+Jiminator@users.noreply.github.com>
Co-authored-by: hnyls2002 <lsyincs@gmail.com>
Co-authored-by: Meng, Hengyu <hengyu.meng@intel.com>
Co-authored-by: Shu Wang <shuw@nvidia.com>
Co-authored-by: Yanbin Jiang <jybsuper@gmail.com>
Co-authored-by: zijiec <zijie.chen@amd.com>
This commit is contained in:
vstone-w
2026-08-13 11:23:14 +08:00
committed by GitHub
co-authored by ClownBin huangzhenyu clown badmer Liangsheng Yin YAMY Shangming Cai Mick Brayden Zhong Brayden Zhong Jimmy Shong Zijie Xia Thomas Wang siyu Yuang Chen Yuang Chen 黄孝君 Xinyuan Tong Xiaoyu Zhang Cursor Claude Fable 5 TobyMint TobyMint Cheng Wan Tan Trinh Lifan Shen Justin Tong Qiaolin Yu AMD-yanfeiwang QIN2DIM Zhiyao Jiang Brayden Zhong DevashishLal-CB Devashish Lal Alex Nails Mohammad Miadh Angkad Baizhou Zhang Michael Gschwind weireweire weireweire Khoa Pham milesial elvischenv cctry Jae B. Michael forrestl EchO Tanmay patil ybyang Hsiu-Chun, Hung Hung HaiShaw Bingxu Chen YC Yen-Ching Tseng Cherry_ming Even Zhou sglang-npu-bot Zhiqiang Xie Tingwei Huang Kaixi github-actions[bot] Faradawn Yang Ryan Stewart gjsheu Jinyan Yi Ke Bao huangtingwei Hanming Lu Jeremy Zhang Dmitrii Sergeev Hao Zhang zhisbug Douglas Yang gongwei1027 ilyasher-harmonic Hanming Lu sglang-bot sglang-bot Jimmy Shong hnyls2002 Meng, Hengyu Shu Wang Yanbin Jiang zijiec
parent d44c836cfd
commit bca8ed4afc
18 changed files with 2314 additions and 175 deletions
+14
View File
@@ -1297,6 +1297,20 @@ class Envs:
SGLANG_MINIMAX_M3_FUSED_SWIGLU_MXFP8 = EnvBool(False)
SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE = EnvBool(False)
# MiniMax M3 NPU prefill MAIN-attention: route the sparse main attention through
# the native Ascend FA op `torch.ops.npu.npu_fused_infer_attention_score` (FIA)
# with a per-query CUSTOM block_table
SGLANG_MINIMAX_NPU_PREFILL_FIA = EnvBool(True)
# MiniMax-M3 NPU sparse INDEXER (decode + verify topk block selection): route
# through the native AscendC packed indexer op instead of the Triton indexer.
SGLANG_MINIMAX_NPU_NATIVE_INDEXER = EnvBool(False)
# MiniMax-M3 NPU sparse MAIN-attention (decode-main + verify-main): route the
# sparse main attention through the native AscendC sparse-attention op with the
# cached block_table override.
SGLANG_MINIMAX_NPU_NATIVE_ATTN = EnvBool(False)
# MiniMax-M3 on ROCm force-disables custom all-reduce in its model override
# (arg_groups/overrides.py) when aiter all-reduce fusion is off. Set this to
# opt back in and keep custom/quick all-reduce enabled -- e.g. to run the
@@ -5,7 +5,9 @@ import torch
from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE
from sglang.srt.environ import envs
from sglang.srt.mem_cache.memory_pool import (
MHATokenToKOnlyPool,
MHATokenToKVPool,
MiniMaxSparseKVPool,
MLATokenToKVPool,
get_tensor_size_bytes,
unwrap_write_loc,
@@ -390,6 +392,135 @@ class NPUMHATokenToKVPool(MHATokenToKVPool):
torch.npu.synchronize()
class NPUMHATokenToKOnlyPool(MHATokenToKOnlyPool):
"""NPU paged K-only cache used by MiniMax sparse index-only layers."""
def __init__(
self,
size: int,
page_size: int,
dtype: torch.dtype,
head_num: int,
head_dim: int,
layer_num: int,
device: str,
enable_memory_saver: bool,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
):
self.use_fia = get_bool_env_var("ASCEND_USE_FIA", "False")
super(MHATokenToKOnlyPool, self).__init__(
size=size,
page_size=page_size,
dtype=dtype,
layer_num=layer_num,
device=device,
enable_memory_saver=enable_memory_saver,
start_layer=start_layer,
end_layer=end_layer,
)
self.head_num = head_num
self.head_dim = head_dim
with self.memory_saver_adapter.region(GPU_MEMORY_TYPE_KV_CACHE):
self.k_buffer = torch.zeros(
(
self.layer_num,
self.size // self.page_size + 1,
self.page_size,
self.head_num,
self.head_dim,
),
dtype=self.store_dtype,
device=self.device,
)
if self.use_fia:
self.k_buffer = [
self.k_buffer[i].view(-1, 1, self.head_num, self.head_dim)
for i in range(self.layer_num)
]
self._finalize_allocation_log(size)
def _get_key_buffer(self, layer_id: int):
k_buffer = self.k_buffer[layer_id - self.start_layer]
if self.store_dtype != self.dtype:
return k_buffer.view(self.dtype)
return k_buffer
def set_k_buffer(
self,
layer_id: int,
loc_info,
cache_k: torch.Tensor,
) -> None:
loc, _, _ = unwrap_write_loc(loc_info)
if cache_k.dtype != self.dtype:
cache_k = cache_k.to(self.dtype)
if self.store_dtype != self.dtype:
cache_k = cache_k.view(self.store_dtype)
k_buffer_layer = self.k_buffer[layer_id - self.start_layer].view(
-1, self.head_num, self.head_dim
)
loc = loc.to(device=cache_k.device, dtype=torch.int32).contiguous()
torch_npu.npu_scatter_nd_update_(
k_buffer_layer,
loc.view(-1, 1),
cache_k.contiguous().view(-1, self.head_num, self.head_dim),
)
def get_contiguous_buf_infos(self):
data_ptrs = [
self.get_key_buffer(i).data_ptr()
for i in range(self.start_layer, self.start_layer + self.layer_num)
]
data_lens = [
self.get_key_buffer(i).nbytes
for i in range(self.start_layer, self.start_layer + self.layer_num)
]
if self.use_fia:
item_lens = [
self.get_key_buffer(i)[0].nbytes * self.page_size
for i in range(self.start_layer, self.start_layer + self.layer_num)
]
else:
item_lens = [
self.get_key_buffer(i)[0].nbytes
for i in range(self.start_layer, self.start_layer + self.layer_num)
]
return data_ptrs, data_lens, item_lens
def get_kv_size_bytes(self):
return get_tensor_size_bytes(self.k_buffer), 0
class NPUMiniMaxSparseKVPool(MiniMaxSparseKVPool):
"""MiniMax sparse wrapper backed by NPU paged MHA/index pools."""
def __init__(self, *args, **kwargs):
super().__init__(
*args,
main_pool_cls=NPUMHATokenToKVPool,
index_kv_pool_cls=NPUMHATokenToKVPool,
index_k_pool_cls=NPUMHATokenToKOnlyPool,
**kwargs,
)
def get_index_k_state_buf_infos(self):
pool = self.index_k_pool
n = pool.layer_num
data_ptrs = [pool.get_key_buffer(i).data_ptr() for i in range(n)]
data_lens = [pool.get_key_buffer(i).nbytes for i in range(n)]
if pool.use_fia:
item_lens = [
pool.get_key_buffer(i)[0].nbytes * pool.page_size for i in range(n)
]
else:
item_lens = [pool.get_key_buffer(i)[0].nbytes for i in range(n)]
return data_ptrs, data_lens, item_lens
class NPUMLATokenToKVPool(MLATokenToKVPool):
def __init__(
@@ -0,0 +1,301 @@
"""NPU patch for MiniMax M3 VL image and video preprocessing.
The MiniMax M3 VL image processor (MiniMaxM3VLImageProcessor) and video
processor (MiniMaxM3VLVideoProcessor) create 10-dimensional tensors during
patch extraction, which exceeds Ascend NPU's 8-dimension limit.
This patch restructures the computation using transform_patches_to_flatten
to stay within 8 dimensions, following the same pattern as the Qwen VL and
GLM-4.6V NPU patches.
"""
import math
from typing import List
import torch
from torchvision.transforms import InterpolationMode
from transformers.image_processing_utils import BatchFeature
from transformers.image_processing_utils_fast import (
group_images_by_shape,
reorder_images,
)
from transformers.image_utils import PILImageResampling, SizeDict
from transformers.utils import TensorType
from transformers.video_utils import group_videos_by_shape, reorder_videos
from sglang.srt.hardware_backend.npu.modules.qwen_vl_processor import (
transform_patches_to_flatten,
)
MAX_RATIO = 200
def _round_by_factor(number: int, factor: int) -> int:
return round(number / factor) * factor
def _ceil_by_factor(number: int, factor: int) -> int:
return math.ceil(number / factor) * factor
def _floor_by_factor(number: int, factor: int) -> int:
return math.floor(number / factor) * factor
def _smart_resize(
height: int,
width: int,
factor: int = 28,
min_pixels: int = 4 * 28 * 28,
max_pixels: int = 451584,
) -> tuple[int, int]:
if max(height, width) / min(height, width) > MAX_RATIO:
raise ValueError(
f"absolute aspect ratio must be smaller than {MAX_RATIO}, "
f"got {max(height, width) / min(height, width)}"
)
h_bar = max(factor, _round_by_factor(height, factor))
w_bar = max(factor, _round_by_factor(width, factor))
if h_bar * w_bar > max_pixels:
beta = math.sqrt((height * width) / max_pixels)
h_bar = _floor_by_factor(height / beta, factor)
w_bar = _floor_by_factor(width / beta, factor)
elif h_bar * w_bar < min_pixels:
beta = math.sqrt(min_pixels / (height * width))
h_bar = _ceil_by_factor(height * beta, factor)
w_bar = _ceil_by_factor(width * beta, factor)
return h_bar, w_bar
def npu_wrapper_minimax_m3_image_preprocess(func):
def _preprocess(
self,
images: List[torch.Tensor],
do_resize: bool,
size: SizeDict,
resample: PILImageResampling | InterpolationMode | int | None,
do_rescale: bool,
rescale_factor: float,
do_normalize: bool,
image_mean: float | List[float] | None,
image_std: float | List[float] | None,
patch_size: int,
temporal_patch_size: int,
merge_size: int,
max_pixels: int,
disable_grouping: bool | None,
return_tensors: str | TensorType | None,
**kwargs,
) -> BatchFeature:
grouped_images, grouped_images_index = group_images_by_shape(
images, disable_grouping=disable_grouping
)
resized_images_grouped = {}
factor = patch_size * merge_size
for shape, stacked_images in grouped_images.items():
height, width = stacked_images.shape[-2:]
if do_resize:
resized_height, resized_width = _smart_resize(
height,
width,
factor=factor,
max_pixels=max_pixels,
)
stacked_images = self.resize(
stacked_images,
size=SizeDict(height=resized_height, width=resized_width),
resample=resample,
)
resized_images_grouped[shape] = stacked_images
resized_images = reorder_images(resized_images_grouped, grouped_images_index)
grouped_images, grouped_images_index = group_images_by_shape(
resized_images, disable_grouping=disable_grouping
)
processed_images_grouped = {}
processed_grids = {}
for shape, stacked_images in grouped_images.items():
resized_height, resized_width = stacked_images.shape[-2:]
patches = self.rescale_and_normalize(
stacked_images,
do_rescale,
rescale_factor,
do_normalize,
image_mean,
image_std,
)
if patches.ndim == 4:
patches = patches.unsqueeze(1)
if patches.shape[1] % temporal_patch_size != 0:
repeats = patches[:, -1:].repeat(
1,
temporal_patch_size - (patches.shape[1] % temporal_patch_size),
1,
1,
1,
)
patches = torch.cat([patches, repeats], dim=1)
batch_size, grid_t, channel = patches.shape[:3]
grid_t = grid_t // temporal_patch_size
grid_h, grid_w = resized_height // patch_size, resized_width // patch_size
flatten_patches = transform_patches_to_flatten(
patches,
batch_size,
grid_t,
temporal_patch_size,
channel,
grid_h,
grid_w,
patch_size,
merge_size,
)
processed_images_grouped[shape] = flatten_patches
processed_grids[shape] = [[grid_t, grid_h, grid_w]] * batch_size
processed_images = reorder_images(
processed_images_grouped, grouped_images_index
)
processed_grids = reorder_images(processed_grids, grouped_images_index)
pixel_values = torch.cat(processed_images, dim=0)
image_grid_thw = torch.tensor(processed_grids, dtype=torch.long)
return BatchFeature(
data={"pixel_values": pixel_values, "image_grid_thw": image_grid_thw},
tensor_type=return_tensors,
)
return _preprocess
def npu_wrapper_minimax_m3_video_preprocess(func):
def _preprocess(
self,
videos: List[torch.Tensor],
do_convert_rgb: bool,
do_resize: bool,
size: SizeDict,
resample: PILImageResampling | InterpolationMode | int | None,
do_rescale: bool,
rescale_factor: float,
do_normalize: bool,
image_mean: float | List[float] | None,
image_std: float | List[float] | None,
patch_size: int,
temporal_patch_size: int,
merge_size: int,
min_pixels: int,
max_pixels: int,
return_tensors: str | TensorType | None = None,
**kwargs,
) -> BatchFeature:
grouped_videos, grouped_videos_index = group_videos_by_shape(videos)
resized_videos_grouped = {}
factor = patch_size * merge_size
for shape, stacked_videos in grouped_videos.items():
batch_size, num_frames, channels, height, width = stacked_videos.shape
resized_height, resized_width = height, width
if do_resize:
resized_height, resized_width = _smart_resize(
height,
width,
factor=factor,
min_pixels=min_pixels,
max_pixels=max_pixels,
)
stacked_videos = stacked_videos.view(
batch_size * num_frames, channels, height, width
)
stacked_videos = self.resize(
stacked_videos,
size=SizeDict(height=resized_height, width=resized_width),
resample=resample,
)
stacked_videos = stacked_videos.view(
batch_size,
num_frames,
channels,
resized_height,
resized_width,
)
resized_videos_grouped[shape] = stacked_videos
resized_videos = reorder_videos(resized_videos_grouped, grouped_videos_index)
grouped_videos, grouped_videos_index = group_videos_by_shape(resized_videos)
processed_videos_grouped = {}
processed_grids = {}
for shape, stacked_videos in grouped_videos.items():
resized_height, resized_width = stacked_videos.shape[-2:]
patches = self.rescale_and_normalize(
stacked_videos,
do_rescale,
rescale_factor,
do_normalize,
image_mean,
image_std,
)
if pad := -patches.shape[1] % temporal_patch_size:
repeats = patches[:, -1:].expand(-1, pad, -1, -1, -1)
patches = torch.cat([patches, repeats], dim=1)
batch_size, grid_t, channels = patches.shape[:3]
grid_t = grid_t // temporal_patch_size
grid_h, grid_w = resized_height // patch_size, resized_width // patch_size
flatten_patches = transform_patches_to_flatten(
patches,
batch_size,
grid_t,
temporal_patch_size,
channels,
grid_h,
grid_w,
patch_size,
merge_size,
)
processed_videos_grouped[shape] = flatten_patches
processed_grids[shape] = [[grid_t, grid_h, grid_w]] * batch_size
processed_videos = reorder_videos(
processed_videos_grouped, grouped_videos_index
)
processed_grids = reorder_videos(processed_grids, grouped_videos_index)
pixel_values_videos = torch.cat(processed_videos, dim=0)
video_grid_thw = torch.tensor(processed_grids, dtype=torch.long)
return BatchFeature(
data={
"pixel_values_videos": pixel_values_videos,
"video_grid_thw": video_grid_thw,
},
tensor_type=return_tensors,
)
return _preprocess
def npu_apply_minimax_m3_image_preprocess_patch(image_processor):
cls = type(image_processor)
if getattr(cls, "_sglang_npu_patched", False):
return
cls._preprocess = npu_wrapper_minimax_m3_image_preprocess(cls._preprocess)
cls._sglang_npu_patched = True
def npu_apply_minimax_m3_video_preprocess_patch(video_processor):
cls = type(video_processor)
if getattr(cls, "_sglang_npu_video_patched", False):
return
cls._preprocess = npu_wrapper_minimax_m3_video_preprocess(cls._preprocess)
cls._sglang_npu_video_patched = True
@@ -22,7 +22,7 @@ class BaseActivation(ABC):
# =============================================================================
# Concrete activation implementations (unchanged except removed 8.)
# Concrete activation implementations
# =============================================================================
class NPUSwiglu(BaseActivation):
def _apply_activation(self, hidden_states: torch.Tensor):
@@ -65,11 +65,31 @@ class NPUSwigluQuantWithScales(BaseActivation):
class NPUSwigluDeepEPKernel(BaseActivation):
def __init__(self, need_quant: bool = True):
from sgl_kernel_npu.activation.swiglu_quant import swiglu_quant
"""DeepEP grouped SwiGLU for the Ascend MoE runner; picks ``swiglu_quant`` vs the MiniMax
SwiGLU-OAI variant (``swiglu_oai_quant``: ``gate*sigmoid(gate*alpha)*(up+1)`` w/ clamping)
based on whether ``alpha``/``limit`` are given. The runner must forward
``gemm1_alpha``/``gemm1_clamp_limit`` here or experts fall back to wrong SwiGLU."""
self._kernel = swiglu_quant
def __init__(
self,
need_quant: bool = True,
alpha: Optional[float] = None,
limit: Optional[float] = None,
):
self.need_quant = need_quant
self.alpha = alpha
self.limit = limit
self._use_oai = alpha is not None and limit is not None
if self._use_oai:
from sgl_kernel_npu.activation.swiglu_oai_quant import (
swiglu_oai_quant,
)
self._kernel = swiglu_oai_quant
else:
from sgl_kernel_npu.activation.swiglu_quant import swiglu_quant
self._kernel = swiglu_quant
def _apply_activation(
self,
@@ -77,9 +97,19 @@ class NPUSwigluDeepEPKernel(BaseActivation):
group_list: torch.Tensor,
group_list_type: int,
):
hidden_states, per_token_scale = self._kernel(
hidden_states, group_list, group_list_type, need_quant=self.need_quant
)
if self._use_oai:
hidden_states, per_token_scale = self._kernel(
hidden_states,
self.alpha,
self.limit,
need_quant=self.need_quant,
group_list=group_list,
group_list_type=group_list_type,
)
else:
hidden_states, per_token_scale = self._kernel(
hidden_states, group_list, group_list_type, need_quant=self.need_quant
)
if self.need_quant:
return hidden_states, per_token_scale
return hidden_states, None
@@ -99,7 +99,7 @@ def fused_topk_npu(
k_group=topk_config.topk_group if use_grouped_topk else 1,
group_count=topk_config.num_expert_group if use_grouped_topk else 1,
group_select_mode=(1 if use_grouped_topk else 0),
renorm=0,
renorm=renormalize,
# 1 for sigmoid, 0 for softmax
norm_type=(0 if topk_config.scoring_func == "softmax" else 1),
routed_scaling_factor=(
File diff suppressed because it is too large Load Diff
+11 -3
View File
@@ -174,6 +174,8 @@ if _is_npu:
import torch_npu
from sgl_kernel_npu.norm.add_rmsnorm_bias import add_gemma_rms_norm
_NPU_GEMMA_RMS_NORM_TRITON_MAX_HIDDEN_SIZE = 5120
@lru_cache(maxsize=1)
def _get_aiter_per_group_quant():
@@ -1143,9 +1145,15 @@ class GemmaRMSNorm(BaseFusedOp):
if residual is not None:
if post_residual_addition is not None:
residual = residual + post_residual_addition
norm_out, residual = add_gemma_rms_norm(
x, self.weight, residual, self.variance_epsilon
)
if x.shape[-1] > _NPU_GEMMA_RMS_NORM_TRITON_MAX_HIDDEN_SIZE:
gamma = self.gemma_weight.to(x.dtype)
norm_out, _, residual = torch_npu.npu_add_rms_norm(
residual, x, gamma, self.variance_epsilon
)
else:
norm_out, residual = add_gemma_rms_norm(
x, self.weight, residual, self.variance_epsilon
)
return norm_out, residual
x, _ = torch_npu.npu_gemma_rms_norm(x, self.weight, self.variance_epsilon)
@@ -419,6 +419,13 @@ class FusedMoE(torch.nn.Module):
if expert_mask is not None:
self.register_buffer("expert_mask_gpu", expert_mask, persistent=False)
self._use_ascend_fuseep = get_moe_a2a_backend().is_ascend_fuseep()
# Expose swigluoai alpha/clamp on the layer so deepep's W8A8 apply
# (apply_without_routing_weights) picks swiglu_oai_quant instead of plain
# npu_swiglu. fuseep injects these via fuseep_activation (aclnnFusedDeepMoe
# internal); deepep reads them here via getattr(layer, "swiglu_alpha").
# Default None (no-op for non-swigluoai models).
self.swiglu_alpha = gemm1_alpha
self.swiglu_clamp_limit = gemm1_clamp_limit
if (
get_moe_runner_backend().is_flashinfer_trtllm_routed()
@@ -111,7 +111,11 @@ class AscendRunnerCore(MoeRunnerCore):
linear_beta=config.gemm1_clamp_limit,
)
else:
self.activation = NPUSwigluDeepEPKernel(need_quant=is_quant_kernel)
self.activation = NPUSwigluDeepEPKernel(
need_quant=is_quant_kernel,
alpha=config.gemm1_alpha,
limit=config.gemm1_clamp_limit,
)
else:
# NonDeepEP (ascend_tp) path
# 1. Choose the base activation according to the quant method
@@ -105,6 +105,18 @@ class ModelSlimConfig(QuantizationConfig):
for k, v in quant_config.items()
}
# Add an mlp.* alias for each block_sparse_moe.* key but KEEP the original,
# so both module namings resolve
for k in list(quant_config.keys()):
if not isinstance(k, str):
continue
if "block_sparse_moe" in k:
quant_config[
k.replace("block_sparse_moe.experts", "mlp.experts").replace(
"block_sparse_moe.shared_experts", "mlp.shared_experts"
)
] = quant_config[k]
self.quant_description = quant_config
ignore = cast(List[str], quant_config.get("ignore", []))
self.ignore = ignore if ignore is not None else []
@@ -993,6 +993,10 @@ class KVCacheConfigurator:
full_max_total_num_tokens=sizes.full_max_total_num_tokens,
swa_max_total_num_tokens=sizes.swa_max_total_num_tokens,
)
elif is_minimax_sparse(self.model_config.hf_config):
token_to_kv_pool = self._build_ascend_minimax_sparse_kv_pool(
max_total_num_tokens=sizes.max_total_num_tokens,
)
elif self.use_mla_backend:
token_to_kv_pool = self._build_ascend_mla_kv_pool(
max_total_num_tokens=sizes.max_total_num_tokens,
@@ -1238,6 +1242,37 @@ class KVCacheConfigurator:
)
return token_to_kv_pool
def _build_ascend_minimax_sparse_kv_pool(
self, *, max_total_num_tokens: int
) -> KVCache:
_hf_config = self.model_config.hf_config
sparse_cfg = get_minimax_sparse_attention_config(_hf_config)
dense_layer_ids, sparse_layer_ids = get_minimax_sparse_layer_ids(sparse_cfg)
disable_value_sparse_layer_ids = get_minimax_sparse_disable_value_layer_ids(
sparse_cfg
)
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
NPUMiniMaxSparseKVPool,
)
token_to_kv_pool = NPUMiniMaxSparseKVPool(
size=max_total_num_tokens,
page_size=self.server_args.page_size,
dtype=self.kv_cache_dtype,
index_dtype=self.model_dtype,
head_num=self.model_config.get_num_kv_heads(get_parallel().attn_tp_size),
head_dim=self.model_config.head_dim,
idx_head_dim=sparse_cfg["sparse_index_dim"],
dense_layer_ids=dense_layer_ids,
sparse_layer_ids=sparse_layer_ids,
disable_value_sparse_layer_ids=disable_value_sparse_layer_ids,
device=self.device,
enable_memory_saver=self.server_args.enable_memory_saver,
start_layer=self.layer_info.start_layer,
end_layer=self.layer_info.end_layer,
)
return token_to_kv_pool
def _build_ascend_mla_kv_pool(
self, *, max_total_num_tokens: int, is_dsa_model: bool
) -> KVCache:
+19 -7
View File
@@ -4629,6 +4629,18 @@ class MHATokenToKOnlyPool(KVCache):
self.layer_transfer_counter.wait_until(layer_id - self.start_layer)
return self._get_key_buffer(layer_id)
def set_k_buffer(
self,
layer_id: int,
loc: torch.Tensor,
cache_k: torch.Tensor,
) -> None:
if cache_k.dtype != self.dtype:
cache_k = cache_k.to(self.dtype)
if self.store_dtype != self.dtype:
cache_k = cache_k.view(self.store_dtype)
self.k_buffer[layer_id][loc] = cache_k
def get_value_buffer(self, layer_id: int) -> torch.Tensor:
raise NotImplementedError("MHATokenToKOnlyPool does not allocate V")
@@ -4673,6 +4685,9 @@ class MiniMaxSparseKVPool(KVCache):
index_dtype: Optional[torch.dtype] = None,
start_layer: Optional[int] = None,
end_layer: Optional[int] = None,
main_pool_cls=MHATokenToKVPool,
index_kv_pool_cls=MHATokenToKVPool,
index_k_pool_cls=MHATokenToKOnlyPool,
):
# Do not call super().__init__() — delegate to sub-pools instead.
self.size = size
@@ -4714,7 +4729,7 @@ class MiniMaxSparseKVPool(KVCache):
gid: i for i, gid in enumerate(local_k_only_sparse_layer_ids)
}
self.main_pool = MHATokenToKVPool(
self.main_pool = main_pool_cls(
size=size,
page_size=page_size,
dtype=dtype,
@@ -4728,7 +4743,7 @@ class MiniMaxSparseKVPool(KVCache):
)
self.index_kv_pool: Optional[MHATokenToKVPool] = (
MHATokenToKVPool(
index_kv_pool_cls(
size=size,
page_size=page_size,
dtype=index_dtype,
@@ -4743,7 +4758,7 @@ class MiniMaxSparseKVPool(KVCache):
)
self.index_k_pool: Optional[MHATokenToKOnlyPool] = (
MHATokenToKOnlyPool(
index_k_pool_cls(
size=size,
page_size=page_size,
dtype=index_dtype,
@@ -4891,10 +4906,7 @@ class MiniMaxSparseKVPool(KVCache):
if cache_idx_k.dtype != sub_pool.dtype:
if k_scale is not None:
cache_idx_k = cache_idx_k / k_scale
cache_idx_k = cache_idx_k.to(sub_pool.dtype)
if sub_pool.store_dtype != sub_pool.dtype:
cache_idx_k = cache_idx_k.view(sub_pool.store_dtype)
sub_pool.k_buffer[mapped_id][loc] = cache_idx_k
sub_pool.set_k_buffer(mapped_id, loc, cache_idx_k)
def _can_fuse_kv_index_store(
self,
+1
View File
@@ -242,6 +242,7 @@ def _get_quantization_config(
"q_a_proj",
"kv_a_proj_with_mqa",
],
"index_qkv_proj": ["index_q_proj", "index_k_proj"],
},
}
)
+115 -17
View File
@@ -42,7 +42,9 @@ from sglang.srt.layers.communicator import (
ScatterMode,
enable_moe_dense_fully_dp,
)
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.layers.dp_attention import (
is_dp_attention_enabled,
)
from sglang.srt.layers.layernorm import GemmaRMSNorm, RMSNorm
from sglang.srt.layers.linear import (
MergedColumnParallelLinear,
@@ -89,6 +91,7 @@ from sglang.srt.utils import (
get_device_sm,
is_cuda,
is_hip,
is_npu,
log_info_on_rank0,
make_layers,
)
@@ -96,6 +99,7 @@ from sglang.srt.utils.hf_transformers_utils import get_rope_config
_is_cuda = is_cuda()
_is_hip = is_hip()
_is_npu = is_npu()
_device_sm = get_device_sm()
_FP8_KV_DTYPES = (
@@ -121,6 +125,16 @@ if _is_hip:
except ImportError:
_has_rocm_qk_norm_rope = False
if _is_npu:
from sgl_kernel_npu.norm.split_qkv_rmsnorm_rope_pos_cache_half_npu import (
split_qkv_rmsnorm_rope_pos_cache_half_npu,
)
from sglang.srt.hardware_backend.npu.utils import (
process_shared_expert,
wait_share_stream,
)
logger = logging.getLogger(__name__)
@@ -213,6 +227,14 @@ def build_minimax_fused_qkv_index(model: nn.Module) -> None:
class MiniMaxM3MLP(nn.Module):
@staticmethod
def _swigluoai_fused(x: torch.Tensor, alpha: float, limit: float) -> torch.Tensor:
"""swiglu_oai using fused Triton kernel (sgl_kernel_npu), no quant."""
from sgl_kernel_npu.activation.swiglu_oai_quant import swiglu_oai_quant
out, _ = swiglu_oai_quant(x, alpha, limit, need_quant=False)
return out
def __init__(
self,
config: PretrainedConfig,
@@ -249,13 +271,18 @@ class MiniMaxM3MLP(nn.Module):
if hidden_act == "silu":
self.act_fn = SiluAndMul()
elif hidden_act == "swigluoai":
from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import (
swiglu_no_interleaved_with_alpha_and_limit,
)
if _is_npu:
self.act_fn = lambda x: self._swigluoai_fused(
x, config.swiglu_alpha, config.swiglu_limit
)
else:
from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe import (
swiglu_no_interleaved_with_alpha_and_limit,
)
self.act_fn = lambda x: swiglu_no_interleaved_with_alpha_and_limit(
x, config.swiglu_alpha, config.swiglu_limit
)
self.act_fn = lambda x: swiglu_no_interleaved_with_alpha_and_limit(
x, config.swiglu_alpha, config.swiglu_limit
)
else:
raise ValueError(
f"Unsupported activation: {hidden_act}. Only silu is supported for now."
@@ -264,6 +291,7 @@ class MiniMaxM3MLP(nn.Module):
def forward(
self,
x,
forward_batch: Optional[ForwardBatch] = None,
should_allreduce_fusion: bool = False,
use_reduce_scatter: bool = False,
):
@@ -416,10 +444,22 @@ class MiniMaxM3MoE(nn.Module):
def forward_deepep(
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
) -> torch.Tensor:
"""DeepEP MoE forward: routed experts via a2a, shared experts replicated."""
shared_output = None
enable_npu_dual_stream = _is_npu and (
forward_batch.forward_mode.is_extend()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_decode()
)
if hidden_states.shape[0] > 0:
shared_output = self._forward_shared_experts(hidden_states)
router_logits = self._compute_router_logits(hidden_states)
if enable_npu_dual_stream:
# Overlap shared experts with router/experts on a separate stream.
shared_output = process_shared_expert(
hidden_states, self._forward_shared_experts
)
else:
shared_output = self._forward_shared_experts(hidden_states)
topk_output = self.topk(
hidden_states,
router_logits,
@@ -435,6 +475,9 @@ class MiniMaxM3MoE(nn.Module):
# shared experts are replicated (tp_size=1), so both add directly.
final_hidden_states = self.experts(hidden_states, topk_output)
if enable_npu_dual_stream:
wait_share_stream()
if shared_output is not None:
final_hidden_states = final_hidden_states + shared_output
@@ -442,6 +485,9 @@ class MiniMaxM3MoE(nn.Module):
def _compute_router_logits(self, hidden_states: torch.Tensor) -> torch.Tensor:
if self.bf16_router_gemm:
if _is_npu:
# NPU lacks aten::mm.dtype; bf16 mm then cast keeps topk semantics.
return torch.mm(hidden_states, self.gate.weight.t()).float()
return torch.mm(
hidden_states, self.gate.weight.t(), out_dtype=torch.float32
)
@@ -501,6 +547,7 @@ class MiniMaxM3Attention(nn.Module):
self.max_position_embeddings = getattr(config, "max_position_embeddings", 8192)
self.rotary_dim = getattr(config, "rotary_dim", self.head_dim)
self.use_qk_norm = getattr(config, "use_qk_norm", False)
self.qk_norm_type = getattr(config, "qk_norm_type", "per_layer")
self.use_gemma_norm = getattr(config, "use_gemma_norm", False)
@@ -994,6 +1041,44 @@ class MiniMaxM3Attention(nn.Module):
return q, k, idx_q, idx_k
return self._sparse_qk_index_norm_rope(positions, q, k, idx_q, idx_k)
def forward_prepare_npu(
self,
positions: torch.Tensor,
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
):
"""NPU qkv projection + fused norm/RoPE/split; returns (None, fb, inner_state)."""
if hidden_states.shape[0] == 0:
assert (
not self.o_proj.reduce_results
), "short-circuiting allreduce will lead to hangs"
return hidden_states, forward_batch, None
qkv, _ = self.qkv_proj(hidden_states)
q, k, v = split_qkv_rmsnorm_rope_pos_cache_half_npu(
input_tensor=qkv,
positions=positions.reshape(-1),
cos_sin_cache=self.rotary_emb.cos_sin_cache,
q_hidden_size=self.q_size,
kv_hidden_size=self.kv_size,
head_dim=self.head_dim,
eps=self.q_norm.variance_epsilon,
q_weight=self.q_norm.gemma_weight,
k_weight=self.k_norm.gemma_weight,
rope_dim=self.rotary_dim,
cast_norm_to_bf16=True,
)
if self.is_sparse_attention_layer:
idx_qkv, _ = self.index_qkv_proj(hidden_states)
# Index attention disables the V head on all M3 sparse layers, so
# index_qkv_proj emits a 2-way [q|k] tensor.
idx_q, idx_k, idx_v = self._split_index_qkv(idx_qkv)
idx_q, idx_k = self._index_qk_norm_rope(positions, idx_q, idx_k)
inner_state = (q, k, v, idx_q, idx_k, idx_v, forward_batch)
else:
inner_state = (q, k, v, forward_batch)
return None, forward_batch, inner_state
def forward_prepare(
self,
positions: torch.Tensor,
@@ -1110,8 +1195,7 @@ class MiniMaxM3Attention(nn.Module):
output, _ = self.o_proj(attn_output)
if self.disable_index_value:
return output
# idx_replica_size ranks produce identical idx_o; pre-divide idx_o (not the
# o_proj weight) so the TP all-reduce sums right and stays FP8-quant-safe.
# Pre-divide idx_o (not the weight) so the TP all-reduce sums right.
if self.idx_replica_size > 1:
idx_o = idx_o / self.idx_replica_size
idx_output, _ = self.index_o_proj(idx_o)
@@ -1128,11 +1212,18 @@ class MiniMaxM3Attention(nn.Module):
hidden_states: torch.Tensor,
forward_batch: ForwardBatch,
) -> torch.Tensor:
s = self.forward_prepare(
positions=positions,
hidden_states=hidden_states,
forward_batch=forward_batch,
)
if _is_npu:
s = self.forward_prepare_npu(
positions=positions,
hidden_states=hidden_states,
forward_batch=forward_batch,
)
else:
s = self.forward_prepare(
positions=positions,
hidden_states=hidden_states,
forward_batch=forward_batch,
)
return self.forward_core(s)
@@ -1282,8 +1373,9 @@ class MiniMaxM3DecoderLayer(nn.Module):
if self.is_layer_sparse or hidden_states.shape[0] != 0:
hidden_states = self.mlp(
hidden_states,
should_allreduce_fusion,
use_reduce_scatter,
forward_batch=forward_batch,
should_allreduce_fusion=should_allreduce_fusion,
use_reduce_scatter=use_reduce_scatter,
)
if should_allreduce_fusion:
@@ -1513,6 +1605,12 @@ class MiniMaxM3SparseForCausalLM(nn.Module):
else:
self.model.layers_to_capture = [val + 1 for val in layer_ids]
# forward checks the per-layer ``_is_layer_to_capture`` flag, not the id
# list, so set it explicitly (mirrors qwen3_next/qwen2_moe).
for layer_id in self.model.layers_to_capture:
if 0 <= layer_id < len(self.model.layers):
setattr(self.model.layers[layer_id], "_is_layer_to_capture", True)
def get_embed_and_head(self):
return self.model.embed_tokens.weight, self.lm_head.weight
+41
View File
@@ -135,6 +135,9 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module):
self.logits_processor = LogitsProcessor(text_config)
# For EAGLE3 support
self.capture_aux_hidden_states = False
@classmethod
def shared_experts_fusion_disable_reason(cls, hf_config, quant_config):
"""Why this checkpoint cannot fuse its shared expert, or None. Asked by
@@ -197,6 +200,36 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module):
def get_input_embeddings(self):
return self.model.embed_tokens
def get_embed_and_head(self):
# EAGLE3 target interface: share the text embed + lm_head with the draft.
return self.model.embed_tokens.weight, self.lm_head.weight
def set_eagle3_layers_to_capture(self, layer_ids: Optional[list[int]] = None):
# EAGLE3 target interface: select which decoder layers' hidden states the
# draft consumes. Mirrors MiniMaxM3SparseForCausalLM; operates on the inner
# MiniMaxM3Model (self.model), whose forward returns (hidden, aux) once set.
if not self.pp_group.is_last_rank:
return
self.capture_aux_hidden_states = True
# MiniMaxM3Model.forward captures at layer ENTRY (= previous layer's
# output), so to capture layer L's output we must mark layer L+1. Apply
# +1 on both paths so EAGLE3 works out-of-the-box even when the draft
# config omits ``eagle_aux_hidden_state_layer_ids`` (the upstream
# Inferact/MiniMax-M3-EAGLE3 checkpoint does not ship it); otherwise the
# default-path layers are off by one and draft accept collapses.
if layer_ids is None:
num_layers = self.config.text_config.num_hidden_layers
layer_ids = [2, num_layers // 2, num_layers - 3]
self.model.layers_to_capture = [val + 1 for val in layer_ids]
# MiniMaxM3Model.forward checks each layer's ``_is_layer_to_capture``
# attribute (not ``i in layers_to_capture``); set it explicitly so the
# (hidden, aux) tuple is actually returned during capture-enabled forwards.
for layer_id in self.model.layers_to_capture:
if 0 <= layer_id < len(self.model.layers):
setattr(self.model.layers[layer_id], "_is_layer_to_capture", True)
def forward(
self,
input_ids: torch.Tensor,
@@ -217,12 +250,20 @@ class MiniMaxM3SparseForConditionalGeneration(nn.Module):
pp_proxy_tensors=pp_proxy_tensors,
)
# EAGLE3: when layers_to_capture is set, MiniMaxM3Model.forward returns
# (hidden_states, aux_hidden_states) once aux is non-empty; on idle/warmup
# forwards with no captured tokens it returns a bare hidden tensor.
aux_hidden_states = None
if self.capture_aux_hidden_states and isinstance(hidden_states, tuple):
hidden_states, aux_hidden_states = hidden_states
if self.pp_group.is_last_rank and not get_embedding:
return self.logits_processor(
input_ids,
hidden_states,
self.lm_head,
forward_batch,
aux_hidden_states,
)
return hidden_states
@@ -512,6 +512,22 @@ class BaseMultimodalProcessor(ABC):
return "xpu"
if not _is_npu:
return f"cuda:{server_args.base_gpu_id}"
if processor.__class__.__name__ == "MiniMaxVLProcessor":
# MiniMax's image/video processors create 10-dim tensors during
# patch extraction, exceeding the Ascend 8-dim limit; patch them
# (same pattern as qwen-vl / GLM-4.6V) and run on NPU.
from sglang.srt.hardware_backend.npu.modules.minimax_m3_processor import (
npu_apply_minimax_m3_image_preprocess_patch,
npu_apply_minimax_m3_video_preprocess_patch,
)
npu_apply_minimax_m3_image_preprocess_patch(processor.image_processor)
if (
hasattr(processor, "video_processor")
and processor.video_processor is not None
):
npu_apply_minimax_m3_video_preprocess_patch(processor.video_processor)
return "npu"
if processor.__class__.__name__ not in {"Glm4vProcessor", "Glm46VProcessor"}:
# For qwen-vl, the processor hits a reshape issue from the Ascend
# dims restriction.
@@ -0,0 +1,176 @@
import importlib.util
import sys
import types
from contextlib import nullcontext
from pathlib import Path
import torch
class _FakeMemorySaverAdapter:
def region(self, _memory_type):
return nullcontext()
class _FakeKVCache:
def __init__(
self,
size,
page_size,
dtype,
layer_num,
device,
enable_memory_saver,
start_layer=None,
end_layer=None,
):
self.size = size
self.page_size = page_size
self.dtype = dtype
self.store_dtype = dtype
self.layer_num = layer_num
self.device = device
self.enable_memory_saver = enable_memory_saver
self.start_layer = start_layer or 0
self.end_layer = end_layer or layer_num - 1
self.memory_saver_adapter = _FakeMemorySaverAdapter()
self.mem_usage = 0
def _finalize_allocation_log(self, _num_tokens):
pass
class _FakeMHATokenToKVPool(_FakeKVCache):
def __init__(
self,
size,
page_size,
dtype,
head_num,
head_dim,
layer_num,
device,
enable_memory_saver,
v_head_dim=None,
swa_head_num=None,
swa_head_dim=None,
swa_v_head_dim=None,
start_layer=None,
end_layer=None,
**_kwargs,
):
super().__init__(
size,
page_size,
dtype,
layer_num,
device,
enable_memory_saver,
start_layer,
end_layer,
)
self.head_num = swa_head_num if swa_head_num is not None else head_num
self.head_dim = swa_head_dim if swa_head_dim is not None else head_dim
self.v_head_dim = (
swa_v_head_dim
if swa_v_head_dim is not None
else v_head_dim if v_head_dim is not None else head_dim
)
self._create_buffers()
class _FakeMHATokenToKOnlyPool(_FakeKVCache):
pass
class _FakeMiniMaxSparseKVPool:
def __init__(self, *args, **kwargs):
pass
class _FakeMLATokenToKVPool(_FakeKVCache):
pass
def _load_npu_memory_pool_module():
for name in (
"sglang",
"sglang.srt",
"sglang.srt.constants",
"sglang.srt.mem_cache",
"sglang.srt.mem_cache.memory_pool",
"sglang.srt.utils",
"sglang.srt.utils.common",
):
sys.modules.setdefault(name, types.ModuleType(name))
constants = sys.modules["sglang.srt.constants"]
constants.GPU_MEMORY_TYPE_KV_CACHE = "kv_cache"
memory_pool = sys.modules["sglang.srt.mem_cache.memory_pool"]
memory_pool.MHATokenToKVPool = _FakeMHATokenToKVPool
memory_pool.MHATokenToKOnlyPool = _FakeMHATokenToKOnlyPool
memory_pool.MiniMaxSparseKVPool = _FakeMiniMaxSparseKVPool
memory_pool.MLATokenToKVPool = _FakeMLATokenToKVPool
memory_pool.get_tensor_size_bytes = lambda tensor: tensor.nbytes
memory_pool.maybe_detect_oob = lambda *args, **kwargs: None
memory_pool.unwrap_write_loc = lambda loc_info: (loc_info, None, None)
utils = sys.modules["sglang.srt.utils"]
utils.get_bool_env_var = lambda _name, default: default == "True"
common = sys.modules["sglang.srt.utils.common"]
common.is_npu = lambda: False
module_path = (
Path(__file__).resolve().parents[3]
/ "python/sglang/srt/hardware_backend/npu/memory_pool_npu.py"
)
spec = importlib.util.spec_from_file_location(
"_npu_memory_pool_under_test", module_path
)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
return module
def test_npu_minimax_k_only_index_cache_uses_scatter_writer():
npu_memory_pool = _load_npu_memory_pool_module()
calls = []
class FakeTorchNpu:
@staticmethod
def npu_scatter_nd_update_(cache, indices, updates):
assert cache.shape == (10, 1, 4)
assert indices.shape == (2, 1)
assert updates.shape == (2, 1, 4)
calls.append((cache, indices, updates))
@staticmethod
def _npu_reshape_and_cache(*, key, value, key_cache, value_cache, slot_indices):
raise AssertionError("K-only MiniMax index cache should use scatter")
npu_memory_pool.torch_npu = FakeTorchNpu
pool = npu_memory_pool.NPUMHATokenToKOnlyPool(
size=8,
page_size=2,
dtype=torch.bfloat16,
head_num=1,
head_dim=4,
layer_num=1,
device="cpu",
enable_memory_saver=False,
)
loc = torch.tensor([1, 3], dtype=torch.int64)
cache_k = torch.randn((2, 1, 4), dtype=torch.bfloat16)
pool.set_k_buffer(0, loc, cache_k)
assert len(calls) == 1
k_size, v_size = pool.get_kv_size_bytes()
assert k_size > 0
assert v_size == 0
+193
View File
@@ -0,0 +1,193 @@
import importlib.util
import sys
import types
from pathlib import Path
from typing import NamedTuple
import torch
def _install_fake_modules():
for name in (
"sglang",
"sglang.srt",
"sglang.srt.eplb",
"sglang.srt.layers",
"sglang.srt.layers.moe",
"sglang.srt.state_capturer",
):
sys.modules.setdefault(name, types.ModuleType(name))
root = types.ModuleType("sgl_kernel_npu")
norm = types.ModuleType("sgl_kernel_npu.norm")
l1_norm_mod = types.ModuleType("sgl_kernel_npu.norm.l1_norm")
def l1_norm(x):
return x / x.sum(dim=-1, keepdim=True)
l1_norm_mod.l1_norm = l1_norm
sys.modules.setdefault("sgl_kernel_npu", root)
sys.modules.setdefault("sgl_kernel_npu.norm", norm)
sys.modules["sgl_kernel_npu.norm.l1_norm"] = l1_norm_mod
expert_distribution = types.ModuleType("sglang.srt.eplb.expert_distribution")
class Recorder:
@staticmethod
def on_select_experts(topk_ids):
pass
expert_distribution.get_global_expert_distribution_recorder = lambda: Recorder()
sys.modules["sglang.srt.eplb.expert_distribution"] = expert_distribution
expert_location = types.ModuleType("sglang.srt.eplb.expert_location_dispatch")
expert_location.topk_ids_logical_to_physical = lambda topk_ids, info: topk_ids
sys.modules["sglang.srt.eplb.expert_location_dispatch"] = expert_location
moe_topk = types.ModuleType("sglang.srt.layers.moe.topk")
class StandardTopKOutput(NamedTuple):
topk_weights: torch.Tensor
topk_ids: torch.Tensor
router_logits: torch.Tensor
def select_experts(*args, **kwargs):
raise AssertionError("fallback select_experts should not be used")
def capture_routed_experts_if_allowed(*args, **kwargs):
return None
moe_topk.StandardTopKOutput = StandardTopKOutput
moe_topk.select_experts = select_experts
moe_topk.capture_routed_experts_if_allowed = capture_routed_experts_if_allowed
sys.modules["sglang.srt.layers.moe.topk"] = moe_topk
routed_experts = types.ModuleType("sglang.srt.state_capturer.routed_experts")
routed_experts.get_global_experts_capturer = lambda: None
sys.modules["sglang.srt.state_capturer.routed_experts"] = routed_experts
def _load_npu_topk_module():
_install_fake_modules()
module_path = (
Path(__file__).resolve().parents[3]
/ "python/sglang/srt/hardware_backend/npu/moe/topk.py"
)
spec = importlib.util.spec_from_file_location("_npu_topk_under_test", module_path)
module = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(module)
return module
def _make_topk_config(
correction_bias, routed_scaling_factor, renormalize=True
) -> types.SimpleNamespace:
"""M3-shaped TopKConfig: sigmoid scoring, no grouped routing."""
return types.SimpleNamespace(
top_k=2,
use_grouped_topk=False,
correction_bias=correction_bias,
topk_group=None,
num_expert_group=None,
renormalize=renormalize,
scoring_func="sigmoid",
num_fused_shared_experts=0,
routed_scaling_factor=routed_scaling_factor,
apply_routed_scaling_factor_on_output=True,
)
def _run_fused_topk_npu(npu_topk, topk_config, router_logits):
return npu_topk.fused_topk_npu(
hidden_states=torch.zeros((1, 4), dtype=torch.bfloat16),
router_logits=router_logits,
topk_config=topk_config,
)
class _SigmoidFakeNpuOps:
"""npu_moe_gating_top_k mirroring the real sigmoid contract (norm_type=1)."""
def __init__(self, router_logits, routed_scaling_factor, expect_bias=None):
self.router_logits = router_logits
self.routed_scaling_factor = routed_scaling_factor
self.expect_bias = expect_bias
def npu_moe_gating_top_k_softmax(self, *args, **kwargs):
raise AssertionError("sigmoid routing must not use the softmax top-k op")
def npu_moe_gating_top_k(
self,
router_logits,
*,
k,
bias,
renorm,
norm_type,
routed_scaling_factor,
**kwargs,
):
# Contract: sigmoid scoring -> norm_type=1; bias (if any) must reach the
# op; renorm and the routed scaling factor are applied inside the op.
assert norm_type == 1
if self.expect_bias is not None:
assert bias is not None and bias.shape == self.expect_bias.shape
scores = (
(router_logits + bias).sigmoid()
if bias is not None
else router_logits.sigmoid()
)
values, ids = torch.topk(scores, k=k, dim=-1)
if renorm:
values = values / values.sum(dim=-1, keepdim=True)
values = values * routed_scaling_factor
return values, ids.to(torch.int32), None
def test_npu_sigmoid_topk_without_bias_uses_sigmoid_op(monkeypatch):
"""Sigmoid routing without correction bias must NOT fall into the softmax fast path.
Guards fused_topk_npu's fast-path branch: it previously matched
``not use_grouped_topk and correction_bias is None`` without excluding
sigmoid scoring, routing sigmoid models through the softmax op.
"""
npu_topk = _load_npu_topk_module()
router_logits = torch.tensor([[0.0, 1.0, 2.0]], dtype=torch.float32)
routed_scaling_factor = 2.5
fake = _SigmoidFakeNpuOps(router_logits, routed_scaling_factor, expect_bias=None)
monkeypatch.setattr(torch.ops, "npu", fake, raising=False)
topk_output = _run_fused_topk_npu(
npu_topk,
_make_topk_config(None, routed_scaling_factor),
router_logits,
)
raw = router_logits.sigmoid().topk(2, dim=-1).values
expected = raw / raw.sum(dim=-1, keepdim=True) * routed_scaling_factor
torch.testing.assert_close(topk_output.topk_weights, expected)
def test_npu_sigmoid_topk_with_routing_bias_matches_m3_config(monkeypatch):
"""M3 real config (use_routing_bias=True): bias must reach the sigmoid op."""
npu_topk = _load_npu_topk_module()
router_logits = torch.tensor([[0.0, 1.0, 2.0, 3.0]], dtype=torch.float32)
routed_scaling_factor = 2.5
correction_bias = torch.tensor([0.1, -0.2, 0.3, 0.05], dtype=torch.float32)
fake = _SigmoidFakeNpuOps(
router_logits, routed_scaling_factor, expect_bias=correction_bias
)
monkeypatch.setattr(torch.ops, "npu", fake, raising=False)
topk_output = _run_fused_topk_npu(
npu_topk,
_make_topk_config(correction_bias, routed_scaling_factor),
router_logits,
)
raw = (router_logits + correction_bias).sigmoid().topk(2, dim=-1).values
expected = raw / raw.sum(dim=-1, keepdim=True) * routed_scaling_factor
torch.testing.assert_close(topk_output.topk_weights, expected)