diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 12b9b112e..473867493 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -447,12 +447,7 @@ class MoEGate(nn.Module): if config.topk_method == "noaux_tc" and not is_hash_moe: correction_bias_dtype = torch.float32 if quant_config is not None: - if ( - quant_config.get_name() == "modelopt_fp4" - and get_moe_runner_backend().is_flashinfer_trtllm() - ): - correction_bias_dtype = torch.bfloat16 - elif _use_aiter and quant_config.get_name() in ( + if _use_aiter and quant_config.get_name() in ( "fp8", "compressed_tensors", "quark", @@ -511,16 +506,13 @@ class MoEGate(nn.Module): elif _use_aiter: logits = aiter_dsv3_router_gemm(hidden_states, self.weight) - elif _is_npu: + elif not _is_cuda: logits = F.linear(hidden_states, self.weight, None) else: - if self.is_deepseek_v4: - from sglang.jit_kernel.dsv4 import linear_bf16_fp32 + # cuBLAS bf16 x bf16 -> fp32 GEMM (torch.mm's out_dtype kwarg is CUDA-only) + from sglang.jit_kernel.dsv4 import linear_bf16_fp32 - logits = linear_bf16_fp32(hidden_states, self.weight) - else: - # After testing, we may use the faster code in `if deepseek v4` branch - logits = F.linear(hidden_states, self.weight, None) + logits = linear_bf16_fp32(hidden_states, self.weight) return logits