[AMD] Use tgemm.mm for MoEGate router gemm in deepseek_v2.py (#21657)

This commit is contained in:
Thomas Wang
2026-03-31 00:55:40 -07:00
committed by GitHub
parent b4cb31f698
commit 5628e908ae
2 changed files with 8 additions and 32 deletions
+3 -23
View File
@@ -1,10 +1,7 @@
import torch
from aiter.ops.triton.fused_kv_cache import fused_qk_rope_cat_and_cache_mla
from aiter.ops.triton.fused_qk_concat import fused_qk_rope_cat
from aiter.ops.triton.gemm_a16w16 import gemm_a16w16
from aiter.ops.triton.gemm_a16w16_atomic import gemm_a16w16_atomic
from sglang.srt.utils import BumpAllocator
from aiter.tuned_gemm import tgemm
__all__ = ["fused_qk_rope_cat", "fused_qk_rope_cat_and_cache_mla"]
@@ -12,26 +9,9 @@ __all__ = ["fused_qk_rope_cat", "fused_qk_rope_cat_and_cache_mla"]
def aiter_dsv3_router_gemm(
hidden_states: torch.Tensor,
weight: torch.Tensor,
gemm_output_zero_allocator: BumpAllocator = None,
):
M = hidden_states.shape[0]
N = weight.shape[0]
y = None
if M <= 256:
# TODO (cagri): convert to bfloat16 as part of another kernel to save time
# for now it is also coupled with zero allocator.
if gemm_output_zero_allocator != None:
y = gemm_output_zero_allocator.allocate(M * N).view(M, N)
else:
y = torch.zeros((M, N), dtype=torch.float32, device=hidden_states.device)
if y is not None:
logits = gemm_a16w16_atomic(hidden_states, weight, y=y).to(hidden_states.dtype)
else:
logits = gemm_a16w16(hidden_states, weight)
return logits
"""Use aiter tuned GEMM dispatcher (tgemm.mm) to automatically select the GEMM kernel."""
return tgemm.mm(hidden_states, weight, otype=hidden_states.dtype)
def get_dsv3_gemm_output_zero_allocator_size(
+5 -9
View File
@@ -153,9 +153,11 @@ from sglang.srt.utils import (
use_intel_amx_backend,
)
if _use_aiter:
from sglang.srt.layers.rocm_linear_utils import aiter_dsv3_router_gemm
if _use_aiter_gfx95:
from sglang.srt.layers.rocm_linear_utils import (
aiter_dsv3_router_gemm,
get_dsv3_gemm_output_zero_allocator_size,
)
@@ -327,14 +329,8 @@ class MoEGate(nn.Module):
logits = dsv3_router_gemm(
hidden_states, self.weight, out_dtype=torch.float32
)
elif (
_use_aiter_gfx95
and hidden_states.shape[0] <= 256
and self.weight.shape[0] <= 256
):
logits = aiter_dsv3_router_gemm(
hidden_states, self.weight, gemm_output_zero_allocator
)
elif _use_aiter:
logits = aiter_dsv3_router_gemm(hidden_states, self.weight)
else:
logits = F.linear(hidden_states, self.weight, None)