Use fused A GEMM for fc1_latent_proj in NemotronH (#29692)

Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Brayden Zhong
2026-07-15 18:05:22 -07:00
committed by GitHub
co-authored by Brayden Zhong
parent 1cc9493747
commit 0b04e9da83
3 changed files with 70 additions and 38 deletions
+40 -1
View File
@@ -16,7 +16,8 @@ from enum import Enum
import torch
from sglang.srt.utils.common import is_sm120_supported
from sglang.srt.layers.quantization.unquant import get_bf16_gemm_backend
from sglang.srt.utils.common import get_device_sm, is_cuda, is_sm120_supported
class FusedAGemmBackend(str, Enum):
@@ -30,6 +31,44 @@ _AUTO_BACKEND = (
FusedAGemmBackend.CUTEDSL if is_sm120_supported() else FusedAGemmBackend.JIT
)
_IS_CUDA = is_cuda()
_DEVICE_SM = get_device_sm()
def fused_a_gemm_weight_eligible(layer: torch.nn.Module) -> bool:
return (
layer.weight.dtype == torch.bfloat16
and layer.weight.shape[0] % 16 == 0
and layer.weight.shape[1] % 256 == 0
and _IS_CUDA
and _DEVICE_SM >= 90
)
def linear_with_fused_a_gemm(
layer: torch.nn.Module,
hidden_states: torch.Tensor,
*,
backend: "FusedAGemmBackend | str" = FusedAGemmBackend.AUTO,
) -> torch.Tensor:
# LoRA reads weight.T directly, bypassing the adapter, so fall back when active.
cutedsl_backend = get_bf16_gemm_backend().is_cutedsl()
if cutedsl_backend:
from sglang.jit_kernel.cutedsl_bf16_gemm import use_cutedsl_bf16_gemm
if (
not isinstance(hidden_states, tuple)
and 1 <= hidden_states.shape[0] <= 16
and not getattr(layer, "set_lora", False)
and not (
cutedsl_backend
and use_cutedsl_bf16_gemm(
hidden_states.shape[0], layer.weight.shape[0], layer.weight.shape[1]
)
)
):
return dsv3_fused_a_gemm(hidden_states, layer.weight.T, backend=backend)
return layer(hidden_states)[0]
def dsv3_fused_a_gemm(
mat_a: torch.Tensor,
+10 -35
View File
@@ -120,7 +120,6 @@ from sglang.srt.layers.quantization.fp8_utils import (
from sglang.srt.layers.quantization.mxfp4_flashinfer_trtllm_moe import (
maybe_fuse_routed_scale_and_shared_add,
)
from sglang.srt.layers.quantization.unquant import get_bf16_gemm_backend
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.layers.rotary_embedding import get_rope_wrapper
from sglang.srt.layers.utils import PPMissingLayer
@@ -214,7 +213,6 @@ if _is_cuda:
from sglang.jit_kernel.dsv3_router_gemm import (
dsv3_router_gemm as _jit_dsv3_router_gemm,
)
from sglang.jit_kernel.fused_a_gemm import dsv3_fused_a_gemm
elif _is_npu:
from sglang.srt.hardware_backend.npu.modules.deepseek_v2_attention_mla_npu import (
forward_dsa_core_npu,
@@ -224,11 +222,14 @@ elif _is_npu:
forward_mla_core_npu,
forward_mla_prepare_npu,
)
elif _is_musa:
from sgl_kernel import dsv3_fused_a_gemm
else:
pass
from sglang.jit_kernel.fused_a_gemm import (
fused_a_gemm_weight_eligible,
linear_with_fused_a_gemm,
)
logger = logging.getLogger(__name__)
_enable_pcg_dsv2_dual_stream = (
@@ -1782,11 +1783,7 @@ class DeepseekV2AttentionMLA(
self.use_min_latency_fused_a_gemm = (
self.has_fused_proj
and not self.is_packed_weight
and self.fused_qkv_a_proj_with_mqa.weight.dtype == torch.bfloat16
and self.fused_qkv_a_proj_with_mqa.weight.shape[0] % 16 == 0
and self.fused_qkv_a_proj_with_mqa.weight.shape[1] % 256 == 0
and _is_cuda
and _device_sm >= 90
and fused_a_gemm_weight_eligible(self.fused_qkv_a_proj_with_mqa)
)
self.fused_a_gemm_backend = "auto"
@@ -1991,35 +1988,13 @@ class DeepseekV2AttentionMLA(
self, hidden_states: torch.Tensor, forward_batch: ForwardBatch
):
assert self.q_lora_rank is not None
# When the module is wrapped with LoRA, the fused GEMM fast-path would
# bypass the adapter because it reads weight.T directly.
lora_active = getattr(self.fused_qkv_a_proj_with_mqa, "set_lora", False)
cutedsl_backend = get_bf16_gemm_backend().is_cutedsl()
if cutedsl_backend:
from sglang.jit_kernel.cutedsl_bf16_gemm import use_cutedsl_bf16_gemm
if (
(not isinstance(hidden_states, tuple))
and hidden_states.shape[0] >= 1
and hidden_states.shape[0] <= 16
and self.use_min_latency_fused_a_gemm
and not lora_active
and not (
cutedsl_backend
and use_cutedsl_bf16_gemm(
hidden_states.shape[0],
self.fused_qkv_a_proj_with_mqa.weight.shape[0],
self.fused_qkv_a_proj_with_mqa.weight.shape[1],
)
)
):
qkv_latent = dsv3_fused_a_gemm(
if self.use_min_latency_fused_a_gemm:
return linear_with_fused_a_gemm(
self.fused_qkv_a_proj_with_mqa,
hidden_states,
self.fused_qkv_a_proj_with_mqa.weight.T,
backend=self.fused_a_gemm_backend,
)
else:
qkv_latent = self.fused_qkv_a_proj_with_mqa(hidden_states)[0]
return qkv_latent
return self.fused_qkv_a_proj_with_mqa(hidden_states)[0]
def rebuild_cp_kv_cache(self, latent_cache, forward_batch, k_nope, k_pe):
# support allgather+rerrange
+20 -2
View File
@@ -99,6 +99,12 @@ from sglang.utils import logger
_is_cuda = is_cuda()
if _is_cuda:
from sglang.jit_kernel.fused_a_gemm import (
fused_a_gemm_weight_eligible,
linear_with_fused_a_gemm,
)
class NemotronHMLP(nn.Module):
def __init__(
@@ -241,6 +247,18 @@ class NemotronHMoE(nn.Module):
self.fc1_latent_proj = None
self.fc2_latent_proj = None
self.use_min_latency_fc1_gemm = (
self.use_latent_moe
and self.fc1_latent_proj is not None
and _is_cuda
and fused_a_gemm_weight_eligible(self.fc1_latent_proj)
)
def _apply_fc1_latent_proj(self, hidden_states: torch.Tensor) -> torch.Tensor:
if self.use_min_latency_fc1_gemm:
return linear_with_fused_a_gemm(self.fc1_latent_proj, hidden_states)
return self.fc1_latent_proj(hidden_states)[0]
def _forward_core(
self,
hidden_states: torch.Tensor,
@@ -268,7 +286,7 @@ class NemotronHMoE(nn.Module):
shared_output = None
topk_output = self.topk(hidden_states, router_logits)
if self.use_latent_moe:
hidden_states, _ = self.fc1_latent_proj(hidden_states)
hidden_states = self._apply_fc1_latent_proj(hidden_states)
final_hidden_states = self.experts(hidden_states, topk_output)
return final_hidden_states, shared_output
@@ -293,7 +311,7 @@ class NemotronHMoE(nn.Module):
)
topk_output = self.topk(hidden_states, router_logits)
if self.use_latent_moe:
hidden_states, _ = self.fc1_latent_proj(hidden_states)
hidden_states = self._apply_fc1_latent_proj(hidden_states)
final_hidden_states = self.experts(hidden_states, topk_output)
get_current_device_stream_fast().wait_stream(alt_stream)