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