From 0b04e9da83db50fc48dc50c843bec81326add001 Mon Sep 17 00:00:00 2001 From: Brayden Zhong Date: Wed, 15 Jul 2026 18:05:22 -0700 Subject: [PATCH] Use fused A GEMM for `fc1_latent_proj` in NemotronH (#29692) Co-authored-by: Brayden Zhong --- python/sglang/jit_kernel/fused_a_gemm.py | 41 ++++++++++++++++++++- python/sglang/srt/models/deepseek_v2.py | 45 ++++++------------------ python/sglang/srt/models/nemotron_h.py | 22 ++++++++++-- 3 files changed, 70 insertions(+), 38 deletions(-) diff --git a/python/sglang/jit_kernel/fused_a_gemm.py b/python/sglang/jit_kernel/fused_a_gemm.py index 7ab37f31c..7d893444a 100644 --- a/python/sglang/jit_kernel/fused_a_gemm.py +++ b/python/sglang/jit_kernel/fused_a_gemm.py @@ -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, diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index b4b105dc9..73b7a4b05 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -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 diff --git a/python/sglang/srt/models/nemotron_h.py b/python/sglang/srt/models/nemotron_h.py index 2037cdbb2..4fb636a26 100644 --- a/python/sglang/srt/models/nemotron_h.py +++ b/python/sglang/srt/models/nemotron_h.py @@ -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)