[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
@@ -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)