[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 import torch
from aiter.ops.triton.fused_kv_cache import fused_qk_rope_cat_and_cache_mla 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.fused_qk_concat import fused_qk_rope_cat
from aiter.ops.triton.gemm_a16w16 import gemm_a16w16 from aiter.tuned_gemm import tgemm
from aiter.ops.triton.gemm_a16w16_atomic import gemm_a16w16_atomic
from sglang.srt.utils import BumpAllocator
__all__ = ["fused_qk_rope_cat", "fused_qk_rope_cat_and_cache_mla"] __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( def aiter_dsv3_router_gemm(
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
weight: torch.Tensor, weight: torch.Tensor,
gemm_output_zero_allocator: BumpAllocator = None,
): ):
M = hidden_states.shape[0] """Use aiter tuned GEMM dispatcher (tgemm.mm) to automatically select the GEMM kernel."""
N = weight.shape[0] return tgemm.mm(hidden_states, weight, otype=hidden_states.dtype)
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
def get_dsv3_gemm_output_zero_allocator_size( 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, use_intel_amx_backend,
) )
if _use_aiter:
from sglang.srt.layers.rocm_linear_utils import aiter_dsv3_router_gemm
if _use_aiter_gfx95: if _use_aiter_gfx95:
from sglang.srt.layers.rocm_linear_utils import ( from sglang.srt.layers.rocm_linear_utils import (
aiter_dsv3_router_gemm,
get_dsv3_gemm_output_zero_allocator_size, get_dsv3_gemm_output_zero_allocator_size,
) )
@@ -327,14 +329,8 @@ class MoEGate(nn.Module):
logits = dsv3_router_gemm( logits = dsv3_router_gemm(
hidden_states, self.weight, out_dtype=torch.float32 hidden_states, self.weight, out_dtype=torch.float32
) )
elif ( elif _use_aiter:
_use_aiter_gfx95 logits = aiter_dsv3_router_gemm(hidden_states, self.weight)
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
)
else: else:
logits = F.linear(hidden_states, self.weight, None) logits = F.linear(hidden_states, self.weight, None)