keeping router GEMM in fp32 for deterministic inference (DeepSeek V3/V4) (#38176)

Co-authored-by: Brayden Zhong <b8zhong@uwaterloo.ca>
Co-authored-by: Brayden Zhong <brayden@radixark.ai>
This commit is contained in:
Carrie Chen
2026-09-15 16:16:09 +08:00
committed by GitHub
co-authored by Brayden Zhong Brayden Zhong
parent 37ebacb50f
commit 4f52c9948b
2 changed files with 29 additions and 15 deletions
@@ -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
+3 -1
View File
@@ -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(