[Blackwell] Reserve SMs for DeepGEMM MegaMoE grid barriers (#36657)

Co-authored-by: weireweire <20922698+weireweire@users.noreply.github.com>
This commit is contained in:
weireweire
2026-08-28 13:37:36 -07:00
committed by GitHub
co-authored by weireweire
parent 70088aa5db
commit 3254f9b47c
3 changed files with 67 additions and 21 deletions
+3
View File
@@ -1070,6 +1070,9 @@ class Envs:
# DeepGEMM Mega MoE
# ===================================================================
SGLANG_OPT_DEEPGEMM_MEGA_MOE_NUM_MAX_TOKENS_PER_RANK = EnvInt(8192)
# Blackwell MegaMoE uses a whole-grid software barrier. Keep a small
# residency margin so every cluster can launch beside other streams.
SGLANG_OPT_DEEPGEMM_MEGA_MOE_RESERVED_SMS = EnvInt(2)
# ===================================================================
# Top-k kernels
+50 -11
View File
@@ -15,8 +15,9 @@
from __future__ import annotations
import functools
import os
from contextlib import nullcontext
from contextlib import contextmanager, nullcontext
from typing import TYPE_CHECKING, Optional
import torch
@@ -44,6 +45,43 @@ if TYPE_CHECKING:
_MEGA_MOE_SYMM_BUFFER: dict = {}
@functools.lru_cache(maxsize=1)
def _mega_moe_max_num_sms() -> Optional[int]:
if _device_sm < 100:
# The SM90 MegaMoE implementation does not use the whole-grid clustered
# launch that needs a residency margin.
return None
# Physical count, not deep_gemm.get_num_sms(): two-batch overlap and the DSA
# indexer reconfigure that process-wide, so reserving on top would compound.
num_sms = torch.cuda.get_device_properties(device="cuda").multi_processor_count
reserved_num_sms = max(envs.SGLANG_OPT_DEEPGEMM_MEGA_MOE_RESERVED_SMS.get(), 0)
return max(2, num_sms - reserved_num_sms)
@contextmanager
def _configure_mega_moe_deep_gemm_num_sms(deep_gemm):
max_num_sms = _mega_moe_max_num_sms()
if max_num_sms is None:
yield
return
current_num_sms = deep_gemm.get_num_sms()
# Stay under an outer context's budget instead of claiming SMs back from it.
target_num_sms = min(max_num_sms, current_num_sms)
# Round down: the clustered launch needs an even CTA count.
target_num_sms -= target_num_sms % 2
if target_num_sms == current_num_sms:
yield
return
deep_gemm.set_num_sms(target_num_sms)
try:
yield
finally:
deep_gemm.set_num_sms(current_num_sms)
def _get_mega_moe_symm_buffer(
group,
num_experts: int,
@@ -247,16 +285,17 @@ def _run_mega_routed(
device=hidden_states.device,
)
swiglu_limit = getattr(moe.config, "swiglu_limit", None)
deep_gemm.fp8_fp4_mega_moe(
y,
moe.experts.mega_l1_weights,
moe.experts.mega_l2_weights,
buf,
recipe=(1, 1, 32),
activation="swiglu",
activation_clamp=swiglu_limit,
fast_math=True,
)
with _configure_mega_moe_deep_gemm_num_sms(deep_gemm):
deep_gemm.fp8_fp4_mega_moe(
y,
moe.experts.mega_l1_weights,
moe.experts.mega_l2_weights,
buf,
recipe=(1, 1, 32),
activation="swiglu",
activation_clamp=swiglu_limit,
fast_math=True,
)
y = y[:num_tokens]
if not moe.experts.should_fuse_routed_scaling_factor_in_topk:
+14 -10
View File
@@ -780,7 +780,10 @@ class KimiK3MoE(nn.Module):
from sglang.kernels.ops.attention.dsv4 import mega_moe_pre_dispatch
from sglang.srt.distributed.parallel_state import get_moe_ep_group
from sglang.srt.environ import envs
from sglang.srt.layers.moe.mega_moe import _get_mega_moe_symm_buffer
from sglang.srt.layers.moe.mega_moe import (
_configure_mega_moe_deep_gemm_num_sms,
_get_mega_moe_symm_buffer,
)
# In SP-MoE mode (KimiK3DecoderLayer reduce-scatters the o_proj
# output) the incoming rows are already this rank's token shard, so
@@ -833,15 +836,16 @@ class KimiK3MoE(nn.Module):
dtype=torch.bfloat16,
device=routed_input.device,
)
deep_gemm.fp8_fp4_mega_moe(
y,
self.experts.mega_l1_weights,
self.experts.mega_l2_weights,
buf,
recipe=(1, 1, 32),
activation="situ",
fast_math=True,
)
with _configure_mega_moe_deep_gemm_num_sms(deep_gemm):
deep_gemm.fp8_fp4_mega_moe(
y,
self.experts.mega_l1_weights,
self.experts.mega_l2_weights,
buf,
recipe=(1, 1, 32),
activation="situ",
fast_math=True,
)
y = y[:num_tokens]
if not self.experts.should_fuse_routed_scaling_factor_in_topk:
if (