[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:
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
@@ -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
@@ -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:
|
||||
# Non‑DeepEP (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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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"],
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user