diff --git a/python/sglang/kernels/ops/layernorm/hc_combine_norm.py b/python/sglang/kernels/ops/layernorm/hc_combine_norm.py index ab5018d45..a52c3e17a 100644 --- a/python/sglang/kernels/ops/layernorm/hc_combine_norm.py +++ b/python/sglang/kernels/ops/layernorm/hc_combine_norm.py @@ -8,7 +8,16 @@ from sglang.kernels.ops.layernorm.mxfp8_epilogue import mxfp8_epilogue @triton.jit -def _hc_combine_norm(X, P, W, Y, SX: tl.constexpr, SP: tl.constexpr, EPS: tl.constexpr): +def _hc_combine_norm( + X, + P, + W, + Y, + SX: tl.constexpr, + SP: tl.constexpr, + EPS: tl.constexpr, + PARTS: tl.constexpr, +): row, part = tl.program_id(0), tl.program_id(1) h = tl.arange(0, 8192) value = tl.full((8192,), 0, tl.float32) @@ -19,7 +28,7 @@ def _hc_combine_norm(X, P, W, Y, SX: tl.constexpr, SP: tl.constexpr, EPS: tl.con # The unfused combine stores BF16 before RMSNorm reads it. value = value.to(tl.bfloat16).to(tl.float32) inv_rms = tl.rsqrt(tl.sum(value * value, 0) / 5120 + EPS) - mask = (h >= part * 1280) & (h < (part + 1) * 1280) + mask = (h >= part * (5120 // PARTS)) & (h < (part + 1) * (5120 // PARTS)) weight = tl.load(W + h, mask, 0).to(tl.float32) tl.store(Y + row * 5120 + h, value * inv_rms * weight, mask) @@ -48,7 +57,7 @@ def hc_combine_norm( ) -> torch.Tensor: """Fuse four-stream combine and RMSNorm for BF16 batches of width 5120.""" m = x.shape[0] - assert (0 < m <= 8 or 4096 <= m <= 65536) and x.shape == (m, 20480) + assert (0 < m <= 96 or 4096 <= m <= 65536) and x.shape == (m, 20480) assert pre.shape == (m, 4) and pre.stride(1) == 1 assert weight.shape == (5120,) and weight.is_contiguous() assert x.dtype == weight.dtype == torch.bfloat16 and x.stride(1) == 1 @@ -59,9 +68,11 @@ def hc_combine_norm( ) return y # Four CTAs per row trade redundant statistics for more concurrent loads - # when only a few speculative tokens are being processed. - _hc_combine_norm[(m, 4)]( - x, pre, weight, y, x.stride(0), pre.stride(0), eps, num_warps=8 + # when only a few speculative tokens are being processed; wider batches have + # enough rows to split the 5120 columns fewer ways. + parts = 4 if m <= 8 else (2 if m <= 48 else 1) + _hc_combine_norm[(m, parts)]( + x, pre, weight, y, x.stride(0), pre.stride(0), eps, parts, num_warps=8 ) return y diff --git a/python/sglang/srt/models/deepseek_v4.py b/python/sglang/srt/models/deepseek_v4.py index d46f0f90e..a81cfa693 100644 --- a/python/sglang/srt/models/deepseek_v4.py +++ b/python/sglang/srt/models/deepseek_v4.py @@ -3177,7 +3177,7 @@ class DeepseekV4DecoderLayer(nn.Module): x.is_cuda and get_platform().is_blackwell and ( - 0 < x.shape[0] <= 8 + 0 < x.shape[0] <= 96 or ( self.config.model_type == "deepseek_v41" and 4096 <= x.shape[0] <= 65536