Files
sglang/python/sglang/srt/layers/moe/moe_runner/deep_gemm.py
T
2026-08-26 19:54:33 -07:00

1613 lines
58 KiB
Python

from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any, List, Optional, Tuple
import torch
import triton
import triton.language as tl
from sglang.kernels.ops.attention.dsv4 import silu_and_mul_masked_post_quant
from sglang.kernels.ops.quantization import per_token_group_quant
logger = logging.getLogger(__name__)
from sglang.srt.distributed import get_tp_group
from sglang.srt.distributed.device_communicators.pynccl_allocator import (
use_symmetric_memory,
)
from sglang.srt.environ import envs
from sglang.srt.layers import deep_gemm_wrapper
from sglang.srt.layers.dp_attention import is_allocation_symmetric
from sglang.srt.layers.moe.moe_runner.base import (
MoeQuantInfo,
MoeRunnerConfig,
MoeRunnerCore,
RunnerInput,
RunnerOutput,
register_post_permute,
register_pre_permute,
)
from sglang.srt.layers.moe.utils import MoeRunnerBackend, get_moe_a2a_backend
from sglang.srt.runtime_context import get_exec
from sglang.srt.utils import (
ceil_div,
dispose_tensor,
get_bool_env_var,
is_cuda,
is_hip,
is_musa,
is_npu,
)
from sglang.srt.utils.offloader import get_offloader
if TYPE_CHECKING:
from sglang.srt.layers.moe.token_dispatcher.deepep import (
DeepEPLLCombineInput,
DeepEPLLDispatchOutput,
DeepEPNormalCombineInput,
DeepEPNormalDispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import (
DeepEPv2CombineInput,
DeepEPv2DispatchOutput,
)
from sglang.srt.layers.moe.token_dispatcher.standard import (
StandardCombineInput,
StandardDispatchOutput,
)
_is_hip = is_hip()
_is_npu = is_npu()
_is_cuda = is_cuda()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
_is_musa = is_musa()
if not (_is_npu or _is_hip) and _is_cuda:
from sglang.kernels.ops.activation.activation import (
silu_and_mul as _legacy_silu_and_mul,
)
elif _is_musa:
_silu_and_mul_musa = torch.nn.SwishGLU()
else:
_legacy_silu_and_mul = None
_DEEPGEMM_ON_H20 = get_bool_env_var("SGLANG_DEEPGEMM_ON_H20")
_masked_standard_layout_memory_budget_bytes: Optional[int] = None
# TODO(kaixih@nvidia): ideally we should merge this logic into
# `fill_gateup_input_triton_kernel` to directly generate e8m0 scale.
@torch.compile(disable=_is_hip or _is_npu)
def _cast_to_e8m0_with_rounding_up(x: torch.Tensor) -> torch.Tensor:
temp = x.to(torch.float32).view(torch.int32)
exp = torch.bitwise_right_shift(temp, 23)
mant = torch.bitwise_and(temp, 0x7FFFFF)
is_ru = torch.logical_and(
torch.logical_and((mant > 0), (exp != 0xFE)),
~torch.logical_and((exp == 0), (mant <= 0x400000)),
)
exp = torch.where(is_ru, exp + 1, exp)
new_x = exp.to(torch.uint8).view(torch.int)
return new_x.transpose(1, 2).contiguous().transpose(1, 2)
def copy_list_to_gpu_no_ce(arr: List[int]):
from sgl_kernel.elementwise import copy_to_gpu_no_ce
tensor_cpu = torch.tensor(arr, dtype=torch.int32, device="cpu")
tensor_gpu = torch.empty_like(tensor_cpu, device="cuda")
copy_to_gpu_no_ce(tensor_cpu, tensor_gpu)
return tensor_gpu
def set_masked_standard_layout_memory_budget(
available_memory_bytes: int,
) -> int:
"""Cache the masked-layout share of free non-static device memory."""
global _masked_standard_layout_memory_budget_bytes
fraction = envs.SGLANG_DEEPGEMM_MASKED_MEMORY_BUDGET_FRACTION.get()
if not 0.0 < fraction <= 1.0:
raise ValueError(
"SGLANG_DEEPGEMM_MASKED_MEMORY_BUDGET_FRACTION must be in (0, 1]"
)
_masked_standard_layout_memory_budget_bytes = int(available_memory_bytes * fraction)
return _masked_standard_layout_memory_budget_bytes
def _estimate_masked_standard_layout_peak_bytes(
runner_config: MoeRunnerConfig,
quant_info: DeepGemmMoeQuantInfo,
hidden_states: torch.Tensor,
) -> int:
padded_m = (hidden_states.shape[0] // 256 + 1) * 256
activation_dtype = (
torch.bfloat16
if quant_info.w13_weight.dtype == torch.bfloat16
else torch.float8_e4m3fn
)
hidden_size = hidden_states.shape[1]
gateup_size = quant_info.w13_weight.shape[1]
gateup_row_bytes = gateup_size * torch.bfloat16.itemsize
down_output_row_bytes = quant_info.w2_weight.shape[1] * torch.bfloat16.itemsize
input_row_bytes = hidden_size * activation_dtype.itemsize
down_input_row_bytes = gateup_size // 2 * activation_dtype.itemsize
if activation_dtype == torch.bfloat16:
input_scale_row_bytes = 0
down_scale_row_bytes = 0
else:
block_k = quant_info.block_shape[1] if quant_info.block_shape else 128
packed_scales = quant_info.use_mxfp8 or deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
scale_item_bytes = (
torch.uint8.itemsize if packed_scales else torch.float32.itemsize
)
input_scale_row_bytes = ceil_div(hidden_size, block_k) * scale_item_bytes
down_scale_row_bytes = ceil_div(gateup_size // 2, block_k) * scale_item_bytes
peak_row_bytes = max(
input_row_bytes + input_scale_row_bytes + gateup_row_bytes,
gateup_row_bytes + down_input_row_bytes + down_scale_row_bytes,
down_input_row_bytes + down_scale_row_bytes + down_output_row_bytes,
)
return runner_config.num_local_experts * padded_m * peak_row_bytes
def _should_use_masked_standard_layout(
runner_config: MoeRunnerConfig,
quant_info: DeepGemmMoeQuantInfo,
hidden_states: torch.Tensor,
) -> bool:
mode = envs.SGLANG_DEEPGEMM_STANDARD_LAYOUT.get().lower()
if mode not in ("auto", "masked", "compact"):
raise ValueError(
"SGLANG_DEEPGEMM_STANDARD_LAYOUT must be one of: auto, masked, compact"
)
if mode != "auto":
return mode == "masked"
global _masked_standard_layout_memory_budget_bytes
if _masked_standard_layout_memory_budget_bytes is None:
# Serving sets an all-rank budget before capture. Direct eager callers
# fall back to this rank's free memory without querying inside capture.
# Import lazily to avoid a module-initialization cycle through
# runner_utils -> DeepEP -> MoE -> this module.
from sglang.srt.model_executor.runner_utils.capture_mode import (
get_is_capture_mode,
)
if get_is_capture_mode():
return False
free_memory, _ = torch.cuda.mem_get_info(hidden_states.device)
set_masked_standard_layout_memory_budget(free_memory)
return (
_estimate_masked_standard_layout_peak_bytes(
runner_config, quant_info, hidden_states
)
<= _masked_standard_layout_memory_budget_bytes
)
def _get_compact_all_tokens(
num_assignments: int, num_experts: int, block_e: int = 128
) -> int:
"""Return the maximum padded rows over all routings of the assignments."""
max_nonempty_experts = min(num_assignments, num_experts)
return block_e * (
max_nonempty_experts + (num_assignments - max_nonempty_experts) // block_e
)
@dataclass
class DeepGemmRunnerInput(RunnerInput):
hidden_states: torch.Tensor
hidden_states_scale: torch.Tensor
use_masked_gemm: bool
masked_m: Optional[torch.Tensor] = None
expected_m: Optional[int] = None
m_indices: Optional[torch.Tensor] = None
hidden_states_scale_tma_aligned: bool = False
@property
def runner_backend(self) -> MoeRunnerBackend:
return MoeRunnerBackend.DEEP_GEMM
@dataclass
class DeepGemmRunnerOutput(RunnerOutput):
hidden_states: torch.Tensor
@property
def runner_backend(self) -> MoeRunnerBackend:
return MoeRunnerBackend.DEEP_GEMM
@dataclass
class DeepGemmMoeQuantInfo(MoeQuantInfo):
w13_weight: torch.Tensor
w2_weight: torch.Tensor
use_fp8: bool
w13_scale: Optional[torch.Tensor] = None
w2_scale: Optional[torch.Tensor] = None
block_shape: Optional[List[int]] = None
# DSV4 mxfp4 layout flag; selects recipe_a=(1,128)/recipe_b=(1,32) downstream.
is_fp4_experts: bool = False
use_mxfp8: bool = False
def __post_init__(self):
if self.use_mxfp8:
assert self.block_shape == [
1,
32,
], f"MXFP8 requires block_shape [1, 32], got {self.block_shape}"
assert (
deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
), "MXFP8 requires DEEPGEMM_SCALE_UE8M0=True"
class DeepGemmRunnerCore(MoeRunnerCore):
def __init__(self, config: MoeRunnerConfig):
super().__init__(config)
# SiTU (Kimi K3) is applied outside the GEMMs in python, so it only
# needs the masked-gemm activation site to branch (see _run_masked_gemm).
assert self.config.activation in ("silu", "situ")
assert self.config.is_gated
self.swiglu_limit = self.config.swiglu_limit
self.use_swizzle = get_moe_a2a_backend().is_megamoe()
def run(
self,
runner_input: DeepGemmRunnerInput,
quant_info: DeepGemmMoeQuantInfo,
running_state: dict,
hooks: Optional[Any] = None,
) -> DeepGemmRunnerOutput:
weight_dtype = quant_info.w13_weight.dtype
if not runner_input.use_masked_gemm:
if weight_dtype == torch.bfloat16:
hidden_states = self._run_bf16_contiguous_gemm(
runner_input, quant_info, running_state
)
else:
hidden_states = self._run_contiguous_gemm(
runner_input, quant_info, running_state
)
else:
if weight_dtype == torch.bfloat16:
hidden_states = self._run_masked_bf16_gemm(
runner_input, quant_info, running_state
)
else:
hidden_states = self._run_masked_gemm(
runner_input, quant_info, running_state
)
return DeepGemmRunnerOutput(hidden_states=hidden_states)
def _run_contiguous_gemm(
self,
runner_input: DeepGemmRunnerInput,
quant_info: DeepGemmMoeQuantInfo,
running_state: dict,
) -> torch.Tensor:
from sglang.kernels.ops.attention.dsv4 import silu_and_mul_contig_post_quant
from sglang.kernels.ops.moe.ep_moe_kernels import tma_align_input_scale
from sglang.kernels.ops.quantization.fp8_kernel import (
create_per_token_group_quant_fp8_output_scale,
)
hidden_states = runner_input.hidden_states
hidden_states_scale = runner_input.hidden_states_scale
all_tokens = running_state["all_tokens"]
hidden_states_device = running_state["hidden_states_device"]
hidden_states_dtype = running_state["hidden_states_dtype"]
hidden_states_shape = running_state["hidden_states_shape"]
m_indices = runner_input.m_indices
N = quant_info.w13_weight.size(1)
K = hidden_states_shape[1]
scale_block_size = 128
recipe_a, recipe_b = (
((1, 128), (1, 32)) if quant_info.is_fp4_experts else (None, None)
)
w13_weight_fp8 = (
quant_info.w13_weight,
quant_info.w13_scale,
)
w2_weight_fp8 = (quant_info.w2_weight, quant_info.w2_scale)
gateup_output = torch.empty(
(all_tokens, N),
device=hidden_states_device,
dtype=torch.bfloat16,
)
if (
deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES
and not runner_input.hidden_states_scale_tma_aligned
):
hidden_states_scale = tma_align_input_scale(hidden_states_scale)
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_contig(
(hidden_states, hidden_states_scale),
w13_weight_fp8,
gateup_output,
m_indices,
recipe_a=recipe_a,
recipe_b=recipe_b,
)
dispose_tensor(hidden_states)
dispose_tensor(hidden_states_scale)
if self.config.activation == "situ":
situ_beta = self.config.gemm1_alpha
situ_linear_beta = self.config.gemm1_clamp_limit
assert situ_beta is not None and situ_linear_beta is not None
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
# Fused SiTU + per-group fp8 quant over the compacted rows,
# then the proven round-up e8m0 cast (mn-major packed layout).
rows = gateup_output.shape[0]
half_n = N // 2
kg = half_n // scale_block_size
down_input_fp8 = torch.empty(
(rows, half_n),
device=gateup_output.device,
dtype=torch.float8_e4m3fn,
)
s = torch.empty(
(rows, kg), device=gateup_output.device, dtype=torch.float32
)
_situ_mul_quant_contig_kernel[(rows,)](
gateup_output,
down_input_fp8,
s,
half_n,
kg,
situ_beta,
situ_linear_beta,
GROUP=scale_block_size,
KG_POW2=triton.next_power_of_2(kg),
num_warps=8,
)
del gateup_output
down_input_scale = _cast_to_e8m0_with_rounding_up(
s.unsqueeze(0)
).squeeze(0)
else:
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
gate = gateup_output[:, : N // 2].float()
up = gateup_output[:, N // 2 :].float()
gate = situ_beta * torch.tanh(gate / situ_beta) * torch.sigmoid(gate)
up = situ_linear_beta * torch.tanh(up / situ_linear_beta)
down_input = (gate * up).to(torch.bfloat16)
del gateup_output
down_input_fp8, down_input_scale = sglang_per_token_group_quant_fp8(
down_input,
scale_block_size,
column_major_scales=False,
scale_tma_aligned=False,
scale_ue8m0=False,
)
del down_input
elif self.use_swizzle:
swiglu_limit_arg: Optional[float] = self.swiglu_limit
down_input_fp8 = torch.empty(
(all_tokens, N // 2),
device=gateup_output.device,
dtype=torch.float8_e4m3fn,
)
down_input_scale = create_per_token_group_quant_fp8_output_scale(
x_shape=(all_tokens, N // 2),
device=gateup_output.device,
group_size=scale_block_size,
column_major_scales=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
scale_tma_aligned=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
silu_and_mul_contig_post_quant(
input=gateup_output,
output=down_input_fp8,
output_scale=down_input_scale,
quant_group_size=scale_block_size,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
transposed=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
swiglu_limit=swiglu_limit_arg,
swizzle=self.use_swizzle,
)
del gateup_output
else:
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
if self.swiglu_limit is not None:
gateup_output = _apply_swiglu_limit(
gateup_output, swiglu_limit=self.swiglu_limit
)
if not _is_musa:
down_input = torch.empty(
(all_tokens, N // 2),
device=gateup_output.device,
dtype=torch.bfloat16,
)
_legacy_silu_and_mul(gateup_output.view(-1, N), down_input)
else:
down_input = _silu_and_mul_musa(gateup_output.view(-1, N))
del gateup_output
down_input_fp8, down_input_scale = sglang_per_token_group_quant_fp8(
down_input,
scale_block_size,
column_major_scales=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
scale_tma_aligned=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
del down_input
# Allocate the MoE output in the NCCL symmetric memory pool when symmetric
# allocation is required, so the downstream all-reduce takes the low-latency
# symmetric path. Only this final output enters the pool; intermediate
# buffers stay on the default allocator to bound pool occupancy.
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
down_output = torch.empty(
(all_tokens, K),
device=hidden_states_device,
dtype=torch.bfloat16,
)
if deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES:
down_input_scale = tma_align_input_scale(down_input_scale)
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_contig(
(down_input_fp8, down_input_scale),
w2_weight_fp8,
down_output,
m_indices,
recipe_a=recipe_a,
recipe_b=recipe_b,
)
return down_output
def _run_bf16_contiguous_gemm(
self,
runner_input: DeepGemmRunnerInput,
quant_info: DeepGemmMoeQuantInfo,
running_state: dict,
) -> torch.Tensor:
hidden_states = runner_input.hidden_states
all_tokens = running_state["all_tokens"]
hidden_states_device = running_state["hidden_states_device"]
hidden_states_shape = running_state["hidden_states_shape"]
m_indices = runner_input.m_indices
N = quant_info.w13_weight.size(1)
K = hidden_states_shape[1]
w13_weight = quant_info.w13_weight
w2_weight = quant_info.w2_weight
# GroupGemm-1: (M, K) (E, N, K) -> (M, N)
gateup_output = torch.empty(
(all_tokens, N),
device=hidden_states_device,
dtype=torch.bfloat16,
)
deep_gemm_wrapper.grouped_gemm_nt_bf16_contig(
hidden_states,
w13_weight,
gateup_output,
m_indices,
)
dispose_tensor(hidden_states)
# Act: (M, N) -> (M, N/2)
if not _is_musa:
down_input = torch.empty(
(
all_tokens,
N // 2,
),
device=gateup_output.device,
dtype=torch.bfloat16,
)
_legacy_silu_and_mul(gateup_output.view(-1, N), down_input)
else:
down_input = _silu_and_mul_musa(gateup_output.view(-1, N))
del gateup_output
# GroupGemm-2: (M, N/2) (E, K, N/2) -> (M, K)
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
down_output = torch.empty(
(all_tokens, K),
device=hidden_states_device,
dtype=torch.bfloat16,
)
deep_gemm_wrapper.grouped_gemm_nt_bf16_contig(
down_input,
w2_weight,
down_output,
m_indices,
)
return down_output
def _run_masked_gemm(
self,
runner_input: DeepGemmRunnerInput,
quant_info: DeepGemmMoeQuantInfo,
running_state: dict,
) -> torch.Tensor:
from sglang.srt.layers import deep_gemm_wrapper
hidden_states = runner_input.hidden_states
hidden_states_scale = runner_input.hidden_states_scale
masked_m = runner_input.masked_m
expected_m = runner_input.expected_m
w13_weight = quant_info.w13_weight
w2_weight = quant_info.w2_weight
w13_scale = quant_info.w13_scale
w2_scale = quant_info.w2_scale
hidden_states_device = running_state["hidden_states_device"]
use_mxfp8 = quant_info.use_mxfp8
scale_block_size = quant_info.block_shape[1] if quant_info.block_shape else 128
if use_mxfp8:
recipe_b = tuple(quant_info.block_shape)
# gran_k is set by the dispatch path (standard=block_shape[1], DeepEP-LL=128),
# not inferable from K; inferring it silently mis-reads the activation scale.
gran_k_act = running_state.get(
"mxfp8_act_gran_k", quant_info.block_shape[1]
)
_, _, k_for_recipe = hidden_states.shape
act_sf_last = hidden_states_scale.shape[-1]
assert ceil_div(k_for_recipe, gran_k_act * 4) == act_sf_last, (
f"MXFP8 gateup scale mismatch: gran_k={gran_k_act}, K={k_for_recipe}, "
f"act_sf_last={act_sf_last}, expected "
f"{ceil_div(k_for_recipe, gran_k_act * 4)}"
)
recipe_a = (quant_info.block_shape[0], gran_k_act)
elif quant_info.is_fp4_experts:
recipe_a, recipe_b = (1, 128), (1, 32)
else:
recipe_a, recipe_b = None, None
# GroupGemm-0
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
if hidden_states_scale.dtype != torch.int:
b, s_mn, s_k = hidden_states_scale.shape
assert (
s_mn % 4 == 0 and s_k % 4 == 0
), f"scales must be aligned to 4, but got ({b}, {s_mn}, {s_k})"
hidden_states_scale = _cast_to_e8m0_with_rounding_up(
hidden_states_scale
)
elif deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES:
hidden_states_scale = deep_gemm_wrapper.get_mn_major_tma_aligned_tensor(
hidden_states_scale
)
num_groups, m, k = hidden_states.shape
n = w13_weight.size(1)
try:
gateup_output = torch.empty(
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
)
except torch.OutOfMemoryError:
logger.error(
"Masked grouped-GEMM workspace allocation failed "
"(num_groups=%d m=%d n=%d). If this happens under saturated "
"dp-attention prefill, try SGLANG_OPT_DG_MASKED_M_CAP=1.",
num_groups,
m,
n,
)
raise
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked(
(hidden_states, hidden_states_scale),
(w13_weight, w13_scale),
gateup_output,
masked_m,
expected_m,
recipe_a=recipe_a,
recipe_b=recipe_b,
)
dispose_tensor(hidden_states)
dispose_tensor(hidden_states_scale)
swiglu_limit_arg: Optional[float] = None
if self.swiglu_limit is not None:
swiglu_limit_arg = self.swiglu_limit
# Act.
if self.config.activation == "situ":
down_input, down_input_scale = _varlen_deep_gemm_situ_mul_quant(
gateup_output,
masked_m,
group_size=128,
topk=self.config.top_k,
beta=self.config.gemm1_alpha,
linear_beta=self.config.gemm1_clamp_limit,
)
else:
topk_ids_rs = running_state.get("topk_ids")
num_real_tokens = (
topk_ids_rs.shape[0]
if (
use_mxfp8
and deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
and topk_ids_rs is not None
and "src2dst" in running_state
)
else None
)
down_input, down_input_scale = _varlen_deep_gemm_silu_mul_quant(
gateup_output,
masked_m,
group_size=scale_block_size,
topk=self.config.top_k,
swiglu_limit=swiglu_limit_arg,
swizzle=self.use_swizzle,
gemm1_alpha=self.config.gemm1_alpha,
gemm1_clamp_limit=self.config.gemm1_clamp_limit,
num_real_tokens=num_real_tokens,
)
del gateup_output
# Down activation is quantised locally at scale_block_size (never DeepEP-LL),
# so its gran_k differs from gateup recipe_a.
recipe_a_down = recipe_a
if use_mxfp8:
recipe_a_down = (quant_info.block_shape[0], scale_block_size)
# GroupGemm-1
n = w2_weight.shape[1]
if (
use_mxfp8
and deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
and down_input_scale.dtype != torch.int32
):
import deep_gemm.utils.layout
down_input_scale = (
deep_gemm.utils.layout.get_mn_major_tma_aligned_packed_ue8m0_tensor(
down_input_scale
)
)
elif deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES:
down_input_scale = deep_gemm_wrapper.get_mn_major_tma_aligned_tensor(
down_input_scale
)
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
down_output = torch.empty(
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
)
down_gemm_overlap_args = running_state.get("down_gemm_overlap_args", None)
if down_gemm_overlap_args is None:
gemm_overlap_args_dict = {}
else:
down_gemm_overlap_args.start_event.record()
max_block_n = (
160 if (_DEEPGEMM_ON_H20 and runner_input.expected_m <= 64) else 256
)
gemm_overlap_args_dict = {
"overlap_args": down_gemm_overlap_args,
"max_block_n": max_block_n,
}
deep_gemm_return_value = deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_masked(
(down_input, down_input_scale),
(w2_weight, w2_scale),
down_output,
masked_m,
expected_m,
recipe_a=recipe_a_down,
recipe_b=recipe_b,
**gemm_overlap_args_dict,
)
meta_overlap_args = running_state.get("meta_overlap_args", None)
# Returns (block_m, threshold) only with down-gemm overlap, else None;
# meta_overlap_args may be set without overlap, so guard the unpack.
if meta_overlap_args is not None and deep_gemm_return_value is not None:
block_m, threshold = deep_gemm_return_value
meta_overlap_args["block_m"] = block_m
meta_overlap_args["threshold"] = threshold
return down_output
def _run_masked_bf16_gemm(
self,
runner_input: DeepGemmRunnerInput,
quant_info: DeepGemmMoeQuantInfo,
running_state: dict,
) -> torch.Tensor:
from sglang.kernels.ops.moe.ep_moe_kernels import silu_and_mul_masked_fwd
from sglang.srt.layers import deep_gemm_wrapper
hidden_states = runner_input.hidden_states
masked_m = runner_input.masked_m
expected_m = runner_input.expected_m
w13_weight = quant_info.w13_weight
w2_weight = quant_info.w2_weight
hidden_states_device = running_state["hidden_states_device"]
# GroupGemm-0
num_groups, m, k = hidden_states.shape
n = w13_weight.size(1)
gateup_output = torch.empty(
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
)
deep_gemm_wrapper.grouped_gemm_nt_bf16_masked(
hidden_states,
w13_weight,
gateup_output,
masked_m,
expected_m,
)
dispose_tensor(hidden_states)
down_input = torch.empty(
(
gateup_output.shape[0],
gateup_output.shape[1],
gateup_output.shape[2] // 2,
),
device=hidden_states_device,
dtype=torch.bfloat16,
)
# Act
silu_and_mul_masked_fwd(gateup_output, down_input, masked_m)
del gateup_output
# GroupGemm-1
n = w2_weight.shape[1]
with use_symmetric_memory(
get_tp_group(), disabled=not is_allocation_symmetric()
):
down_output = torch.empty(
(num_groups, m, n), device=hidden_states_device, dtype=torch.bfloat16
)
deep_gemm_wrapper.grouped_gemm_nt_bf16_masked(
down_input,
w2_weight,
down_output,
masked_m,
expected_m,
)
# Note: BF16 masked gemm doesn't support overlap_args, so no return value unpack
return down_output
@property
def runner_backend(self) -> MoeRunnerBackend:
return MoeRunnerBackend.DEEP_GEMM
@register_pre_permute("standard", "deep_gemm")
def pre_permute_standard_to_deep_gemm(
dispatch_output: StandardDispatchOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepGemmRunnerInput:
from sglang.kernels.ops.moe.ep_moe_kernels import (
ep_scatter,
fused_moe_dispatch_index,
moe_ep_deepgemm_preprocess,
)
hidden_states, topk_output = (
dispatch_output.hidden_states,
dispatch_output.topk_output,
)
topk_weights, topk_ids, _ = topk_output
hidden_states_shape = hidden_states.shape
hidden_states_dtype = hidden_states.dtype
hidden_states_device = hidden_states.device
hidden_states_ref = hidden_states
topk_weights, topk_ids = topk_weights, topk_ids
if _should_use_masked_standard_layout(runner_config, quant_info, hidden_states):
output_dtype = (
torch.bfloat16
if quant_info.w13_weight.dtype == torch.bfloat16
else torch.float8_e4m3fn
)
masked_m, _, src2dst, hidden_states, hidden_states_scale = (
moe_ep_deepgemm_preprocess(
topk_ids,
runner_config.num_local_experts,
hidden_states,
runner_config.top_k,
quant_info.block_shape,
output_dtype=output_dtype,
use_mxfp8=quant_info.use_mxfp8,
)
)
# Use the global expert count because expected_m is a tuning hint, not
# the per-rank buffer capacity.
expected_m = max(
1,
ceil_div(
hidden_states_shape[0] * runner_config.top_k,
runner_config.num_experts,
),
)
if runner_config.inplace:
dispose_tensor(hidden_states_ref)
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
running_state["hidden_states_shape"] = hidden_states_shape
running_state["hidden_states_dtype"] = hidden_states_dtype
running_state["hidden_states_device"] = hidden_states_device
running_state["src2dst"] = src2dst
running_state["mxfp8_act_gran_k"] = (
quant_info.block_shape[1] if quant_info.block_shape else 128
)
return DeepGemmRunnerInput(
hidden_states=hidden_states,
hidden_states_scale=hidden_states_scale,
use_masked_gemm=True,
masked_m=masked_m,
expected_m=expected_m,
)
# The compact layout avoids scaling masked buffers with the expert count.
# Scatter and post-permute skip non-local experts mapped to -1.
block_e = 128
num_experts = runner_config.num_local_experts
num_assignments = topk_ids.numel()
all_tokens = _get_compact_all_tokens(num_assignments, num_experts, block_e)
tokens_per_expert, unused_masked_dst = fused_moe_dispatch_index(
topk_ids, num_experts, 1
)
dispose_tensor(unused_masked_dst)
valid_tokens_per_expert = tokens_per_expert
tokens_per_expert = (ceil_div(tokens_per_expert, block_e) * block_e).to(torch.int32)
# Keep graph-static shapes by appending padding to the final segment.
# Its m_indices stay -1, so DeepGEMM skips those rows.
tokens_per_expert[-1].add_(all_tokens - tokens_per_expert.sum())
k = hidden_states.size(1)
output_dtype = (
torch.bfloat16
if quant_info.w13_weight.dtype == torch.bfloat16
else torch.float8_e4m3fn
)
if output_dtype == torch.bfloat16:
packed_input_source = hidden_states
packed_input_source_scale = None
packed_input = torch.empty(
(all_tokens, k), device=hidden_states_device, dtype=torch.bfloat16
)
# ep_scatter ignores scales for BF16, but a real tensor keeps its
# Triton signature uniform across the existing DeepEP caller.
packed_input_scale = torch.empty(
(all_tokens, 1), device=hidden_states_device, dtype=torch.float32
)
else:
from sglang.kernels.ops.quantization.fp8_kernel import (
sglang_per_token_group_quant_fp8,
)
block_k = quant_info.block_shape[1] if quant_info.block_shape else 128
packed_input_source, packed_input_source_scale = (
sglang_per_token_group_quant_fp8(
hidden_states,
block_k,
column_major_scales=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
scale_tma_aligned=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
)
packed_input = torch.zeros(
(all_tokens, k),
device=hidden_states_device,
dtype=torch.float8_e4m3fn,
)
scale_width = k // block_k
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
scale_width = ceil_div(scale_width, 4)
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
packed_input_scale = torch.zeros(
(scale_width, all_tokens),
device=hidden_states_device,
dtype=packed_input_source_scale.dtype,
).transpose(0, 1)
else:
packed_input_scale = torch.zeros(
(all_tokens, scale_width),
device=hidden_states_device,
dtype=packed_input_source_scale.dtype,
)
expert_start_loc = torch.empty(
num_experts, device=hidden_states_device, dtype=torch.int32
)
m_indices = torch.empty(all_tokens, device=hidden_states_device, dtype=torch.int32)
src2dst = torch.empty_like(topk_ids, dtype=torch.int32)
ep_scatter(
packed_input_source,
packed_input_source_scale,
topk_ids,
tokens_per_expert,
valid_tokens_per_expert,
expert_start_loc,
packed_input,
packed_input_scale,
m_indices,
src2dst,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
quant_block_size=(quant_info.block_shape[1] if quant_info.block_shape else 128),
)
if packed_input_source is not hidden_states:
dispose_tensor(packed_input_source)
if packed_input_source_scale is not None:
dispose_tensor(packed_input_source_scale)
# Preserve the input when a shared expert or its gate may still use it.
if runner_config.inplace:
dispose_tensor(hidden_states_ref)
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
running_state["hidden_states_shape"] = hidden_states_shape
running_state["hidden_states_dtype"] = hidden_states_dtype
running_state["hidden_states_device"] = hidden_states_device
running_state["src2dst"] = src2dst
running_state["all_tokens"] = all_tokens
running_state["mxfp8_act_gran_k"] = (
quant_info.block_shape[1] if quant_info.block_shape else 128
)
return DeepGemmRunnerInput(
hidden_states=packed_input,
hidden_states_scale=packed_input_scale,
use_masked_gemm=False,
m_indices=m_indices,
)
@register_post_permute("deep_gemm", "standard")
def post_permute_deep_gemm_to_standard(
runner_output: DeepGemmRunnerOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> StandardCombineInput:
from sglang.kernels.ops.moe.ep_moe_kernels import post_reorder_deepgemm
from sglang.srt.layers.moe.token_dispatcher.standard import StandardCombineInput
hidden_states_shape = running_state["hidden_states_shape"]
hidden_states_dtype = running_state["hidden_states_dtype"]
hidden_states_device = running_state["hidden_states_device"]
topk_ids = running_state["topk_ids"]
topk_weights = running_state["topk_weights"]
src2dst = running_state["src2dst"]
with use_symmetric_memory(get_tp_group(), disabled=not is_allocation_symmetric()):
output = torch.empty(
hidden_states_shape, dtype=hidden_states_dtype, device=hidden_states_device
)
post_reorder_deepgemm(
runner_output.hidden_states,
output,
src2dst,
topk_ids,
topk_weights,
runner_config.top_k,
hidden_states_shape[0],
hidden_states_shape[1],
(
runner_config.routed_scaling_factor
if runner_config.routed_scaling_factor is not None
else 1.0
),
)
dispose_tensor(runner_output.hidden_states)
return StandardCombineInput(
hidden_states=output,
)
@register_pre_permute("deepep_ll", "deep_gemm")
def pre_permute_deepep_ll_to_deep_gemm(
dispatch_output: DeepEPLLDispatchOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepGemmRunnerInput:
hidden_states, hidden_states_scale, topk_ids, topk_weights, masked_m, expected_m = (
dispatch_output
)
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
running_state["hidden_states_shape"] = hidden_states.shape
running_state["hidden_states_dtype"] = hidden_states.dtype
running_state["hidden_states_device"] = hidden_states.device
# DeepEP-LL FP8 dispatch quantises activations at a fixed 128 block, not the checkpoint block_shape.
running_state["mxfp8_act_gran_k"] = 128
return DeepGemmRunnerInput(
hidden_states=hidden_states,
hidden_states_scale=hidden_states_scale,
use_masked_gemm=True,
masked_m=masked_m,
expected_m=expected_m,
)
@register_post_permute("deep_gemm", "deepep_ll")
def post_permute_deep_gemm_to_deepep_ll(
runner_output: DeepGemmRunnerOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepEPLLCombineInput:
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPLLCombineInput
return DeepEPLLCombineInput(
hidden_states=runner_output.hidden_states,
topk_ids=running_state["topk_ids"],
topk_weights=running_state["topk_weights"],
)
@register_pre_permute("deepep_normal", "deep_gemm")
def pre_permute_deepep_normal_to_deep_gemm(
dispatch_output: DeepEPNormalDispatchOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepGemmRunnerInput:
from sglang.kernels.ops.moe.ep_moe_kernels import ep_scatter
(
hidden_states,
hidden_states_scale,
topk_ids,
topk_weights,
num_recv_tokens_per_expert,
) = dispatch_output
assert runner_config.activation in ("silu", "situ")
all_tokens = sum(num_recv_tokens_per_expert)
running_state["all_tokens"] = all_tokens
K = hidden_states.shape[1]
hidden_states_shape = hidden_states.shape
hidden_states_device = hidden_states.device
hidden_states_dtype = hidden_states.dtype
running_state["hidden_states_shape"] = hidden_states_shape
running_state["hidden_states_device"] = hidden_states_device
running_state["hidden_states_dtype"] = hidden_states_dtype
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
# Deterministic inference zero-fills the scatter buffers: expert-alignment
# padding leaves slots that ep_scatter never writes, and pad garbage in
# input_tensor would leak batch-dependent values into the grouped GEMM.
# The scale buffer only matters for FP8 activations sharing this
# pre-permute (ep_scatter skips scales entirely for BF16 dispatch).
deterministic = get_exec().deterministic.enable_deterministic_inference
buffer_init = torch.zeros if deterministic else torch.empty
input_tensor = buffer_init(
(all_tokens, K),
device=hidden_states.device,
dtype=hidden_states.dtype,
)
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
# TODO check whether need `zeros`
input_tensor_scale = torch.zeros(
(ceil_div(K // 128, 4), all_tokens),
device=hidden_states.device,
dtype=torch.int,
).transpose(0, 1)
else:
input_tensor_scale = buffer_init(
(all_tokens, K // 128),
device=hidden_states.device,
dtype=torch.float32,
)
m_indices = buffer_init(all_tokens, device=hidden_states.device, dtype=torch.int32)
output_index = torch.empty_like(topk_ids)
if get_offloader().forbid_copy_engine_usage:
num_recv_tokens_per_expert_gpu = copy_list_to_gpu_no_ce(
num_recv_tokens_per_expert
)
else:
num_recv_tokens_per_expert_gpu = torch.tensor(
num_recv_tokens_per_expert,
dtype=torch.int32,
pin_memory=True,
device="cpu",
).cuda(non_blocking=True)
expert_start_loc = torch.empty_like(num_recv_tokens_per_expert_gpu)
ep_scatter(
hidden_states,
hidden_states_scale,
topk_ids,
num_recv_tokens_per_expert_gpu,
num_recv_tokens_per_expert_gpu,
expert_start_loc,
input_tensor,
input_tensor_scale,
m_indices,
output_index,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
dispose_tensor(hidden_states)
if hidden_states_scale is not None:
dispose_tensor(hidden_states_scale)
running_state["output_index"] = output_index
return DeepGemmRunnerInput(
hidden_states=input_tensor,
hidden_states_scale=input_tensor_scale,
use_masked_gemm=False,
m_indices=m_indices,
)
@register_post_permute("deep_gemm", "deepep_normal")
def post_permute_deep_gemm_to_deepep_normal(
runner_output: DeepGemmRunnerOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepEPNormalCombineInput:
from sglang.kernels.ops.moe.ep_moe_kernels import ep_gather
from sglang.srt.layers.moe.token_dispatcher.deepep import DeepEPNormalCombineInput
hidden_states = runner_output.hidden_states
topk_ids = running_state["topk_ids"]
topk_weights = running_state["topk_weights"]
output_index = running_state["output_index"]
gather_out = torch.empty(
running_state["hidden_states_shape"],
device=running_state["hidden_states_device"],
dtype=torch.bfloat16,
)
ep_gather(hidden_states, topk_ids, topk_weights, output_index, gather_out)
return DeepEPNormalCombineInput(
hidden_states=gather_out,
topk_ids=running_state["topk_ids"],
topk_weights=running_state["topk_weights"],
)
def _varlen_deep_gemm_situ_mul_quant(
gateup_output: torch.Tensor,
masked_m: torch.Tensor,
group_size: int,
topk: int,
beta: float,
linear_beta: float,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Fused SiTU activation + per-group fp8 quant via CUDA JIT kernel."""
from sglang.kernels.ops.kimi_k3 import situ_and_mul_masked_post_quant
E, N, D_2 = gateup_output.shape
D = D_2 // 2
G = D // group_size
packed_ue8m0 = deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
down_input = torch.empty(
(E, N, D), device=gateup_output.device, dtype=torch.float8_e4m3fn
)
if packed_ue8m0:
down_input_scale = torch.empty(
(E, G // 4, N), device=gateup_output.device, dtype=torch.int32
)
else:
down_input_scale = torch.empty(
(E, N, G), device=gateup_output.device, dtype=torch.float32
)
situ_and_mul_masked_post_quant(
gateup_output,
down_input,
down_input_scale,
group_size,
masked_m,
beta=beta,
linear_beta=linear_beta,
scale_ue8m0=packed_ue8m0,
topk=topk,
transposed=packed_ue8m0,
)
if packed_ue8m0:
down_input_scale = down_input_scale.transpose(-1, -2)
return down_input, down_input_scale
def _varlen_deep_gemm_silu_mul_quant(
gateup_output: torch.Tensor,
masked_m: Optional[torch.Tensor],
group_size: int,
topk: int,
swiglu_limit: Optional[float] = None,
swizzle: bool = False,
gemm1_alpha: Optional[float] = None,
gemm1_clamp_limit: Optional[float] = None,
num_real_tokens: Optional[int] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
assert masked_m is not None
hidden_states_device = gateup_output.device
E, N, D_2 = gateup_output.shape
D = D_2 // 2
del D_2
G = D // group_size
# oai-swiglu (gemm1_alpha) stays on the Triton kernel until
# per_token_group_quant grows an activation-kind axis. The output_scale dtype picks the schedule: packed
# int32 UE8M0 (no follow-up transform; needs G % 4 == 0 and the
# num_real_tokens grid bound) when eligible, row-major fp32 otherwise.
if gemm1_alpha is not None:
assert (
swiglu_limit is None
), "swiglu_limit and gemm1_alpha are mutually exclusive"
assert not swizzle, "swizzle is not supported with gemm1_alpha"
from sglang.kernels.ops.moe.ep_moe_kernels import (
silu_and_mul_masked_post_quant_fwd,
)
use_packed = (
deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
and num_real_tokens is not None
and G % 4 == 0
and D % (group_size * 4) == 0
)
down_input = torch.empty(
(E, N, D), device=hidden_states_device, dtype=torch.float8_e4m3fn
)
down_input_scale = torch.empty(
(E, G // 4, N) if use_packed else (E, N, G),
device=hidden_states_device,
dtype=torch.int32 if use_packed else torch.float32,
)
silu_and_mul_masked_post_quant_fwd(
gateup_output,
down_input,
down_input_scale,
group_size,
masked_m,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
gemm1_alpha=gemm1_alpha,
gemm1_clamp_limit=gemm1_clamp_limit or 0.0,
num_real_tokens=num_real_tokens,
topk=topk,
)
if use_packed:
down_input_scale = down_input_scale.transpose(-1, -2)
return down_input, down_input_scale
# DSV4-specific activations (clamped swiglu, swizzled gate|up layout) stay
# on the DSV4 JIT kernel; it is the only implementation carrying them.
if swiglu_limit is not None or swizzle:
assert N % 4 == 0 and G % 4 == 0 and D // 8 >= E, (
"DSV4 JIT activation requires N % 4 == 0, G % 4 == 0 and "
f"D // 8 >= num_experts, got N={N} G={G} D={D} E={E}"
)
packed_ue8m0 = deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0
down_input = torch.empty(
(E, N, D), device=hidden_states_device, dtype=torch.float8_e4m3fn
)
down_input_scale = torch.empty(
(E, G // 4, N) if packed_ue8m0 else (E, N, G),
device=hidden_states_device,
dtype=torch.int32 if packed_ue8m0 else torch.float32,
)
silu_and_mul_masked_post_quant(
gateup_output,
down_input,
down_input_scale,
group_size,
masked_m,
scale_ue8m0=packed_ue8m0,
topk=topk,
transposed=packed_ue8m0,
swiglu_limit=swiglu_limit,
swizzle=swizzle,
)
if packed_ue8m0:
down_input_scale = down_input_scale.transpose(-1, -2)
return down_input, down_input_scale
# Default plain-silu path: the unified JIT masked fused quant. It allocates
# the outputs itself, with scales directly in the layout deep_gemm consumes
# (packed-int32 col-major for UE8M0, TMA-aligned col-major fp32 otherwise),
# so the caller's get_mn_major transform short-circuits.
expected_m = ceil_div(num_real_tokens * topk, E) if num_real_tokens else None
return per_token_group_quant(
gateup_output,
group_size=group_size,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
fuse_silu_and_mul=True,
masked_m=masked_m,
expected_m=expected_m,
column_major_scales=True,
)
@triton.jit
def _situ_mul_quant_contig_kernel(
g_ptr, # [rows, 2N] bf16, non-interleaved [gate; up] halves
q_ptr, # [rows, N] fp8 out
s_ptr, # [rows, KG] fp32 scales out
N,
KG,
situ_beta,
situ_linear_beta,
GROUP: tl.constexpr,
KG_POW2: tl.constexpr,
):
row = tl.program_id(0).to(tl.int64)
rows2d = tl.arange(0, KG_POW2)[:, None]
cols = tl.arange(0, GROUP)[None, :]
offs = rows2d * GROUP + cols
mask = rows2d < KG
gate = tl.load(g_ptr + row * 2 * N + offs, mask=mask, other=0.0).to(tl.float32)
up = tl.load(g_ptr + row * 2 * N + N + offs, mask=mask, other=0.0).to(tl.float32)
# tanh(x) == 2*sigmoid(2x) - 1 (avoids a libdevice dependency)
gate_t = 2.0 * tl.sigmoid(2.0 * gate / situ_beta) - 1.0
gate = situ_beta * gate_t * tl.sigmoid(gate)
up_t = 2.0 * tl.sigmoid(2.0 * up / situ_linear_beta) - 1.0
y = gate * situ_linear_beta * up_t
amax = tl.clamp(tl.max(tl.abs(y), axis=1), min=1e-10, max=float("inf"))
q = (y * (448.0 / amax)[:, None]).to(tl.float8e4nv)
tl.store(q_ptr + row * N + offs, q, mask=mask)
srow = tl.arange(0, KG_POW2)
tl.store(s_ptr + row * KG + srow, amax / 448.0, mask=srow < KG)
def _apply_swiglu_limit(
gateup_output: torch.Tensor, swiglu_limit: float
) -> torch.Tensor:
assert swiglu_limit == 10
num_tokens, hidden_size_x2 = gateup_output.shape
assert gateup_output.dtype == torch.bfloat16
gate, up = torch.chunk(gateup_output, chunks=2, dim=-1)
assert gate.shape == (num_tokens, hidden_size_x2 // 2)
assert up.shape == (num_tokens, hidden_size_x2 // 2)
up = torch.clamp(up, min=-swiglu_limit, max=swiglu_limit)
gate = torch.clamp(gate, max=swiglu_limit)
out = torch.cat([gate, up], dim=-1)
assert out.shape == (num_tokens, hidden_size_x2)
return out
@register_pre_permute("deepep_v2", "deep_gemm")
def pre_permute_deepep_v2_to_deep_gemm(
dispatch_output: DeepEPv2DispatchOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepGemmRunnerInput:
from sglang.kernels.ops.moe.ep_moe_kernels import (
ep_expand_init_m_indices_from_psum,
ep_scatter_from_psum,
)
hidden_states = dispatch_output.hidden_states
hidden_states_scale = dispatch_output.hidden_states_scale
topk_ids = dispatch_output.topk_ids
topk_weights = dispatch_output.topk_weights
psum_num_recv_tokens_per_expert = dispatch_output.psum_num_recv_tokens_per_expert
is_expanded = dispatch_output.is_expanded
hidden_states_scale_tma_aligned = dispatch_output.hidden_states_scale_tma_aligned
deepep_v2_use_masked = dispatch_output.use_masked_gemm
deepep_v2_expected_m = dispatch_output.expected_m
deepep_v2_masked_max_m = dispatch_output.masked_max_m
deepep_v2_total_expanded = dispatch_output.total_expanded
deepep_v2_expert_alignment = dispatch_output.expert_alignment
if hidden_states_scale is None:
raise RuntimeError(
"DeepEP v2 -> DeepGEMM requires FP8 dispatch output with activation "
"scales, but the dispatch output carried none."
)
assert runner_config.activation == "silu"
if is_expanded:
if psum_num_recv_tokens_per_expert is None:
raise RuntimeError(
"DeepEP v2 requires the native expert prefix sums from the "
"ElasticBuffer dispatch handle."
)
all_tokens = hidden_states.shape[0]
running_state["all_tokens"] = all_tokens
running_state["hidden_states_shape"] = hidden_states.shape
running_state["hidden_states_device"] = hidden_states.device
running_state["hidden_states_dtype"] = hidden_states.dtype
running_state["topk_ids"] = None
running_state["topk_weights"] = topk_weights
running_state["deepep_v2_expanded"] = True
if deepep_v2_use_masked:
# masked_m bounds each expert independently of buffer capacity.
from sglang.kernels.ops.moe.ep_moe_kernels import expand_to_masked_slab
num_local_experts = psum_num_recv_tokens_per_expert.shape[0]
input_tensor, input_tensor_scale, masked_m = expand_to_masked_slab(
hidden_states,
hidden_states_scale,
psum_num_recv_tokens_per_expert,
num_local_experts,
deepep_v2_masked_max_m,
deepep_v2_expert_alignment,
)
running_state["deepep_v2_masked"] = True
running_state["deepep_v2_psum"] = psum_num_recv_tokens_per_expert
running_state["deepep_v2_total_expanded"] = deepep_v2_total_expanded
running_state["deepep_v2_expert_alignment"] = deepep_v2_expert_alignment
return DeepGemmRunnerInput(
hidden_states=input_tensor,
hidden_states_scale=input_tensor_scale,
use_masked_gemm=True,
masked_m=masked_m,
expected_m=deepep_v2_expected_m,
)
# Mark aligned expert rows and leave the unused receive tail at -1.
m_indices = torch.full(
(all_tokens,), -1, device=hidden_states.device, dtype=torch.int32
)
ep_expand_init_m_indices_from_psum(psum_num_recv_tokens_per_expert, m_indices)
return DeepGemmRunnerInput(
hidden_states=hidden_states,
hidden_states_scale=hidden_states_scale,
use_masked_gemm=False,
m_indices=m_indices,
hidden_states_scale_tma_aligned=hidden_states_scale_tma_aligned,
)
all_tokens = int(psum_num_recv_tokens_per_expert[-1].item())
K = hidden_states.shape[1]
running_state["all_tokens"] = all_tokens
running_state["hidden_states_shape"] = hidden_states.shape
running_state["hidden_states_device"] = hidden_states.device
running_state["hidden_states_dtype"] = hidden_states.dtype
running_state["topk_ids"] = topk_ids
running_state["topk_weights"] = topk_weights
input_tensor = torch.empty(
(all_tokens, K), device=hidden_states.device, dtype=hidden_states.dtype
)
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
# Packed UE8M0 scales require zero padding lanes.
input_tensor_scale = torch.zeros(
(ceil_div(K // 128, 4), all_tokens),
device=hidden_states.device,
dtype=torch.int,
).transpose(0, 1)
else:
input_tensor_scale = torch.empty(
(all_tokens, K // 128), device=hidden_states.device, dtype=torch.float32
)
m_indices = torch.empty(all_tokens, device=hidden_states.device, dtype=torch.int32)
output_index = torch.empty_like(topk_ids)
# Contiguous psum already includes the 128-row expert alignment.
expert_start_loc = torch.empty_like(psum_num_recv_tokens_per_expert)
ep_scatter_from_psum(
hidden_states,
hidden_states_scale,
topk_ids,
psum_num_recv_tokens_per_expert,
expert_start_loc,
input_tensor,
input_tensor_scale,
m_indices,
output_index,
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
)
dispose_tensor(hidden_states)
dispose_tensor(hidden_states_scale)
running_state["output_index"] = output_index
return DeepGemmRunnerInput(
hidden_states=input_tensor,
hidden_states_scale=input_tensor_scale,
use_masked_gemm=False,
m_indices=m_indices,
)
@register_post_permute("deep_gemm", "deepep_v2")
def post_permute_deep_gemm_to_deepep_v2(
runner_output: DeepGemmRunnerOutput,
quant_info: DeepGemmMoeQuantInfo,
runner_config: MoeRunnerConfig,
running_state: dict,
) -> DeepEPv2CombineInput:
from sglang.kernels.ops.moe.ep_moe_kernels import ep_gather
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import DeepEPv2CombineInput
if running_state.get("deepep_v2_expanded", False):
hidden_states = runner_output.hidden_states
topk_weights = running_state["topk_weights"]
if running_state.get("deepep_v2_masked", False):
# Expanded combine does not consume top-k weights.
from sglang.kernels.ops.moe.ep_moe_kernels import masked_slab_to_expand
hidden_states = masked_slab_to_expand(
hidden_states,
running_state["deepep_v2_psum"],
running_state["deepep_v2_total_expanded"],
running_state["deepep_v2_expert_alignment"],
topk_weights=topk_weights,
)
return DeepEPv2CombineInput(hidden_states, None)
if topk_weights is not None:
# Expanded combine does not consume top-k weights.
hidden_states = hidden_states * topk_weights.to(
hidden_states.dtype
).unsqueeze(-1)
return DeepEPv2CombineInput(hidden_states, None)
hidden_states = runner_output.hidden_states
topk_ids = running_state["topk_ids"]
topk_weights = running_state["topk_weights"]
output_index = running_state["output_index"]
gather_out = torch.empty(
running_state["hidden_states_shape"],
device=running_state["hidden_states_device"],
dtype=torch.bfloat16,
)
ep_gather(hidden_states, topk_ids, topk_weights, output_index, gather_out)
return DeepEPv2CombineInput(
hidden_states=gather_out,
topk_weights=topk_weights,
)