From f57f8a8afd845ca5f03e229f41d891b727890fd3 Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Sat, 6 Jun 2026 14:14:25 +0800 Subject: [PATCH] Optimize Gemma4 H200 MoE and extend attention (#26588) Co-authored-by: BBuf --- .../attention/triton_ops/extend_attention.py | 4 +- python/sglang/srt/layers/gemma4_fused_ops.py | 136 +++++++++++++----- .../E=128,N=704,device_name=NVIDIA_H200.json | 114 +++++++++++++++ ...28,N=704,device_name=NVIDIA_H200_down.json | 114 +++++++++++++++ 4 files changed, 335 insertions(+), 33 deletions(-) create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200.json create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200_down.json diff --git a/python/sglang/srt/layers/attention/triton_ops/extend_attention.py b/python/sglang/srt/layers/attention/triton_ops/extend_attention.py index e6a353e9b..870ca88cd 100644 --- a/python/sglang/srt/layers/attention/triton_ops/extend_attention.py +++ b/python/sglang/srt/layers/attention/triton_ops/extend_attention.py @@ -82,8 +82,10 @@ def _get_block_sizes_for_extend_attention(Lq: int, Lv: int): BLOCK_M, BLOCK_N = (16, 64) elif _is_cuda and CUDA_CAPABILITY[0] >= 9: # Hopper architecture (H100, etc.) - if Lq <= 256: + if Lq <= 128: BLOCK_M, BLOCK_N = (128, 64) + elif Lq <= 256: + BLOCK_M, BLOCK_N = (64, 64) else: BLOCK_M, BLOCK_N = (32, 64) elif _is_cuda and CUDA_CAPABILITY[0] >= 8: diff --git a/python/sglang/srt/layers/gemma4_fused_ops.py b/python/sglang/srt/layers/gemma4_fused_ops.py index 11a0ddd5f..fc1902fdb 100644 --- a/python/sglang/srt/layers/gemma4_fused_ops.py +++ b/python/sglang/srt/layers/gemma4_fused_ops.py @@ -132,6 +132,29 @@ def _gemma_dual_rmsnorm_residual_kernel( tl.store(Out_ptr + row * stride_o + cols, out.to(x1.dtype), mask=mask) +@triton.jit +def _gemma_qkv_rmsnorm_store( + X_ptr, + W_ptr, + stride_m, + m, + h, + cols, + mask, + HEAD_DIM: tl.constexpr, + eps, + HAS_WEIGHT: tl.constexpr, +): + off = m * stride_m + h * HEAD_DIM + cols + x = tl.load(X_ptr + off, mask=mask, other=0.0).to(tl.float32) + rrms = tl.rsqrt(tl.sum(x * x, axis=0) / HEAD_DIM + eps) + out = x * rrms + if HAS_WEIGHT: + w = tl.load(W_ptr + cols, mask=mask, other=0.0).to(tl.float32) + out = out * w + tl.store(X_ptr + off, out.to(X_ptr.dtype.element_ty), mask=mask) + + @triton.jit def _gemma_qkv_rmsnorm_kernel( Q_ptr, @@ -147,48 +170,75 @@ def _gemma_qkv_rmsnorm_kernel( HEAD_DIM: tl.constexpr, eps, HAS_KV: tl.constexpr, + BY_HEAD: tl.constexpr, BLOCK: tl.constexpr, ): - """Per-token fused RMSNorm of Q (with q_w), K (with k_w), V (no scale). + """Fused per-head RMSNorm for Q, K, V. - Layout assumption: each tensor's last dim packs (num_heads, head_dim) contiguously - so per-head offset is `h * HEAD_DIM`. The token (M) stride is taken from - stride_*_m so the kernel works on strided views (e.g. slices of a larger - qkv buffer produced by `qkv.split`) without requiring `.contiguous()` copies. - V uses `weight=ones` semantics so the multiply-by-weight is omitted. + The same kernel supports two launch shapes: + - BY_HEAD=True: grid is (M, total_heads), one program per token/head. + - BY_HEAD=False: grid is (M,), one program per token looping over heads. + + Layout assumption: each tensor's last dim packs (num_heads, head_dim) + contiguously so per-head offset is `h * HEAD_DIM`. The token (M) stride is + taken from stride_*_m so the kernel works on strided views (e.g. slices of a + larger qkv buffer produced by `qkv.split`) without requiring `.contiguous()` + copies. V uses `weight=ones` semantics so the multiply-by-weight is omitted. """ m = tl.program_id(0) cols = tl.arange(0, BLOCK) mask = cols < HEAD_DIM - qw = tl.load(Q_w_ptr + cols, mask=mask, other=0.0).to(tl.float32) + if BY_HEAD: + h_all = tl.program_id(1) + if h_all < NUM_Q_HEADS: + _gemma_qkv_rmsnorm_store( + Q_ptr, Q_w_ptr, stride_q_m, m, h_all, cols, mask, HEAD_DIM, eps, True + ) + elif HAS_KV and h_all < NUM_Q_HEADS + NUM_KV_HEADS: + h = h_all - NUM_Q_HEADS + _gemma_qkv_rmsnorm_store( + K_ptr, K_w_ptr, stride_k_m, m, h, cols, mask, HEAD_DIM, eps, True + ) + elif HAS_KV: + h = h_all - NUM_Q_HEADS - NUM_KV_HEADS + _gemma_qkv_rmsnorm_store( + V_ptr, Q_w_ptr, stride_v_m, m, h, cols, mask, HEAD_DIM, eps, False + ) + else: + for h in tl.static_range(NUM_Q_HEADS): + _gemma_qkv_rmsnorm_store( + Q_ptr, Q_w_ptr, stride_q_m, m, h, cols, mask, HEAD_DIM, eps, True + ) - # Q heads - for h in tl.static_range(NUM_Q_HEADS): - off = m * stride_q_m + h * HEAD_DIM + cols - x = tl.load(Q_ptr + off, mask=mask, other=0.0).to(tl.float32) - rrms = tl.rsqrt(tl.sum(x * x, axis=0) / HEAD_DIM + eps) - out = x * rrms * qw - tl.store(Q_ptr + off, out.to(Q_ptr.dtype.element_ty), mask=mask) + if HAS_KV: + for h in tl.static_range(NUM_KV_HEADS): + _gemma_qkv_rmsnorm_store( + K_ptr, + K_w_ptr, + stride_k_m, + m, + h, + cols, + mask, + HEAD_DIM, + eps, + True, + ) - if HAS_KV: - kw = tl.load(K_w_ptr + cols, mask=mask, other=0.0).to(tl.float32) - - # K heads - for h in tl.static_range(NUM_KV_HEADS): - off = m * stride_k_m + h * HEAD_DIM + cols - x = tl.load(K_ptr + off, mask=mask, other=0.0).to(tl.float32) - rrms = tl.rsqrt(tl.sum(x * x, axis=0) / HEAD_DIM + eps) - out = x * rrms * kw - tl.store(K_ptr + off, out.to(K_ptr.dtype.element_ty), mask=mask) - - # V heads (no scaling: V-norm uses weight=ones) - for h in tl.static_range(NUM_KV_HEADS): - off = m * stride_v_m + h * HEAD_DIM + cols - x = tl.load(V_ptr + off, mask=mask, other=0.0).to(tl.float32) - rrms = tl.rsqrt(tl.sum(x * x, axis=0) / HEAD_DIM + eps) - out = x * rrms - tl.store(V_ptr + off, out.to(V_ptr.dtype.element_ty), mask=mask) + for h in tl.static_range(NUM_KV_HEADS): + _gemma_qkv_rmsnorm_store( + V_ptr, + Q_w_ptr, + stride_v_m, + m, + h, + cols, + mask, + HEAD_DIM, + eps, + False, + ) def gemma_qkv_rmsnorm( @@ -227,6 +277,27 @@ def gemma_qkv_rmsnorm( assert k.stride(-1) == 1 and v.stride(-1) == 1 assert k_weight is not None and k_weight.shape[-1] == head_dim + if M <= 256: + total_heads = num_q_heads + (2 * num_kv_heads if has_kv else 0) + _gemma_qkv_rmsnorm_kernel[(M, total_heads)]( + q, + k if has_kv else q, + v if has_kv else q, + q_weight, + k_weight if has_kv else q_weight, + q.stride(0), + k.stride(0) if has_kv else 0, + v.stride(0) if has_kv else 0, + NUM_Q_HEADS=num_q_heads, + NUM_KV_HEADS=num_kv_heads if has_kv else 0, + HEAD_DIM=head_dim, + eps=eps, + HAS_KV=has_kv, + BY_HEAD=True, + BLOCK=BLOCK, + ) + return + _gemma_qkv_rmsnorm_kernel[(M,)]( q, k if has_kv else q, @@ -241,6 +312,7 @@ def gemma_qkv_rmsnorm( HEAD_DIM=head_dim, eps=eps, HAS_KV=has_kv, + BY_HEAD=False, BLOCK=BLOCK, ) diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200.json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200.json new file mode 100644 index 000000000..f16acab33 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200.json @@ -0,0 +1,114 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 5 + }, + "2": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "4": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "8": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "16": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "24": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "32": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "48": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "64": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "96": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "128": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "256": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 2 + }, + "512": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "1024": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200_down.json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200_down.json new file mode 100644 index 000000000..f16acab33 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=128,N=704,device_name=NVIDIA_H200_down.json @@ -0,0 +1,114 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 5 + }, + "2": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "4": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "8": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "16": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "24": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "32": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "48": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "64": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "96": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "128": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 4 + }, + "256": { + "BLOCK_SIZE_M": 32, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 2 + }, + "512": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + }, + "1024": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 64, + "GROUP_SIZE_M": 1, + "num_warps": 4, + "num_stages": 3 + } +}