diff --git a/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py b/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py index b360d364a..509bd4156 100644 --- a/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py +++ b/python/sglang/srt/batch_invariant_ops/batch_invariant_ops.py @@ -171,7 +171,10 @@ def matmul_kernel_persistent( def _matmul_persistent_triton( - a: torch.Tensor, b: torch.Tensor, bias: torch.Tensor | None = None + a: torch.Tensor, + b: torch.Tensor, + out_dtype: torch.dtype, + bias: torch.Tensor | None = None, ): # Check constraints. assert a.shape[1] == b.shape[0], "Incompatible dimensions" @@ -182,9 +185,8 @@ def _matmul_persistent_triton( NUM_SMS = get_device_core_count() M, K = a.shape K, N = b.shape - dtype = a.dtype # Allocates output. - c = torch.empty((M, N), device=a.device, dtype=dtype) + c = torch.empty((M, N), device=a.device, dtype=out_dtype) # 1D launch kernel where each block gets its own program. def grid(META): @@ -242,18 +244,20 @@ def _matmul_persistent_triton( B_LARGE=b.numel() > 2**31, C_LARGE=c.numel() > 2**31, HAS_BIAS=bias is not None, - **configs[dtype], + **configs[a.dtype], ) return c def _matmul_persistent_deepgemm( - a: torch.Tensor, b: torch.Tensor, bias: torch.Tensor | None = None + a: torch.Tensor, + b: torch.Tensor, + out_dtype: torch.dtype, + bias: torch.Tensor | None = None, ): M, K = a.shape K, N = b.shape - dtype = a.dtype - out = torch.empty((M, N), device=a.device, dtype=dtype) + out = torch.empty((M, N), device=a.device, dtype=out_dtype) # DeepGEMM 0.2 defaults BF16 GEMMs to cuBLASLt, whose reduction can # depend on the batch size. Older wheels always use the invariant kernel. @@ -282,9 +286,13 @@ def _matmul_persistent_deepgemm( def matmul_persistent( - a: torch.Tensor, b: torch.Tensor, bias: torch.Tensor | None = None + a: torch.Tensor, + b: torch.Tensor, + bias: torch.Tensor | None = None, + out_dtype: torch.dtype | None = None, ): K, N = b.shape + out_dtype = out_dtype or a.dtype # DeepGEMM has minimum dimension requirements for TMA descriptors MIN_DEEPGEMM_DIM = 16 @@ -299,8 +307,12 @@ def matmul_persistent( and N >= MIN_DEEPGEMM_DIM ): if _ENABLE_MM_COMPARISON_TEST: - out_triton = _matmul_persistent_triton(a=a, b=b, bias=bias) - out_deepgemm = _matmul_persistent_deepgemm(a=a, b=b, bias=bias) + out_triton = _matmul_persistent_triton( + a=a, b=b, bias=bias, out_dtype=out_dtype + ) + out_deepgemm = _matmul_persistent_deepgemm( + a=a, b=b, bias=bias, out_dtype=out_dtype + ) diff = calc_diff(out_triton, out_deepgemm) assert diff < 0.0001, f"{diff=} {out_triton=} {out_deepgemm=}" # can be enabled for debugging @@ -313,15 +325,15 @@ def matmul_persistent( # print(f"{a=} {b=} {bias=} {out_triton=} {out_deepgemm=}") return out_deepgemm - return _matmul_persistent_deepgemm(a=a, b=b, bias=bias) + return _matmul_persistent_deepgemm(a=a, b=b, bias=bias, out_dtype=out_dtype) - if _ENABLE_MM_FALLBACK_VARIANT: + if _ENABLE_MM_FALLBACK_VARIANT and out_dtype == a.dtype: out = torch.einsum("ik,kj->ij", a, b) if bias is not None: out += bias return out - return _matmul_persistent_triton(a=a, b=b, bias=bias) + return _matmul_persistent_triton(a=a, b=b, bias=bias, out_dtype=out_dtype) @triton.jit @@ -974,7 +986,7 @@ def _rms_norm_aten_compat(input, normalized_shape, weight=None, eps=None): def _mm_dtype_compat(self, mat2, out_dtype): - return matmul_persistent(self.contiguous(), mat2.contiguous()).to(out_dtype) + return matmul_persistent(self.contiguous(), mat2.contiguous(), out_dtype=out_dtype) _batch_invariant_MODE = False diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 41148eb40..bae009ce8 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -518,7 +518,9 @@ class MoEGate(nn.Module): ) if get_exec().deterministic.enable_deterministic_inference: - return F.linear(hidden_states, self.weight, None) + if _is_cuda or _is_hip: + return torch.mm(hidden_states, self.weight.t(), out_dtype=torch.float32) + return F.linear(hidden_states.float(), self.weight.float(), None) if hidden_states.shape[0] <= self.tiny_router_gemm_max_tokens: logits = tiny_gemm_bf16(