Support MXFP8 and deferred route weighting in DeepEP v2 (#40030)
Co-authored-by: metamergebot <324680979+metamergebot@users.noreply.github.com> Co-authored-by: Xingyu Liu <38244988+charlotte12l@users.noreply.github.com> Co-authored-by: pranjalssh <14260275+pranjalssh@users.noreply.github.com>
This commit is contained in:
co-authored by
metamergebot
Xingyu Liu
pranjalssh
parent
6cc9090d1f
commit
0e5347db82
@@ -99,17 +99,16 @@ SGL_DEVICE CTAWork get_work(const SiluMulQuantVarlenParams& params) {
|
|||||||
return result;
|
return result;
|
||||||
}
|
}
|
||||||
|
|
||||||
template <bool kScaleUE8M0, bool kTransposed, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
template <uint32_t kGroupSize, bool kScaleUE8M0, bool kTransposed, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
||||||
__global__ __launch_bounds__(1024, 2) void // maximize occupancy
|
__global__ __launch_bounds__(1024, 2) void // maximize occupancy
|
||||||
silu_mul_quant_varlen_kernel(const SiluMulQuantVarlenParams __grid_constant__ params) {
|
silu_mul_quant_varlen_kernel(const SiluMulQuantVarlenParams __grid_constant__ params) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
|
|
||||||
constexpr uint32_t kGroupSize = 128u;
|
constexpr uint32_t kWorkThreads = kGroupSize / 8u;
|
||||||
constexpr uint32_t kWorkThreads = 16u;
|
|
||||||
// each thread will handle 8 elements
|
// each thread will handle 8 elements
|
||||||
using InputVec = AlignedVector<bf16x2_t, 4>;
|
using InputVec = AlignedVector<bf16x2_t, 4>;
|
||||||
using OutputVec = AlignedVector<fp8x2_e4m3_t, 4>;
|
using OutputVec = AlignedVector<fp8x2_e4m3_t, 4>;
|
||||||
static_assert(8 * kWorkThreads == 128, "Invalid tiling");
|
static_assert(kGroupSize == 32 || kGroupSize == 128, "unsupported group_size");
|
||||||
static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0");
|
static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0");
|
||||||
|
|
||||||
const auto [expert_id, token_id, valid] = get_work(params);
|
const auto [expert_id, token_id, valid] = get_work(params);
|
||||||
@@ -250,11 +249,11 @@ __global__ __launch_bounds__(1024, 2) void // maximize occupancy
|
|||||||
|
|
||||||
template <int64_t kGroupSize, bool kScaleUE8M0, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
template <int64_t kGroupSize, bool kScaleUE8M0, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
||||||
struct SiluAndMulMaskedPostQuantKernel {
|
struct SiluAndMulMaskedPostQuantKernel {
|
||||||
static_assert(kGroupSize == 128);
|
static_assert(kGroupSize == 32 || kGroupSize == 128);
|
||||||
static constexpr auto kernel_normal =
|
static constexpr auto kernel_normal =
|
||||||
silu_mul_quant_varlen_kernel<kScaleUE8M0, false, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
silu_mul_quant_varlen_kernel<kGroupSize, kScaleUE8M0, false, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
||||||
static constexpr auto kernel_transposed =
|
static constexpr auto kernel_transposed =
|
||||||
silu_mul_quant_varlen_kernel<true, true, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
silu_mul_quant_varlen_kernel<kGroupSize, true, true, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
||||||
|
|
||||||
static void
|
static void
|
||||||
run(const tvm::ffi::TensorView input,
|
run(const tvm::ffi::TensorView input,
|
||||||
@@ -385,16 +384,15 @@ struct SiluMulQuantContigParams {
|
|||||||
uint32_t scale_row_stride_int32; // only used when kTransposed=true
|
uint32_t scale_row_stride_int32; // only used when kTransposed=true
|
||||||
};
|
};
|
||||||
|
|
||||||
template <bool kScaleUE8M0, bool kTransposed, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
template <uint32_t kGroupSize, bool kScaleUE8M0, bool kTransposed, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
||||||
__global__ __launch_bounds__(1024, 2) void // maximize occupancy
|
__global__ __launch_bounds__(1024, 2) void // maximize occupancy
|
||||||
silu_mul_quant_contig_kernel(const SiluMulQuantContigParams __grid_constant__ params) {
|
silu_mul_quant_contig_kernel(const SiluMulQuantContigParams __grid_constant__ params) {
|
||||||
using namespace device;
|
using namespace device;
|
||||||
|
|
||||||
constexpr uint32_t kGroupSize = 128u;
|
constexpr uint32_t kWorkThreads = kGroupSize / 8u;
|
||||||
constexpr uint32_t kWorkThreads = 16u;
|
|
||||||
using InputVec = AlignedVector<bf16x2_t, 4>;
|
using InputVec = AlignedVector<bf16x2_t, 4>;
|
||||||
using OutputVec = AlignedVector<fp8x2_e4m3_t, 4>;
|
using OutputVec = AlignedVector<fp8x2_e4m3_t, 4>;
|
||||||
static_assert(8 * kWorkThreads == 128, "Invalid tiling");
|
static_assert(kGroupSize == 32 || kGroupSize == 128, "unsupported group_size");
|
||||||
static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0");
|
static_assert(!(kTransposed && !kScaleUE8M0), "transposed layout only supports ue8m0");
|
||||||
|
|
||||||
const auto token_id = blockIdx.x;
|
const auto token_id = blockIdx.x;
|
||||||
@@ -473,11 +471,11 @@ __global__ __launch_bounds__(1024, 2) void // maximize occupancy
|
|||||||
|
|
||||||
template <int64_t kGroupSize, bool kScaleUE8M0, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
template <int64_t kGroupSize, bool kScaleUE8M0, bool kSwizzle, bool kUsePDL, bool kApplySwigluLimit>
|
||||||
struct SiluAndMulContigPostQuantKernel {
|
struct SiluAndMulContigPostQuantKernel {
|
||||||
static_assert(kGroupSize == 128);
|
static_assert(kGroupSize == 32 || kGroupSize == 128);
|
||||||
static constexpr auto kernel_normal =
|
static constexpr auto kernel_normal =
|
||||||
silu_mul_quant_contig_kernel<kScaleUE8M0, false, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
silu_mul_quant_contig_kernel<kGroupSize, kScaleUE8M0, false, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
||||||
static constexpr auto kernel_transposed =
|
static constexpr auto kernel_transposed =
|
||||||
silu_mul_quant_contig_kernel<true, true, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
silu_mul_quant_contig_kernel<kGroupSize, true, true, kSwizzle, kUsePDL, kApplySwigluLimit>;
|
||||||
|
|
||||||
static void
|
static void
|
||||||
run(const tvm::ffi::TensorView input,
|
run(const tvm::ffi::TensorView input,
|
||||||
|
|||||||
@@ -1270,13 +1270,13 @@ def ep_scatter_from_psum(
|
|||||||
m_indices: torch.Tensor,
|
m_indices: torch.Tensor,
|
||||||
output_index: torch.Tensor,
|
output_index: torch.Tensor,
|
||||||
scale_ue8m0: bool = False,
|
scale_ue8m0: bool = False,
|
||||||
|
quant_block_size: int = 128,
|
||||||
):
|
):
|
||||||
BLOCK_E = 128
|
BLOCK_E = 128
|
||||||
BLOCK_D = 128
|
|
||||||
num_warps = 8
|
num_warps = 8
|
||||||
num_experts = psum_num_recv_tokens_per_expert.shape[0]
|
num_experts = psum_num_recv_tokens_per_expert.shape[0]
|
||||||
hidden_size = recv_x.shape[1]
|
hidden_size = recv_x.shape[1]
|
||||||
scale_hidden_size = hidden_size // BLOCK_D
|
scale_hidden_size = hidden_size // quant_block_size
|
||||||
if scale_ue8m0:
|
if scale_ue8m0:
|
||||||
scale_hidden_size = ceil_div(scale_hidden_size, 4)
|
scale_hidden_size = ceil_div(scale_hidden_size, 4)
|
||||||
|
|
||||||
|
|||||||
@@ -52,13 +52,13 @@ _is_xpu = is_xpu()
|
|||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
|
||||||
|
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sglang.kernels.ops.quantization import (
|
from sglang.kernels.ops.quantization import sgl_per_token_quant_fp8
|
||||||
per_token_group_quant,
|
|
||||||
sgl_per_token_quant_fp8,
|
|
||||||
)
|
|
||||||
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
|
from sglang.kernels.ops.quantization.per_tensor_quant_fp8 import (
|
||||||
per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8,
|
per_tensor_quant_fp8 as sgl_per_tensor_quant_fp8,
|
||||||
)
|
)
|
||||||
|
from sglang.kernels.ops.quantization.per_token_group_quant import (
|
||||||
|
per_token_group_quant,
|
||||||
|
)
|
||||||
elif _is_xpu:
|
elif _is_xpu:
|
||||||
from sgl_kernel import sgl_per_tensor_quant_fp8, sgl_per_token_quant_fp8
|
from sgl_kernel import sgl_per_tensor_quant_fp8, sgl_per_token_quant_fp8
|
||||||
|
|
||||||
|
|||||||
@@ -14,7 +14,9 @@ from sglang.srt.utils import get_device_name, is_cuda, is_hip
|
|||||||
_is_cuda = is_cuda()
|
_is_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
if _is_cuda:
|
if _is_cuda:
|
||||||
from sglang.kernels.ops.quantization import per_token_group_quant
|
from sglang.kernels.ops.quantization.per_token_group_quant import (
|
||||||
|
per_token_group_quant,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|||||||
@@ -22,6 +22,7 @@ from sglang.srt.arg_groups.overrides import (
|
|||||||
resolving_view,
|
resolving_view,
|
||||||
run_post_process_pass,
|
run_post_process_pass,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.configs.moe_model_registry import model_supports_deepep_v2
|
||||||
from sglang.srt.connector import ConnectorType
|
from sglang.srt.connector import ConnectorType
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
|
from sglang.srt.model_executor.cuda_graph_config import Backend, Phase, with_phase
|
||||||
@@ -587,7 +588,7 @@ def validate_deepep_v2_dispatch_token_budget(server_args: Any) -> None:
|
|||||||
|
|
||||||
|
|
||||||
def validate_deepep_v2_model_architecture(server_args: Any) -> None:
|
def validate_deepep_v2_model_architecture(server_args: Any) -> None:
|
||||||
"""Allow DeepEP v2 only where its model workflow is validated."""
|
"""Allow DeepEP v2 only for registered model architectures."""
|
||||||
|
|
||||||
if (
|
if (
|
||||||
parse_connector_type(resolved_view(server_args).model_path)
|
parse_connector_type(resolved_view(server_args).model_path)
|
||||||
@@ -599,24 +600,15 @@ def validate_deepep_v2_model_architecture(server_args: Any) -> None:
|
|||||||
"--moe-a2a-backend deepep."
|
"--moe-a2a-backend deepep."
|
||||||
)
|
)
|
||||||
|
|
||||||
architectures = (
|
hf_config = model_config_of(server_args).hf_config
|
||||||
getattr(model_config_of(server_args).hf_config, "architectures", None) or []
|
if not model_supports_deepep_v2(hf_config):
|
||||||
)
|
architectures = getattr(hf_config, "architectures", None) or []
|
||||||
|
architecture = architectures[0] if architectures else None
|
||||||
architecture = architectures[0] if architectures else None
|
|
||||||
# These architectures take the A2A MoE path and skip post-expert
|
|
||||||
# all-reduce.
|
|
||||||
validated_architectures = (
|
|
||||||
"DeepseekV3ForCausalLM",
|
|
||||||
"DeepseekV4ForCausalLM",
|
|
||||||
"Qwen3MoeForCausalLM",
|
|
||||||
)
|
|
||||||
if architecture not in validated_architectures:
|
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"DeepEP v2 MoE is not validated for {architecture!r}; supported "
|
f"DeepEP v2 MoE is not validated for {architecture!r}. The model "
|
||||||
f"architectures are {sorted(validated_architectures)}. "
|
"package must register its architecture with "
|
||||||
"Other model workflows may require an all-reduce after A2A "
|
"register_deepep_v2_model, because its combine and post-expert "
|
||||||
"combine. Use --moe-a2a-backend deepep."
|
"reduction semantics must be validated first."
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,46 @@
|
|||||||
|
"""Model eligibility and MXFP8 precision policy for the DeepEP v2 adapter.
|
||||||
|
|
||||||
|
External model packages register architecture names before server-argument
|
||||||
|
validation, without adding private architecture names to SGLang core.
|
||||||
|
Only the primary architecture controls eligibility and precision.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
# Membership enables DeepEP v2; the value requests FP32 SiLU intermediates.
|
||||||
|
_DEEPEP_V2_MODELS: dict[str, bool] = {
|
||||||
|
"DeepseekV3ForCausalLM": False,
|
||||||
|
"DeepseekV4ForCausalLM": False,
|
||||||
|
"Qwen3MoeForCausalLM": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def register_deepep_v2_model(
|
||||||
|
architecture: str, *, silu_mul_keep_fp32: bool = False
|
||||||
|
) -> None:
|
||||||
|
"""Register a model architecture with validated DeepEP v2 semantics."""
|
||||||
|
|
||||||
|
if (
|
||||||
|
architecture in _DEEPEP_V2_MODELS
|
||||||
|
and _DEEPEP_V2_MODELS[architecture] != silu_mul_keep_fp32
|
||||||
|
):
|
||||||
|
raise ValueError(f"Conflicting DeepEP v2 precision policy for {architecture}")
|
||||||
|
_DEEPEP_V2_MODELS[architecture] = silu_mul_keep_fp32
|
||||||
|
|
||||||
|
|
||||||
|
def model_supports_deepep_v2(hf_config: Any) -> bool:
|
||||||
|
"""Check model eligibility, independently of runtime layout and topology."""
|
||||||
|
|
||||||
|
architectures = getattr(hf_config, "architectures", None) or ()
|
||||||
|
architecture = architectures[0] if architectures else None
|
||||||
|
return architecture is not None and architecture in _DEEPEP_V2_MODELS
|
||||||
|
|
||||||
|
|
||||||
|
def model_requires_fp32_silu_mul(hf_config: Any) -> bool:
|
||||||
|
"""Whether this model opts into FP32 intermediates for DeepEP v2 MXFP8."""
|
||||||
|
|
||||||
|
architectures = getattr(hf_config, "architectures", None) or ()
|
||||||
|
architecture = architectures[0] if architectures else None
|
||||||
|
return architecture is not None and _DEEPEP_V2_MODELS.get(architecture, False)
|
||||||
@@ -12,6 +12,7 @@ from torch.nn.parameter import UninitializedParameter
|
|||||||
|
|
||||||
from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs
|
from sglang.srt.batch_overlap.single_batch_overlap import DownGemmOverlapArgs
|
||||||
from sglang.srt.batch_overlap.two_batch_overlap import MaybeTboDeepEPDispatcher
|
from sglang.srt.batch_overlap.two_batch_overlap import MaybeTboDeepEPDispatcher
|
||||||
|
from sglang.srt.configs.moe_model_registry import model_requires_fp32_silu_mul
|
||||||
from sglang.srt.distributed import (
|
from sglang.srt.distributed import (
|
||||||
get_moe_ep_group,
|
get_moe_ep_group,
|
||||||
get_tp_group,
|
get_tp_group,
|
||||||
@@ -78,6 +79,7 @@ from sglang.srt.runtime_context import (
|
|||||||
get_global_dwdp_manager,
|
get_global_dwdp_manager,
|
||||||
get_parallel,
|
get_parallel,
|
||||||
get_server_args,
|
get_server_args,
|
||||||
|
process_model_config,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
cpu_has_amx_support,
|
cpu_has_amx_support,
|
||||||
@@ -209,6 +211,11 @@ def create_moe_dispatcher(
|
|||||||
hidden_size=moe_runner_config.hidden_size,
|
hidden_size=moe_runner_config.hidden_size,
|
||||||
params_dtype=moe_runner_config.params_dtype,
|
params_dtype=moe_runner_config.params_dtype,
|
||||||
use_fp8_dispatch=output_dtype is DispatcherOutputDtype.FP8,
|
use_fp8_dispatch=output_dtype is DispatcherOutputDtype.FP8,
|
||||||
|
activation_scale_block_size=(
|
||||||
|
32
|
||||||
|
if isinstance(quant_method, Fp8MoEMethod) and quant_method.use_mxfp8
|
||||||
|
else 128
|
||||||
|
),
|
||||||
)
|
)
|
||||||
elif a2a_backend.is_flashinfer():
|
elif a2a_backend.is_flashinfer():
|
||||||
return FlashinferDispatcher(
|
return FlashinferDispatcher(
|
||||||
@@ -267,18 +274,19 @@ def _validate_deepep_v2_quant_method(quant_method) -> None:
|
|||||||
reason = None
|
reason = None
|
||||||
if not isinstance(quant_method, Fp8MoEMethod):
|
if not isinstance(quant_method, Fp8MoEMethod):
|
||||||
reason = f"selected {type(quant_method).__name__}"
|
reason = f"selected {type(quant_method).__name__}"
|
||||||
elif quant_method.use_mxfp8:
|
|
||||||
reason = "selected MXFP8 weights"
|
|
||||||
elif quant_method.is_fp4_expert:
|
elif quant_method.is_fp4_expert:
|
||||||
reason = "selected FP4 experts"
|
reason = "selected FP4 experts"
|
||||||
elif list(quant_method.weight_block_size or []) != [128, 128]:
|
elif list(quant_method.weight_block_size or []) != (
|
||||||
reason = f"has weight_block_size={quant_method.weight_block_size}"
|
[1, 32] if quant_method.use_mxfp8 else [128, 128]
|
||||||
|
):
|
||||||
|
quant_format = "MXFP8 " if quant_method.use_mxfp8 else ""
|
||||||
|
reason = f"has {quant_format}weight_block_size={quant_method.weight_block_size}"
|
||||||
elif config.activation_scheme != "dynamic":
|
elif config.activation_scheme != "dynamic":
|
||||||
reason = f"has activation_scheme={config.activation_scheme!r}"
|
reason = f"has activation_scheme={config.activation_scheme!r}"
|
||||||
|
|
||||||
if reason is not None:
|
if reason is not None:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"--moe-a2a-backend deepep_v2 requires either 128x128 blockwise FP8 "
|
"--moe-a2a-backend deepep_v2 requires 128x128 blockwise FP8 or 1x32 MXFP8 "
|
||||||
"experts with dynamic activation scaling or unquantized BF16 "
|
"experts with dynamic activation scaling or unquantized BF16 "
|
||||||
f"experts, but this layer {reason}. Use a compatible checkpoint or "
|
f"experts, but this layer {reason}. Use a compatible checkpoint or "
|
||||||
"--moe-a2a-backend deepep."
|
"--moe-a2a-backend deepep."
|
||||||
@@ -479,6 +487,14 @@ class FusedMoE(torch.nn.Module):
|
|||||||
)
|
)
|
||||||
_validate_hpc_ops_quant_method(self.quant_method)
|
_validate_hpc_ops_quant_method(self.quant_method)
|
||||||
_validate_deepep_v2_quant_method(self.quant_method)
|
_validate_deepep_v2_quant_method(self.quant_method)
|
||||||
|
if (
|
||||||
|
get_moe_a2a_backend().is_deepep_v2()
|
||||||
|
and isinstance(self.quant_method, Fp8MoEMethod)
|
||||||
|
and self.quant_method.use_mxfp8
|
||||||
|
):
|
||||||
|
self.moe_runner_config.silu_mul_keep_fp32 = model_requires_fp32_silu_mul(
|
||||||
|
process_model_config().hf_config
|
||||||
|
)
|
||||||
nvfp4_deferred = envs.SGLANG_ENABLE_MOE_DEFERRED_FINALIZE.get() and isinstance(
|
nvfp4_deferred = envs.SGLANG_ENABLE_MOE_DEFERRED_FINALIZE.get() and isinstance(
|
||||||
self.quant_method, ModelOptNvFp4FusedMoEMethod
|
self.quant_method, ModelOptNvFp4FusedMoEMethod
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -64,6 +64,9 @@ class MoeRunnerConfig:
|
|||||||
gate_up_interleaved: bool = True
|
gate_up_interleaved: bool = True
|
||||||
layer: Optional[torch.nn.Module] = None
|
layer: Optional[torch.nn.Module] = None
|
||||||
use_tp_all_gather_activation: bool = False
|
use_tp_all_gather_activation: bool = False
|
||||||
|
# Request FP32 SiLU/multiply intermediates until FP8 quantization.
|
||||||
|
# False preserves backend defaults, including their existing FP32 paths.
|
||||||
|
silu_mul_keep_fp32: bool = False
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
|
|||||||
@@ -13,7 +13,7 @@ from sglang.kernels.ops.attention.dsv4 import (
|
|||||||
silu_and_mul_masked_post_quant,
|
silu_and_mul_masked_post_quant,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.moe.triton_pad_expert_counts import pad_expert_counts
|
from sglang.kernels.ops.moe.triton_pad_expert_counts import pad_expert_counts
|
||||||
from sglang.kernels.ops.quantization import per_token_group_quant
|
from sglang.kernels.ops.quantization.per_token_group_quant import per_token_group_quant
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -269,6 +269,9 @@ class DeepGemmRunnerInput(RunnerInput):
|
|||||||
expected_m: Optional[int] = None
|
expected_m: Optional[int] = None
|
||||||
m_indices: Optional[torch.Tensor] = None
|
m_indices: Optional[torch.Tensor] = None
|
||||||
hidden_states_scale_tma_aligned: bool = False
|
hidden_states_scale_tma_aligned: bool = False
|
||||||
|
# Number of activation elements sharing one scale along K.
|
||||||
|
# Records the actual input quantization group, independently of weight scales.
|
||||||
|
activation_scale_block_size: Optional[int] = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def runner_backend(self) -> MoeRunnerBackend:
|
def runner_backend(self) -> MoeRunnerBackend:
|
||||||
@@ -306,10 +309,43 @@ class DeepGemmMoeQuantInfo(MoeQuantInfo):
|
|||||||
"MXFP8 requires DEEPGEMM_SCALE_UE8M0=True"
|
"MXFP8 requires DEEPGEMM_SCALE_UE8M0=True"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
def scale_recipes(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
activation_block_size: Optional[int],
|
||||||
|
hidden_size: int,
|
||||||
|
activation_scale_width: int,
|
||||||
|
) -> tuple[Optional[tuple[int, int]], Optional[tuple[int, int]]]:
|
||||||
|
"""Return DeepGEMM A/B scale recipes from explicit layout metadata."""
|
||||||
|
if self.use_mxfp8:
|
||||||
|
assert self.block_shape is not None
|
||||||
|
weight_recipe = (self.block_shape[0], self.block_shape[1])
|
||||||
|
activation_block_size = activation_block_size or self.block_shape[1]
|
||||||
|
assert ceil_div(hidden_size, activation_block_size * 4) == (
|
||||||
|
activation_scale_width
|
||||||
|
), (
|
||||||
|
"MXFP8 activation scale mismatch: "
|
||||||
|
f"block_size={activation_block_size}, K={hidden_size}, "
|
||||||
|
f"scale_width={activation_scale_width}, expected "
|
||||||
|
f"{ceil_div(hidden_size, activation_block_size * 4)}"
|
||||||
|
)
|
||||||
|
return (self.block_shape[0], activation_block_size), weight_recipe
|
||||||
|
if self.is_fp4_experts:
|
||||||
|
return (1, 128), (1, 32)
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
|
||||||
class DeepGemmRunnerCore(MoeRunnerCore):
|
class DeepGemmRunnerCore(MoeRunnerCore):
|
||||||
def __init__(self, config: MoeRunnerConfig):
|
def __init__(self, config: MoeRunnerConfig):
|
||||||
super().__init__(config)
|
super().__init__(config)
|
||||||
|
if config.silu_mul_keep_fp32 and (
|
||||||
|
config.activation != "silu"
|
||||||
|
or not config.is_gated
|
||||||
|
or config.gemm1_alpha is not None
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"silu_mul_keep_fp32 requires gated SiLU without gemm1_alpha"
|
||||||
|
)
|
||||||
# SiTU (Kimi K3) is applied outside the GEMMs in python, so it only
|
# 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).
|
# needs the masked-gemm activation site to branch (see _run_masked_gemm).
|
||||||
assert self.config.activation in ("silu", "situ")
|
assert self.config.activation in ("silu", "situ")
|
||||||
@@ -329,6 +365,8 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
hooks: Optional[Any] = None,
|
hooks: Optional[Any] = None,
|
||||||
) -> DeepGemmRunnerOutput:
|
) -> DeepGemmRunnerOutput:
|
||||||
weight_dtype = quant_info.w13_weight.dtype
|
weight_dtype = quant_info.w13_weight.dtype
|
||||||
|
if self.config.silu_mul_keep_fp32 and weight_dtype != torch.float8_e4m3fn:
|
||||||
|
raise ValueError("silu_mul_keep_fp32 requires FP8 expert weights")
|
||||||
alignment = (
|
alignment = (
|
||||||
running_state.get("contiguous_layout_alignment")
|
running_state.get("contiguous_layout_alignment")
|
||||||
if not runner_input.use_masked_gemm
|
if not runner_input.use_masked_gemm
|
||||||
@@ -392,12 +430,11 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
(0, K), device=hidden_states_device, dtype=torch.bfloat16
|
(0, K), device=hidden_states_device, dtype=torch.bfloat16
|
||||||
)
|
)
|
||||||
|
|
||||||
if quant_info.use_mxfp8:
|
recipe_a, recipe_b = quant_info.scale_recipes(
|
||||||
recipe_a = recipe_b = tuple(quant_info.block_shape)
|
activation_block_size=runner_input.activation_scale_block_size,
|
||||||
elif quant_info.is_fp4_experts:
|
hidden_size=K,
|
||||||
recipe_a, recipe_b = (1, 128), (1, 32)
|
activation_scale_width=hidden_states_scale.shape[-1],
|
||||||
else:
|
)
|
||||||
recipe_a, recipe_b = None, None
|
|
||||||
|
|
||||||
w13_weight_fp8 = (
|
w13_weight_fp8 = (
|
||||||
quant_info.w13_weight,
|
quant_info.w13_weight,
|
||||||
@@ -485,7 +522,7 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
scale_ue8m0=False,
|
scale_ue8m0=False,
|
||||||
)
|
)
|
||||||
del down_input
|
del down_input
|
||||||
elif self.use_swizzle:
|
elif self.use_swizzle or self.config.silu_mul_keep_fp32:
|
||||||
swiglu_limit_arg: Optional[float] = self.swiglu_limit
|
swiglu_limit_arg: Optional[float] = self.swiglu_limit
|
||||||
use_contig_swizzle = self.use_swizzle and not running_state.get(
|
use_contig_swizzle = self.use_swizzle and not running_state.get(
|
||||||
"deepep_v2_disable_contig_swizzle", False
|
"deepep_v2_disable_contig_swizzle", False
|
||||||
@@ -569,10 +606,11 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
if deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES:
|
if deep_gemm_wrapper.DEEPGEMM_NEED_TMA_ALIGNED_SCALES:
|
||||||
down_input_scale = tma_align_input_scale(down_input_scale)
|
down_input_scale = tma_align_input_scale(down_input_scale)
|
||||||
|
|
||||||
recipe_a_down = (
|
# The down activation is quantized here, independently of dispatch.
|
||||||
(quant_info.block_shape[0], scale_block_size)
|
recipe_a_down, _ = quant_info.scale_recipes(
|
||||||
if quant_info.use_mxfp8
|
activation_block_size=scale_block_size,
|
||||||
else recipe_a
|
hidden_size=down_input_fp8.shape[-1],
|
||||||
|
activation_scale_width=down_input_scale.shape[-1],
|
||||||
)
|
)
|
||||||
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_contig(
|
deep_gemm_wrapper.grouped_gemm_nt_f8f8bf16_contig(
|
||||||
(down_input_fp8, down_input_scale),
|
(down_input_fp8, down_input_scale),
|
||||||
@@ -695,25 +733,11 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
use_mxfp8 = quant_info.use_mxfp8
|
use_mxfp8 = quant_info.use_mxfp8
|
||||||
scale_block_size = quant_info.block_shape[1] if quant_info.block_shape else 128
|
scale_block_size = quant_info.block_shape[1] if quant_info.block_shape else 128
|
||||||
|
|
||||||
if use_mxfp8:
|
recipe_a, recipe_b = quant_info.scale_recipes(
|
||||||
recipe_b = tuple(quant_info.block_shape)
|
activation_block_size=runner_input.activation_scale_block_size,
|
||||||
# gran_k is set by the dispatch path (standard=block_shape[1], DeepEP-LL=128),
|
hidden_size=hidden_states.shape[-1],
|
||||||
# not inferable from K; inferring it silently mis-reads the activation scale.
|
activation_scale_width=hidden_states_scale.shape[-1],
|
||||||
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
|
# GroupGemm-0
|
||||||
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
|
if deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
|
||||||
@@ -767,10 +791,11 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
|
|
||||||
# Act.
|
# Act.
|
||||||
if self.config.activation == "situ":
|
if self.config.activation == "situ":
|
||||||
|
scale_block_size = 128
|
||||||
down_input, down_input_scale = _varlen_deep_gemm_situ_mul_quant(
|
down_input, down_input_scale = _varlen_deep_gemm_situ_mul_quant(
|
||||||
gateup_output,
|
gateup_output,
|
||||||
masked_m,
|
masked_m,
|
||||||
group_size=128,
|
group_size=scale_block_size,
|
||||||
topk=self.config.top_k,
|
topk=self.config.top_k,
|
||||||
beta=self.config.gemm1_alpha,
|
beta=self.config.gemm1_alpha,
|
||||||
linear_beta=self.config.gemm1_clamp_limit,
|
linear_beta=self.config.gemm1_clamp_limit,
|
||||||
@@ -797,18 +822,13 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
gemm1_alpha=self.config.gemm1_alpha,
|
gemm1_alpha=self.config.gemm1_alpha,
|
||||||
gemm1_clamp_limit=self.config.gemm1_clamp_limit,
|
gemm1_clamp_limit=self.config.gemm1_clamp_limit,
|
||||||
num_real_tokens=num_real_tokens,
|
num_real_tokens=num_real_tokens,
|
||||||
|
silu_mul_keep_fp32=self.config.silu_mul_keep_fp32,
|
||||||
)
|
)
|
||||||
if trace_deepep_v2_masked:
|
if trace_deepep_v2_masked:
|
||||||
torch.cuda.synchronize()
|
torch.cuda.synchronize()
|
||||||
logger.warning("DeepEP v2 masked runner activation returned")
|
logger.warning("DeepEP v2 masked runner activation returned")
|
||||||
del gateup_output
|
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
|
# GroupGemm-1
|
||||||
n = w2_weight.shape[1]
|
n = w2_weight.shape[1]
|
||||||
|
|
||||||
@@ -829,6 +849,11 @@ class DeepGemmRunnerCore(MoeRunnerCore):
|
|||||||
down_input_scale
|
down_input_scale
|
||||||
)
|
)
|
||||||
|
|
||||||
|
recipe_a_down, _ = quant_info.scale_recipes(
|
||||||
|
activation_block_size=scale_block_size,
|
||||||
|
hidden_size=down_input.shape[-1],
|
||||||
|
activation_scale_width=down_input_scale.shape[-1],
|
||||||
|
)
|
||||||
with use_symmetric_memory(
|
with use_symmetric_memory(
|
||||||
get_tp_group(), disabled=not is_allocation_symmetric()
|
get_tp_group(), disabled=not is_allocation_symmetric()
|
||||||
):
|
):
|
||||||
@@ -1024,16 +1049,15 @@ def pre_permute_standard_to_deep_gemm(
|
|||||||
running_state["hidden_states_dtype"] = hidden_states_dtype
|
running_state["hidden_states_dtype"] = hidden_states_dtype
|
||||||
running_state["hidden_states_device"] = hidden_states_device
|
running_state["hidden_states_device"] = hidden_states_device
|
||||||
running_state["src2dst"] = src2dst
|
running_state["src2dst"] = src2dst
|
||||||
running_state["mxfp8_act_gran_k"] = (
|
|
||||||
quant_info.block_shape[1] if quant_info.block_shape else 128
|
|
||||||
)
|
|
||||||
|
|
||||||
return DeepGemmRunnerInput(
|
return DeepGemmRunnerInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
hidden_states_scale=hidden_states_scale,
|
hidden_states_scale=hidden_states_scale,
|
||||||
use_masked_gemm=True,
|
use_masked_gemm=True,
|
||||||
masked_m=masked_m,
|
masked_m=masked_m,
|
||||||
expected_m=expected_m,
|
expected_m=expected_m,
|
||||||
|
activation_scale_block_size=(
|
||||||
|
quant_info.block_shape[1] if quant_info.block_shape else 128
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# The compact layout avoids scaling masked buffers with the expert count.
|
# The compact layout avoids scaling masked buffers with the expert count.
|
||||||
@@ -1157,15 +1181,15 @@ def pre_permute_standard_to_deep_gemm(
|
|||||||
running_state["src2dst"] = src2dst
|
running_state["src2dst"] = src2dst
|
||||||
running_state["all_tokens"] = all_tokens
|
running_state["all_tokens"] = all_tokens
|
||||||
running_state["contiguous_layout_alignment"] = block_e
|
running_state["contiguous_layout_alignment"] = block_e
|
||||||
running_state["mxfp8_act_gran_k"] = (
|
|
||||||
quant_info.block_shape[1] if quant_info.block_shape else 128
|
|
||||||
)
|
|
||||||
|
|
||||||
return DeepGemmRunnerInput(
|
return DeepGemmRunnerInput(
|
||||||
hidden_states=packed_input,
|
hidden_states=packed_input,
|
||||||
hidden_states_scale=packed_input_scale,
|
hidden_states_scale=packed_input_scale,
|
||||||
use_masked_gemm=False,
|
use_masked_gemm=False,
|
||||||
m_indices=m_indices,
|
m_indices=m_indices,
|
||||||
|
activation_scale_block_size=(
|
||||||
|
quant_info.block_shape[1] if quant_info.block_shape else 128
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -1301,15 +1325,13 @@ def pre_permute_deepep_ll_to_deep_gemm(
|
|||||||
running_state["hidden_states_shape"] = hidden_states.shape
|
running_state["hidden_states_shape"] = hidden_states.shape
|
||||||
running_state["hidden_states_dtype"] = hidden_states.dtype
|
running_state["hidden_states_dtype"] = hidden_states.dtype
|
||||||
running_state["hidden_states_device"] = hidden_states.device
|
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(
|
return DeepGemmRunnerInput(
|
||||||
hidden_states=hidden_states,
|
hidden_states=hidden_states,
|
||||||
hidden_states_scale=hidden_states_scale,
|
hidden_states_scale=hidden_states_scale,
|
||||||
use_masked_gemm=True,
|
use_masked_gemm=True,
|
||||||
masked_m=masked_m,
|
masked_m=masked_m,
|
||||||
expected_m=expected_m,
|
expected_m=expected_m,
|
||||||
|
activation_scale_block_size=128,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -1428,6 +1450,7 @@ def pre_permute_deepep_normal_to_deep_gemm(
|
|||||||
hidden_states_scale=input_tensor_scale,
|
hidden_states_scale=input_tensor_scale,
|
||||||
use_masked_gemm=False,
|
use_masked_gemm=False,
|
||||||
m_indices=m_indices,
|
m_indices=m_indices,
|
||||||
|
activation_scale_block_size=128,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -1517,6 +1540,7 @@ def _varlen_deep_gemm_silu_mul_quant(
|
|||||||
gemm1_alpha: Optional[float] = None,
|
gemm1_alpha: Optional[float] = None,
|
||||||
gemm1_clamp_limit: Optional[float] = None,
|
gemm1_clamp_limit: Optional[float] = None,
|
||||||
num_real_tokens: Optional[int] = None,
|
num_real_tokens: Optional[int] = None,
|
||||||
|
silu_mul_keep_fp32: bool = False,
|
||||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
assert masked_m is not None
|
assert masked_m is not None
|
||||||
hidden_states_device = gateup_output.device
|
hidden_states_device = gateup_output.device
|
||||||
@@ -1525,6 +1549,9 @@ def _varlen_deep_gemm_silu_mul_quant(
|
|||||||
del D_2
|
del D_2
|
||||||
G = D // group_size
|
G = D // group_size
|
||||||
|
|
||||||
|
if silu_mul_keep_fp32 and gemm1_alpha is not None:
|
||||||
|
raise ValueError("silu_mul_keep_fp32 does not support gemm1_alpha")
|
||||||
|
|
||||||
# oai-swiglu (gemm1_alpha) stays on the Triton kernel until
|
# 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
|
# 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
|
# int32 UE8M0 (no follow-up transform; needs G % 4 == 0 and the
|
||||||
@@ -1568,9 +1595,9 @@ def _varlen_deep_gemm_silu_mul_quant(
|
|||||||
down_input_scale = down_input_scale.transpose(-1, -2)
|
down_input_scale = down_input_scale.transpose(-1, -2)
|
||||||
return down_input, down_input_scale
|
return down_input, down_input_scale
|
||||||
|
|
||||||
# DSV4-specific activations (clamped swiglu, swizzled gate|up layout) stay
|
# Only explicit precision requests opt additional callers into this kernel;
|
||||||
# on the DSV4 JIT kernel; it is the only implementation carrying them.
|
# the generic fused quantizer rounds its SiLU intermediates to BF16.
|
||||||
if swiglu_limit is not None or swizzle:
|
if swiglu_limit is not None or swizzle or silu_mul_keep_fp32:
|
||||||
assert N % 4 == 0 and G % 4 == 0 and D // 8 >= E, (
|
assert N % 4 == 0 and G % 4 == 0 and D // 8 >= E, (
|
||||||
"DSV4 JIT activation requires N % 4 == 0, G % 4 == 0 and "
|
"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}"
|
f"D // 8 >= num_experts, got N={N} G={G} D={D} E={E}"
|
||||||
@@ -1739,6 +1766,9 @@ def pre_permute_deepep_v2_to_deep_gemm(
|
|||||||
use_masked_gemm=True,
|
use_masked_gemm=True,
|
||||||
masked_m=masked_m,
|
masked_m=masked_m,
|
||||||
expected_m=deepep_v2_expected_m,
|
expected_m=deepep_v2_expected_m,
|
||||||
|
activation_scale_block_size=(
|
||||||
|
dispatch_output.activation_scale_block_size
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
# Mark aligned expert rows and leave the unused receive tail at -1.
|
# Mark aligned expert rows and leave the unused receive tail at -1.
|
||||||
@@ -1752,10 +1782,12 @@ def pre_permute_deepep_v2_to_deep_gemm(
|
|||||||
use_masked_gemm=False,
|
use_masked_gemm=False,
|
||||||
m_indices=m_indices,
|
m_indices=m_indices,
|
||||||
hidden_states_scale_tma_aligned=hidden_states_scale_tma_aligned,
|
hidden_states_scale_tma_aligned=hidden_states_scale_tma_aligned,
|
||||||
|
activation_scale_block_size=dispatch_output.activation_scale_block_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
all_tokens = int(psum_num_recv_tokens_per_expert[-1].item())
|
all_tokens = int(psum_num_recv_tokens_per_expert[-1].item())
|
||||||
K = hidden_states.shape[1]
|
K = hidden_states.shape[1]
|
||||||
|
scale_block_size = dispatch_output.activation_scale_block_size
|
||||||
running_state["all_tokens"] = all_tokens
|
running_state["all_tokens"] = all_tokens
|
||||||
running_state["hidden_states_shape"] = hidden_states.shape
|
running_state["hidden_states_shape"] = hidden_states.shape
|
||||||
running_state["hidden_states_device"] = hidden_states.device
|
running_state["hidden_states_device"] = hidden_states.device
|
||||||
@@ -1771,13 +1803,15 @@ def pre_permute_deepep_v2_to_deep_gemm(
|
|||||||
elif deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
|
elif deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0:
|
||||||
# Packed UE8M0 scales require zero padding lanes.
|
# Packed UE8M0 scales require zero padding lanes.
|
||||||
input_tensor_scale = torch.zeros(
|
input_tensor_scale = torch.zeros(
|
||||||
(ceil_div(K // 128, 4), all_tokens),
|
(ceil_div(K // scale_block_size, 4), all_tokens),
|
||||||
device=hidden_states.device,
|
device=hidden_states.device,
|
||||||
dtype=torch.int,
|
dtype=torch.int,
|
||||||
).transpose(0, 1)
|
).transpose(0, 1)
|
||||||
else:
|
else:
|
||||||
input_tensor_scale = torch.empty(
|
input_tensor_scale = torch.empty(
|
||||||
(all_tokens, K // 128), device=hidden_states.device, dtype=torch.float32
|
(all_tokens, K // scale_block_size),
|
||||||
|
device=hidden_states.device,
|
||||||
|
dtype=torch.float32,
|
||||||
)
|
)
|
||||||
m_indices = torch.empty(all_tokens, device=hidden_states.device, dtype=torch.int32)
|
m_indices = torch.empty(all_tokens, device=hidden_states.device, dtype=torch.int32)
|
||||||
output_index = torch.empty_like(topk_ids)
|
output_index = torch.empty_like(topk_ids)
|
||||||
@@ -1794,6 +1828,7 @@ def pre_permute_deepep_v2_to_deep_gemm(
|
|||||||
m_indices,
|
m_indices,
|
||||||
output_index,
|
output_index,
|
||||||
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
|
scale_ue8m0=deep_gemm_wrapper.DEEPGEMM_SCALE_UE8M0,
|
||||||
|
quant_block_size=scale_block_size,
|
||||||
)
|
)
|
||||||
dispose_tensor(hidden_states)
|
dispose_tensor(hidden_states)
|
||||||
if hidden_states_scale is not None:
|
if hidden_states_scale is not None:
|
||||||
@@ -1805,6 +1840,7 @@ def pre_permute_deepep_v2_to_deep_gemm(
|
|||||||
hidden_states_scale=input_tensor_scale,
|
hidden_states_scale=input_tensor_scale,
|
||||||
use_masked_gemm=False,
|
use_masked_gemm=False,
|
||||||
m_indices=m_indices,
|
m_indices=m_indices,
|
||||||
|
activation_scale_block_size=dispatch_output.activation_scale_block_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -1816,23 +1852,42 @@ def post_permute_deep_gemm_to_deepep_v2(
|
|||||||
running_state: dict,
|
running_state: dict,
|
||||||
) -> DeepEPv2CombineInput:
|
) -> DeepEPv2CombineInput:
|
||||||
from sglang.kernels.ops.moe.ep_moe_kernels import ep_gather
|
from sglang.kernels.ops.moe.ep_moe_kernels import ep_gather
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.base import RoutewiseLayout
|
||||||
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import DeepEPv2CombineInput
|
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import DeepEPv2CombineInput
|
||||||
|
|
||||||
|
return_unweighted_routes = runner_config.no_combine
|
||||||
if running_state.get("deepep_v2_expanded", False):
|
if running_state.get("deepep_v2_expanded", False):
|
||||||
hidden_states = runner_output.hidden_states
|
hidden_states = runner_output.hidden_states
|
||||||
topk_weights = running_state["topk_weights"]
|
topk_weights = running_state["topk_weights"]
|
||||||
if running_state.get("deepep_v2_masked", False):
|
if running_state.get("deepep_v2_masked", False):
|
||||||
# Expanded combine does not consume top-k weights.
|
# A routewise finalizer must run before router weighting. Preserve
|
||||||
|
# one raw row per route and carry its 1-D weight to that finalizer.
|
||||||
from sglang.kernels.ops.moe.ep_moe_kernels import masked_slab_to_expand
|
from sglang.kernels.ops.moe.ep_moe_kernels import masked_slab_to_expand
|
||||||
|
|
||||||
|
output_capacity = running_state["deepep_v2_total_expanded"]
|
||||||
|
if topk_weights.ndim != 1 or topk_weights.shape[0] < output_capacity:
|
||||||
|
raise ValueError(
|
||||||
|
"DeepEP v2 expanded output exceeds router-weight capacity"
|
||||||
|
)
|
||||||
hidden_states = masked_slab_to_expand(
|
hidden_states = masked_slab_to_expand(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
running_state["deepep_v2_psum"],
|
running_state["deepep_v2_psum"],
|
||||||
running_state["deepep_v2_total_expanded"],
|
output_capacity,
|
||||||
running_state["deepep_v2_expert_alignment"],
|
running_state["deepep_v2_expert_alignment"],
|
||||||
topk_weights=topk_weights,
|
topk_weights=None if return_unweighted_routes else topk_weights,
|
||||||
|
)
|
||||||
|
if not return_unweighted_routes:
|
||||||
|
return DeepEPv2CombineInput(hidden_states, None)
|
||||||
|
# Match the communication-capacity weights to the output slab.
|
||||||
|
return DeepEPv2CombineInput(
|
||||||
|
hidden_states=hidden_states,
|
||||||
|
topk_weights=topk_weights[: hidden_states.shape[0]],
|
||||||
|
routewise_layout=RoutewiseLayout.EXPANDED,
|
||||||
|
)
|
||||||
|
if return_unweighted_routes:
|
||||||
|
return DeepEPv2CombineInput(
|
||||||
|
hidden_states, topk_weights, RoutewiseLayout.EXPANDED
|
||||||
)
|
)
|
||||||
return DeepEPv2CombineInput(hidden_states, None)
|
|
||||||
if topk_weights is not None:
|
if topk_weights is not None:
|
||||||
# Expanded combine does not consume top-k weights.
|
# Expanded combine does not consume top-k weights.
|
||||||
hidden_states = hidden_states * topk_weights.to(
|
hidden_states = hidden_states * topk_weights.to(
|
||||||
@@ -1844,6 +1899,24 @@ def post_permute_deep_gemm_to_deepep_v2(
|
|||||||
topk_ids = running_state["topk_ids"]
|
topk_ids = running_state["topk_ids"]
|
||||||
topk_weights = running_state["topk_weights"]
|
topk_weights = running_state["topk_weights"]
|
||||||
output_index = running_state["output_index"]
|
output_index = running_state["output_index"]
|
||||||
|
if return_unweighted_routes:
|
||||||
|
# Restore the route dimension required by a routewise finalizer.
|
||||||
|
# output_index maps each received token/expert slot back to the compact
|
||||||
|
# expert-sorted DeepGEMM output; -1 denotes a non-local route.
|
||||||
|
valid = output_index >= 0
|
||||||
|
if hidden_states.shape[0] == 0:
|
||||||
|
route_out = hidden_states.new_zeros(
|
||||||
|
(*output_index.shape, hidden_states.shape[-1])
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
safe_output_index = output_index.clamp_min(0).to(torch.int64)
|
||||||
|
route_out = hidden_states[safe_output_index]
|
||||||
|
route_out.masked_fill_(~valid.unsqueeze(-1), 0)
|
||||||
|
return DeepEPv2CombineInput(
|
||||||
|
hidden_states=route_out,
|
||||||
|
topk_weights=topk_weights,
|
||||||
|
routewise_layout=RoutewiseLayout.TOKEN_TOPK,
|
||||||
|
)
|
||||||
gather_out = torch.empty(
|
gather_out = torch.empty(
|
||||||
running_state["hidden_states_shape"],
|
running_state["hidden_states_shape"],
|
||||||
device=running_state["hidden_states_device"],
|
device=running_state["hidden_states_device"],
|
||||||
|
|||||||
@@ -195,4 +195,5 @@ def _pre_permute_standard_contig(
|
|||||||
hidden_states_scale=input_tensor_scale,
|
hidden_states_scale=input_tensor_scale,
|
||||||
use_masked_gemm=False,
|
use_masked_gemm=False,
|
||||||
m_indices=m_indices,
|
m_indices=m_indices,
|
||||||
|
activation_scale_block_size=128,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -61,6 +61,12 @@ class MoeRunner:
|
|||||||
self.config = config
|
self.config = config
|
||||||
self.lora_enabled = lora_enabled
|
self.lora_enabled = lora_enabled
|
||||||
|
|
||||||
|
if config.silu_mul_keep_fp32 and not runner_backend.is_deep_gemm():
|
||||||
|
raise ValueError(
|
||||||
|
"silu_mul_keep_fp32 is currently supported only by deep_gemm, "
|
||||||
|
f"got {runner_backend.value}"
|
||||||
|
)
|
||||||
|
|
||||||
# --moe-runner-backend hpc_ops makes the standard dispatcher keep
|
# --moe-runner-backend hpc_ops makes the standard dispatcher keep
|
||||||
# global expert ids (skip_local_expert_mapping), so every MoE layer
|
# global expert ids (skip_local_expert_mapping), so every MoE layer
|
||||||
# must actually run the hpc_ops runner. A quant method that falls
|
# must actually run the hpc_ops runner. A quant method that falls
|
||||||
|
|||||||
@@ -12,6 +12,8 @@ from sglang.srt.layers.moe.token_dispatcher.base import (
|
|||||||
DispatchOutput,
|
DispatchOutput,
|
||||||
DispatchOutputChecker,
|
DispatchOutputChecker,
|
||||||
DispatchOutputFormat,
|
DispatchOutputFormat,
|
||||||
|
RoutewiseCombineInput,
|
||||||
|
RoutewiseLayout,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.token_dispatcher.deepep import (
|
from sglang.srt.layers.moe.token_dispatcher.deepep import (
|
||||||
DeepEPConfig,
|
DeepEPConfig,
|
||||||
@@ -67,6 +69,8 @@ __all__ = [
|
|||||||
"DispatchOutput",
|
"DispatchOutput",
|
||||||
"DispatchOutputFormat",
|
"DispatchOutputFormat",
|
||||||
"DispatchOutputChecker",
|
"DispatchOutputChecker",
|
||||||
|
"RoutewiseCombineInput",
|
||||||
|
"RoutewiseLayout",
|
||||||
"FlashinferDispatchOutput",
|
"FlashinferDispatchOutput",
|
||||||
"FlashinferDispatcher",
|
"FlashinferDispatcher",
|
||||||
"MooncakeCombineInput",
|
"MooncakeCombineInput",
|
||||||
|
|||||||
@@ -214,6 +214,16 @@ class DispatchOutput(Protocol):
|
|||||||
|
|
||||||
|
|
||||||
class CombineInputChecker:
|
class CombineInputChecker:
|
||||||
|
@staticmethod
|
||||||
|
def needs_model_route_finalization(
|
||||||
|
combine_input: CombineInput,
|
||||||
|
) -> TypeGuard[RoutewiseCombineInput]:
|
||||||
|
"""Whether expert output still needs routewise model finalization."""
|
||||||
|
return (
|
||||||
|
isinstance(combine_input, RoutewiseCombineInput)
|
||||||
|
and combine_input.routewise_layout is not None
|
||||||
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def format_is_standard(
|
def format_is_standard(
|
||||||
combine_input: CombineInput,
|
combine_input: CombineInput,
|
||||||
@@ -269,6 +279,23 @@ class CombineInputFormat(Enum):
|
|||||||
ASCEND_TP = "ascend_tp"
|
ASCEND_TP = "ascend_tp"
|
||||||
|
|
||||||
|
|
||||||
|
class RoutewiseLayout(Enum):
|
||||||
|
"""Layout of unweighted routes awaiting model finalization; H is hidden size.
|
||||||
|
|
||||||
|
TOKEN_TOPK: outputs [T, K, H], router weights [T, K]. T counts received
|
||||||
|
token rows on this EP rank; K is router top-k. Model finalization reduces
|
||||||
|
K to produce [T, H] before dispatcher combine.
|
||||||
|
|
||||||
|
EXPANDED: outputs [R, H], router weights [R]. Each valid row is one
|
||||||
|
token-expert route; R includes alignment padding and unused capacity.
|
||||||
|
Model finalization preserves row positions; dispatcher combine maps and
|
||||||
|
sums valid routes back to the original sender's tokens.
|
||||||
|
"""
|
||||||
|
|
||||||
|
TOKEN_TOPK = "token_topk"
|
||||||
|
EXPANDED = "expanded"
|
||||||
|
|
||||||
|
|
||||||
@runtime_checkable
|
@runtime_checkable
|
||||||
class CombineInput(Protocol):
|
class CombineInput(Protocol):
|
||||||
"""Protocol for combine inputs in different formats."""
|
"""Protocol for combine inputs in different formats."""
|
||||||
@@ -279,6 +306,15 @@ class CombineInput(Protocol):
|
|||||||
def format(self) -> CombineInputFormat: ...
|
def format(self) -> CombineInputFormat: ...
|
||||||
|
|
||||||
|
|
||||||
|
@runtime_checkable
|
||||||
|
class RoutewiseCombineInput(CombineInput, Protocol):
|
||||||
|
"""Combine input whose router weighting is deferred to the model layer."""
|
||||||
|
|
||||||
|
hidden_states: torch.Tensor
|
||||||
|
topk_weights: torch.Tensor
|
||||||
|
routewise_layout: RoutewiseLayout
|
||||||
|
|
||||||
|
|
||||||
# ------------------------------ Base Dispatcher -------------------------------------
|
# ------------------------------ Base Dispatcher -------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -15,6 +15,7 @@ from sglang.srt.layers.moe.token_dispatcher.base import (
|
|||||||
CombineInputFormat,
|
CombineInputFormat,
|
||||||
DispatchOutput,
|
DispatchOutput,
|
||||||
DispatchOutputFormat,
|
DispatchOutputFormat,
|
||||||
|
RoutewiseLayout,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.topk import TopKOutput
|
from sglang.srt.layers.moe.topk import TopKOutput
|
||||||
from sglang.srt.layers.moe.utils import (
|
from sglang.srt.layers.moe.utils import (
|
||||||
@@ -35,6 +36,7 @@ _deepep_v2_import_error: Optional[BaseException] = None
|
|||||||
_fp8_quant_import_error: Optional[BaseException] = None
|
_fp8_quant_import_error: Optional[BaseException] = None
|
||||||
sglang_per_token_group_quant_fp8 = None
|
sglang_per_token_group_quant_fp8 = None
|
||||||
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from deep_ep import ElasticBuffer
|
from deep_ep import ElasticBuffer
|
||||||
|
|
||||||
@@ -55,7 +57,7 @@ if use_deepep_v2:
|
|||||||
class DeepEPv2DispatchOutput(NamedTuple):
|
class DeepEPv2DispatchOutput(NamedTuple):
|
||||||
hidden_states: torch.Tensor
|
hidden_states: torch.Tensor
|
||||||
hidden_states_scale: Optional[torch.Tensor]
|
hidden_states_scale: Optional[torch.Tensor]
|
||||||
topk_ids: Optional[torch.Tensor]
|
topk_ids: Optional[torch.Tensor] # Receiver-local expert IDs, or -1.
|
||||||
topk_weights: torch.Tensor
|
topk_weights: torch.Tensor
|
||||||
psum_num_recv_tokens_per_expert: Optional[torch.Tensor] = None
|
psum_num_recv_tokens_per_expert: Optional[torch.Tensor] = None
|
||||||
is_expanded: bool = False
|
is_expanded: bool = False
|
||||||
@@ -65,6 +67,7 @@ class DeepEPv2DispatchOutput(NamedTuple):
|
|||||||
masked_max_m: int = 0
|
masked_max_m: int = 0
|
||||||
total_expanded: int = 0
|
total_expanded: int = 0
|
||||||
expert_alignment: int = 128
|
expert_alignment: int = 128
|
||||||
|
activation_scale_block_size: int = _SCALE_BLOCK_SIZE
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def format(self) -> DispatchOutputFormat:
|
def format(self) -> DispatchOutputFormat:
|
||||||
@@ -74,6 +77,7 @@ class DeepEPv2DispatchOutput(NamedTuple):
|
|||||||
class DeepEPv2CombineInput(NamedTuple):
|
class DeepEPv2CombineInput(NamedTuple):
|
||||||
hidden_states: torch.Tensor
|
hidden_states: torch.Tensor
|
||||||
topk_weights: Optional[torch.Tensor]
|
topk_weights: Optional[torch.Tensor]
|
||||||
|
routewise_layout: Optional[RoutewiseLayout] = None
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def format(self) -> CombineInputFormat:
|
def format(self) -> CombineInputFormat:
|
||||||
@@ -121,12 +125,14 @@ def _get_allow_hybrid_mode() -> bool:
|
|||||||
|
|
||||||
|
|
||||||
def _quantize_for_deepep_v2_dispatch(
|
def _quantize_for_deepep_v2_dispatch(
|
||||||
hidden_states: torch.Tensor, scale_format: DeepEPv2Fp8ScaleFormat
|
hidden_states: torch.Tensor,
|
||||||
|
scale_format: DeepEPv2Fp8ScaleFormat,
|
||||||
|
activation_scale_block_size: int = _SCALE_BLOCK_SIZE,
|
||||||
):
|
):
|
||||||
_ensure_fp8_quant_available()
|
_ensure_fp8_quant_available()
|
||||||
return sglang_per_token_group_quant_fp8(
|
return sglang_per_token_group_quant_fp8(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
_SCALE_BLOCK_SIZE,
|
activation_scale_block_size,
|
||||||
column_major_scales=scale_format.tma_aligned,
|
column_major_scales=scale_format.tma_aligned,
|
||||||
scale_tma_aligned=scale_format.tma_aligned,
|
scale_tma_aligned=scale_format.tma_aligned,
|
||||||
scale_ue8m0=scale_format.ue8m0,
|
scale_ue8m0=scale_format.ue8m0,
|
||||||
@@ -227,6 +233,7 @@ class _DeepEPv2Impl:
|
|||||||
scale_format: DeepEPv2Fp8ScaleFormat,
|
scale_format: DeepEPv2Fp8ScaleFormat,
|
||||||
num_max_dispatch_tokens_per_rank: int,
|
num_max_dispatch_tokens_per_rank: int,
|
||||||
use_fp8_dispatch: bool,
|
use_fp8_dispatch: bool,
|
||||||
|
activation_scale_block_size: int = _SCALE_BLOCK_SIZE,
|
||||||
):
|
):
|
||||||
self.group = group
|
self.group = group
|
||||||
self.router_topk = router_topk
|
self.router_topk = router_topk
|
||||||
@@ -234,9 +241,9 @@ class _DeepEPv2Impl:
|
|||||||
self.num_local_experts = num_local_experts
|
self.num_local_experts = num_local_experts
|
||||||
self.hidden_size = hidden_size
|
self.hidden_size = hidden_size
|
||||||
self.scale_format = scale_format
|
self.scale_format = scale_format
|
||||||
|
self.activation_scale_block_size = activation_scale_block_size
|
||||||
self.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
|
self.num_max_dispatch_tokens_per_rank = num_max_dispatch_tokens_per_rank
|
||||||
self.use_fp8_dispatch = use_fp8_dispatch
|
self.use_fp8_dispatch = use_fp8_dispatch
|
||||||
self.rank = dist.get_rank(group)
|
|
||||||
self._handle = None
|
self._handle = None
|
||||||
self._pad_empty_combine = False
|
self._pad_empty_combine = False
|
||||||
|
|
||||||
@@ -267,10 +274,10 @@ class _DeepEPv2Impl:
|
|||||||
f"DeepEP v2 hidden size mismatch: expected {self.hidden_size}, "
|
f"DeepEP v2 hidden size mismatch: expected {self.hidden_size}, "
|
||||||
f"got {hidden_states.shape[1]}"
|
f"got {hidden_states.shape[1]}"
|
||||||
)
|
)
|
||||||
if self.hidden_size % _SCALE_BLOCK_SIZE != 0:
|
if self.hidden_size % self.activation_scale_block_size != 0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"DeepEP v2 requires hidden_size multiple of "
|
"DeepEP v2 requires hidden_size multiple of "
|
||||||
f"{_SCALE_BLOCK_SIZE}, got {self.hidden_size}"
|
f"{self.activation_scale_block_size}, got {self.hidden_size}"
|
||||||
)
|
)
|
||||||
if topk_ids.shape[1] != self.router_topk:
|
if topk_ids.shape[1] != self.router_topk:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
@@ -312,7 +319,7 @@ class _DeepEPv2Impl:
|
|||||||
_ue8m0 = self.scale_format.ue8m0
|
_ue8m0 = self.scale_format.ue8m0
|
||||||
dispatch_x = sglang_per_token_group_quant_fp8(
|
dispatch_x = sglang_per_token_group_quant_fp8(
|
||||||
hidden_states,
|
hidden_states,
|
||||||
_SCALE_BLOCK_SIZE,
|
self.activation_scale_block_size,
|
||||||
column_major_scales=_ue8m0,
|
column_major_scales=_ue8m0,
|
||||||
scale_tma_aligned=_ue8m0,
|
scale_tma_aligned=_ue8m0,
|
||||||
scale_ue8m0=_ue8m0,
|
scale_ue8m0=_ue8m0,
|
||||||
@@ -320,7 +327,7 @@ class _DeepEPv2Impl:
|
|||||||
use_tma_aligned_col_major_sf = _ue8m0
|
use_tma_aligned_col_major_sf = _ue8m0
|
||||||
else:
|
else:
|
||||||
dispatch_x = _quantize_for_deepep_v2_dispatch(
|
dispatch_x = _quantize_for_deepep_v2_dispatch(
|
||||||
hidden_states, self.scale_format
|
hidden_states, self.scale_format, self.activation_scale_block_size
|
||||||
)
|
)
|
||||||
use_tma_aligned_col_major_sf = self.scale_format.tma_aligned
|
use_tma_aligned_col_major_sf = self.scale_format.tma_aligned
|
||||||
|
|
||||||
@@ -344,7 +351,6 @@ class _DeepEPv2Impl:
|
|||||||
do_cpu_sync=do_cpu_sync_val,
|
do_cpu_sync=do_cpu_sync_val,
|
||||||
do_expand=use_expand_layout,
|
do_expand=use_expand_layout,
|
||||||
)
|
)
|
||||||
self._handle = handle
|
|
||||||
local_tokens = hidden_states.shape[0]
|
local_tokens = hidden_states.shape[0]
|
||||||
if event.event is not None:
|
if event.event is not None:
|
||||||
event.current_stream_wait()
|
event.current_stream_wait()
|
||||||
@@ -368,12 +374,20 @@ class _DeepEPv2Impl:
|
|||||||
if recv_hidden_states_scale is not None:
|
if recv_hidden_states_scale is not None:
|
||||||
recv_hidden_states_scale = recv_hidden_states_scale[:num_recv_tokens]
|
recv_hidden_states_scale = recv_hidden_states_scale[:num_recv_tokens]
|
||||||
|
|
||||||
|
# ElasticBuffer already converts global router IDs to receiver-local
|
||||||
|
# expert IDs; applying this rank's offset again would discard routes.
|
||||||
local_topk_ids = recv_topk_idx
|
local_topk_ids = recv_topk_idx
|
||||||
|
|
||||||
expected_m = 0
|
expected_m = 0
|
||||||
masked_max_m = 0
|
masked_max_m = 0
|
||||||
total_expanded = 0
|
total_expanded = 0
|
||||||
if use_masked:
|
if use_masked:
|
||||||
|
recv_capacity = recv_hidden_states.shape[0]
|
||||||
|
if recv_topk_weights.shape != (recv_capacity,):
|
||||||
|
raise ValueError(
|
||||||
|
"DeepEP v2 expanded activations and router weights must "
|
||||||
|
"have the same receive capacity"
|
||||||
|
)
|
||||||
# expected_m is only a schedule hint; masked_m is the actual bound.
|
# expected_m is only a schedule hint; masked_m is the actual bound.
|
||||||
ep_group_size = max(1, self.num_experts // self.num_local_experts)
|
ep_group_size = max(1, self.num_experts // self.num_local_experts)
|
||||||
expected_m = max(
|
expected_m = max(
|
||||||
@@ -381,10 +395,12 @@ class _DeepEPv2Impl:
|
|||||||
(local_tokens * ep_group_size * self.router_topk + self.num_experts)
|
(local_tokens * ep_group_size * self.router_topk + self.num_experts)
|
||||||
// self.num_experts,
|
// self.num_experts,
|
||||||
)
|
)
|
||||||
# Account for the worst case where every rank targets one local expert.
|
# Every rank can send its full dispatch capacity to one expert.
|
||||||
masked_max_m = self.num_max_dispatch_tokens_per_rank * ep_group_size
|
masked_max_m = self.num_max_dispatch_tokens_per_rank * ep_group_size
|
||||||
total_expanded = recv_hidden_states.shape[0]
|
total_expanded = recv_capacity
|
||||||
|
|
||||||
|
# Publish ownership only after the returned metadata is validated.
|
||||||
|
self._handle = handle
|
||||||
return DeepEPv2DispatchOutput(
|
return DeepEPv2DispatchOutput(
|
||||||
recv_hidden_states,
|
recv_hidden_states,
|
||||||
recv_hidden_states_scale,
|
recv_hidden_states_scale,
|
||||||
@@ -398,6 +414,7 @@ class _DeepEPv2Impl:
|
|||||||
masked_max_m,
|
masked_max_m,
|
||||||
total_expanded,
|
total_expanded,
|
||||||
_EXPERT_ALIGNMENT,
|
_EXPERT_ALIGNMENT,
|
||||||
|
activation_scale_block_size=self.activation_scale_block_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
def combine(self, combine_input: DeepEPv2CombineInput) -> torch.Tensor:
|
def combine(self, combine_input: DeepEPv2CombineInput) -> torch.Tensor:
|
||||||
@@ -407,6 +424,10 @@ class _DeepEPv2Impl:
|
|||||||
)
|
)
|
||||||
# Release the single-use handle even when combine fails.
|
# Release the single-use handle even when combine fails.
|
||||||
try:
|
try:
|
||||||
|
if combine_input.routewise_layout is not None:
|
||||||
|
raise ValueError(
|
||||||
|
"DeepEP v2 combine requires model route finalization first"
|
||||||
|
)
|
||||||
buffer = self._get_buffer()
|
buffer = self._get_buffer()
|
||||||
combined_x, _, event = buffer.combine(
|
combined_x, _, event = buffer.combine(
|
||||||
combine_input.hidden_states,
|
combine_input.hidden_states,
|
||||||
@@ -433,6 +454,7 @@ class DeepEPv2Dispatcher(BaseDispatcher):
|
|||||||
hidden_size: int,
|
hidden_size: int,
|
||||||
params_dtype: torch.dtype,
|
params_dtype: torch.dtype,
|
||||||
use_fp8_dispatch: bool,
|
use_fp8_dispatch: bool,
|
||||||
|
activation_scale_block_size: int = _SCALE_BLOCK_SIZE,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
if params_dtype != torch.bfloat16:
|
if params_dtype != torch.bfloat16:
|
||||||
@@ -454,6 +476,7 @@ class DeepEPv2Dispatcher(BaseDispatcher):
|
|||||||
scale_format=scale_format,
|
scale_format=scale_format,
|
||||||
num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank,
|
num_max_dispatch_tokens_per_rank=self.num_max_dispatch_tokens_per_rank,
|
||||||
use_fp8_dispatch=use_fp8_dispatch,
|
use_fp8_dispatch=use_fp8_dispatch,
|
||||||
|
activation_scale_block_size=activation_scale_block_size,
|
||||||
)
|
)
|
||||||
|
|
||||||
def dispatch(
|
def dispatch(
|
||||||
|
|||||||
@@ -35,6 +35,35 @@ register_cuda_ci(est_time=20, stage="base-b-kernel-unit", runner_config="4-gpu-b
|
|||||||
dev = "cuda"
|
dev = "cuda"
|
||||||
|
|
||||||
|
|
||||||
|
def test_sm120_mxfp8_dispatch_preserves_activation_scale_recipe(monkeypatch):
|
||||||
|
"""SM120 group-128 activations must not use the MXFP8 weight-scale recipe."""
|
||||||
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
|
from sglang.srt.layers.moe.moe_runner import deep_gemm_sm120
|
||||||
|
|
||||||
|
monkeypatch.setattr(deep_gemm_sm120, "_is_sm120", True)
|
||||||
|
monkeypatch.setattr(deep_gemm_wrapper, "DEEPGEMM_SCALE_UE8M0", True)
|
||||||
|
config = MoeRunnerConfig(
|
||||||
|
num_experts=2, num_local_experts=2, top_k=1, hidden_size=512
|
||||||
|
)
|
||||||
|
quant = DeepGemmMoeQuantInfo(
|
||||||
|
torch.empty(1, dtype=torch.float8_e4m3fn),
|
||||||
|
None,
|
||||||
|
True,
|
||||||
|
block_shape=[1, 32],
|
||||||
|
use_mxfp8=True,
|
||||||
|
)
|
||||||
|
x = torch.randn(1024, 512, device=dev, dtype=torch.bfloat16)
|
||||||
|
ids = (torch.arange(1024, device=dev, dtype=torch.int32) % 2).view(-1, 1)
|
||||||
|
weights = torch.ones(1024, 1, device=dev)
|
||||||
|
result = deep_gemm_sm120.maybe_pre_permute(x, ids, weights, quant, config, {})
|
||||||
|
|
||||||
|
assert quant.scale_recipes(
|
||||||
|
activation_block_size=result.activation_scale_block_size,
|
||||||
|
hidden_size=result.hidden_states.shape[-1],
|
||||||
|
activation_scale_width=result.hidden_states_scale.shape[-1],
|
||||||
|
) == ((1, 128), (1, 32))
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("num_tokens", [1, 7, 64, 256])
|
@pytest.mark.parametrize("num_tokens", [1, 7, 64, 256])
|
||||||
@pytest.mark.parametrize("topk", [4, 5, 8])
|
@pytest.mark.parametrize("topk", [4, 5, 8])
|
||||||
@pytest.mark.parametrize("hidden,group", [(6144, 32), (2048, 32), (4096, 128)])
|
@pytest.mark.parametrize("hidden,group", [(6144, 32), (2048, 32), (4096, 128)])
|
||||||
|
|||||||
@@ -160,7 +160,9 @@ def test_ue8m0_bitexact(dtype, num_tokens, hidden):
|
|||||||
assert torch.equal(exp, exp_ref), "exponent bytes differ"
|
assert torch.equal(exp, exp_ref), "exponent bytes differ"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("group_size", get_ci_test_range([16, 32, 64, 128], [16, 64]))
|
@pytest.mark.parametrize(
|
||||||
|
"group_size", get_ci_test_range([16, 32, 64, 128], [16, 32, 64])
|
||||||
|
)
|
||||||
def test_ue8m0_group_sizes(group_size):
|
def test_ue8m0_group_sizes(group_size):
|
||||||
"""Group size is a template axis (v2 dispatched a runtime switch). Each size
|
"""Group size is a template axis (v2 dispatched a runtime switch). Each size
|
||||||
maps a group onto a different subwarp lane count; codes/exponents must stay
|
maps a group onto a different subwarp lane count; codes/exponents must stay
|
||||||
@@ -311,6 +313,104 @@ def _ref_silu_mul(x, hidden):
|
|||||||
return torch.nn.functional.silu(gate.float()).to(x.dtype) * up
|
return torch.nn.functional.silu(gate.float()).to(x.dtype) * up
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("group_size,hidden", [(32, 1792), (32, 6144), (128, 1024)])
|
||||||
|
@pytest.mark.parametrize("swiglu_limit", [None, 10.0])
|
||||||
|
def test_fp32_silu_post_quant(group_size, hidden, swiglu_limit):
|
||||||
|
"""The post-quant kernels keep SiLU and the multiply in FP32 until FP8.
|
||||||
|
|
||||||
|
Unlike the generic fused quantizer below, there is no intermediate BF16
|
||||||
|
round. Reuse the independent UE8M0 oracle, and compare both layouts only
|
||||||
|
on active rows; an empty expert and a partial slab exercise masked counts.
|
||||||
|
"""
|
||||||
|
from sglang.kernels.ops.attention.dsv4 import (
|
||||||
|
silu_and_mul_contig_post_quant,
|
||||||
|
silu_and_mul_masked_post_quant,
|
||||||
|
)
|
||||||
|
|
||||||
|
torch.manual_seed(123 + hidden)
|
||||||
|
experts, capacity = 3, 32
|
||||||
|
x = (
|
||||||
|
torch.randn(experts, capacity, hidden * 2, device="cuda", dtype=torch.bfloat16)
|
||||||
|
* 5
|
||||||
|
)
|
||||||
|
x[2, 0].zero_()
|
||||||
|
counts = torch.tensor([0, 17, 9], device="cuda", dtype=torch.int32)
|
||||||
|
gate, up = x.float().chunk(2, dim=-1)
|
||||||
|
if swiglu_limit is not None:
|
||||||
|
gate = gate.clamp_max(swiglu_limit)
|
||||||
|
up = up.clamp(-swiglu_limit, swiglu_limit)
|
||||||
|
activation = gate * torch.sigmoid(gate) * up
|
||||||
|
q_ref, exp_ref = ref_fp8_ue8m0(activation, group_size)
|
||||||
|
|
||||||
|
flat = x.flatten(0, 1)
|
||||||
|
q = torch.empty(experts * capacity, hidden, device="cuda", dtype=fp8_dtype)
|
||||||
|
scale = create_per_token_group_quant_fp8_output_scale(
|
||||||
|
x_shape=q.shape,
|
||||||
|
device="cuda",
|
||||||
|
group_size=group_size,
|
||||||
|
column_major_scales=True,
|
||||||
|
scale_tma_aligned=True,
|
||||||
|
scale_ue8m0=True,
|
||||||
|
)
|
||||||
|
silu_and_mul_contig_post_quant(
|
||||||
|
flat,
|
||||||
|
q,
|
||||||
|
scale,
|
||||||
|
group_size,
|
||||||
|
scale_ue8m0=True,
|
||||||
|
transposed=True,
|
||||||
|
swiglu_limit=swiglu_limit,
|
||||||
|
)
|
||||||
|
masked_q = torch.empty_like(q).view(experts, capacity, hidden)
|
||||||
|
masked_scale = torch.empty(
|
||||||
|
experts,
|
||||||
|
hidden // group_size // 4,
|
||||||
|
capacity,
|
||||||
|
device="cuda",
|
||||||
|
dtype=torch.int32,
|
||||||
|
)
|
||||||
|
silu_and_mul_masked_post_quant(
|
||||||
|
x,
|
||||||
|
masked_q,
|
||||||
|
masked_scale,
|
||||||
|
group_size,
|
||||||
|
counts,
|
||||||
|
scale_ue8m0=True,
|
||||||
|
transposed=True,
|
||||||
|
swiglu_limit=swiglu_limit,
|
||||||
|
)
|
||||||
|
exp = _decode_packed_exp(scale, hidden // group_size).view_as(exp_ref)
|
||||||
|
masked_exp = _decode_packed_exp(masked_scale.transpose(1, 2), hidden // group_size)
|
||||||
|
q = q.view_as(q_ref)
|
||||||
|
for expert, count in enumerate(counts.tolist()):
|
||||||
|
assert torch.equal(exp[expert, :count], exp_ref[expert, :count])
|
||||||
|
assert torch.equal(masked_exp[expert, :count], exp[expert, :count])
|
||||||
|
assert torch.equal(
|
||||||
|
masked_q[expert, :count].view(torch.uint8),
|
||||||
|
q[expert, :count].view(torch.uint8),
|
||||||
|
)
|
||||||
|
# Fast sigmoid may differ from torch by an FP32 ULP at an FP8
|
||||||
|
# rounding boundary; bound the resulting error, not arbitrary bytes.
|
||||||
|
torch.testing.assert_close(
|
||||||
|
q[expert, :count].float(),
|
||||||
|
q_ref[expert, :count].float(),
|
||||||
|
rtol=0.125,
|
||||||
|
atol=2**-9,
|
||||||
|
)
|
||||||
|
if count:
|
||||||
|
mismatch = (
|
||||||
|
(
|
||||||
|
q[expert, :count].view(torch.uint8)
|
||||||
|
!= q_ref[expert, :count].view(torch.uint8)
|
||||||
|
)
|
||||||
|
.float()
|
||||||
|
.mean()
|
||||||
|
)
|
||||||
|
# A rare fast-math boundary flip is allowed, but systematic BF16
|
||||||
|
# intermediate rounding (the other fused path) must fail this gate.
|
||||||
|
assert mismatch.item() < 1e-4
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize("column_major", [True, False])
|
@pytest.mark.parametrize("column_major", [True, False])
|
||||||
@pytest.mark.parametrize("scale_ue8m0", [True, False])
|
@pytest.mark.parametrize("scale_ue8m0", [True, False])
|
||||||
def test_fused_silu(scale_ue8m0, column_major):
|
def test_fused_silu(scale_ue8m0, column_major):
|
||||||
|
|||||||
@@ -144,6 +144,7 @@ class TestDeepEPv2BufferLifecycle(CustomTestCase):
|
|||||||
impl = object.__new__(deepep_v2._DeepEPv2Impl)
|
impl = object.__new__(deepep_v2._DeepEPv2Impl)
|
||||||
impl.num_max_dispatch_tokens_per_rank = 4
|
impl.num_max_dispatch_tokens_per_rank = 4
|
||||||
impl.hidden_size = 128
|
impl.hidden_size = 128
|
||||||
|
impl.activation_scale_block_size = 128
|
||||||
impl.router_topk = 2
|
impl.router_topk = 2
|
||||||
impl._validate_common(torch.empty(4, 128), torch.zeros(4, 2))
|
impl._validate_common(torch.empty(4, 128), torch.zeros(4, 2))
|
||||||
with self.assertRaisesRegex(ValueError, "per-rank buffer capacity"):
|
with self.assertRaisesRegex(ValueError, "per-rank buffer capacity"):
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
"""Tests for the DeepEP v2 expanded/masked repack kernels."""
|
"""Tests for the DeepEP v2 expanded/masked repack kernels."""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -113,6 +114,80 @@ class TestDeepEPv2MaskedSlab(CustomTestCase):
|
|||||||
def test_empty_experts(self):
|
def test_empty_experts(self):
|
||||||
self._check_expand_roundtrip([0, 0, 0, 0], torch.bfloat16, with_scale=False)
|
self._check_expand_roundtrip([0, 0, 0, 0], torch.bfloat16, with_scale=False)
|
||||||
|
|
||||||
|
def test_runner_defers_expanded_route_weighting(self):
|
||||||
|
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||||
|
from sglang.srt.layers.moe.moe_runner.deep_gemm import (
|
||||||
|
DeepGemmRunnerOutput,
|
||||||
|
post_permute_deep_gemm_to_deepep_v2,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.base import RoutewiseLayout
|
||||||
|
|
||||||
|
counts = [3, 0, 2]
|
||||||
|
recv_x, _, psum, starts, total = _build_layout(
|
||||||
|
counts, self.ALIGN, self.HIDDEN, torch.bfloat16
|
||||||
|
)
|
||||||
|
masked_x, _, _ = expand_to_masked_slab(
|
||||||
|
recv_x, None, psum, len(counts), self.MAX_M, self.ALIGN
|
||||||
|
)
|
||||||
|
weights = torch.full((total,), 0.25, device=DEVICE)
|
||||||
|
state = {
|
||||||
|
"deepep_v2_expanded": True,
|
||||||
|
"deepep_v2_masked": True,
|
||||||
|
"deepep_v2_psum": psum,
|
||||||
|
"deepep_v2_total_expanded": total,
|
||||||
|
"deepep_v2_expert_alignment": self.ALIGN,
|
||||||
|
"topk_weights": weights,
|
||||||
|
}
|
||||||
|
rows = _real_rows(starts, counts)
|
||||||
|
for no_combine in (False, True):
|
||||||
|
with self.subTest(no_combine=no_combine):
|
||||||
|
output = post_permute_deep_gemm_to_deepep_v2(
|
||||||
|
DeepGemmRunnerOutput(masked_x),
|
||||||
|
None,
|
||||||
|
MoeRunnerConfig(no_combine=no_combine),
|
||||||
|
state,
|
||||||
|
)
|
||||||
|
expected = recv_x[rows] if no_combine else recv_x[rows] * 0.25
|
||||||
|
self.assertTrue(torch.equal(output.hidden_states[rows], expected))
|
||||||
|
self.assertEqual(
|
||||||
|
output.routewise_layout,
|
||||||
|
RoutewiseLayout.EXPANDED if no_combine else None,
|
||||||
|
)
|
||||||
|
if no_combine:
|
||||||
|
self.assertTrue(torch.equal(output.topk_weights, weights))
|
||||||
|
|
||||||
|
def test_runner_restores_token_topk_routes_and_masks_nonlocal_slots(self):
|
||||||
|
from sglang.srt.layers.moe.moe_runner.base import MoeRunnerConfig
|
||||||
|
from sglang.srt.layers.moe.moe_runner.deep_gemm import (
|
||||||
|
DeepGemmRunnerOutput,
|
||||||
|
post_permute_deep_gemm_to_deepep_v2,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.base import RoutewiseLayout
|
||||||
|
|
||||||
|
hidden = torch.tensor([[3.0], [5.0], [7.0]], device=DEVICE)
|
||||||
|
weights = torch.tensor([[0.25, 0.0], [0.5, 0.75]], device=DEVICE)
|
||||||
|
state = {
|
||||||
|
"topk_ids": torch.tensor([[0, -1], [1, 0]], device=DEVICE),
|
||||||
|
"topk_weights": weights,
|
||||||
|
"output_index": torch.tensor([[1, -1], [0, 2]], device=DEVICE),
|
||||||
|
}
|
||||||
|
for empty in (False, True):
|
||||||
|
with self.subTest(empty=empty):
|
||||||
|
if empty:
|
||||||
|
state["output_index"] = torch.full((2, 2), -1, device=DEVICE)
|
||||||
|
output = post_permute_deep_gemm_to_deepep_v2(
|
||||||
|
DeepGemmRunnerOutput(hidden[:0] if empty else hidden),
|
||||||
|
None,
|
||||||
|
MoeRunnerConfig(no_combine=True),
|
||||||
|
state,
|
||||||
|
)
|
||||||
|
expected = torch.tensor([[[5.0], [0.0]], [[3.0], [7.0]]], device=DEVICE)
|
||||||
|
if empty:
|
||||||
|
expected.zero_()
|
||||||
|
self.assertTrue(torch.equal(output.hidden_states, expected))
|
||||||
|
self.assertIs(output.topk_weights, weights)
|
||||||
|
self.assertEqual(output.routewise_layout, RoutewiseLayout.TOKEN_TOPK)
|
||||||
|
|
||||||
def test_single_hot_expert(self):
|
def test_single_hot_expert(self):
|
||||||
self._check_expand_roundtrip(
|
self._check_expand_roundtrip(
|
||||||
[0, self.MAX_M, 0, 0], torch.bfloat16, with_scale=False, topk=True
|
[0, self.MAX_M, 0, 0], torch.bfloat16, with_scale=False, topk=True
|
||||||
@@ -133,7 +208,7 @@ class TestDeepEPv2MaskedSlab(CustomTestCase):
|
|||||||
recv_x, None, psum, len(counts), self.MAX_M, self.ALIGN
|
recv_x, None, psum, len(counts), self.MAX_M, self.ALIGN
|
||||||
)
|
)
|
||||||
|
|
||||||
def _production_packed_ue8m0_layout(self, counts):
|
def _production_packed_ue8m0_layout(self, counts, group_size):
|
||||||
"""Build expanded rows with the production packed UE8M0 quantizer."""
|
"""Build expanded rows with the production packed UE8M0 quantizer."""
|
||||||
from sglang.kernels.ops.quantization.fp8_kernel import (
|
from sglang.kernels.ops.quantization.fp8_kernel import (
|
||||||
sglang_per_token_group_quant_fp8,
|
sglang_per_token_group_quant_fp8,
|
||||||
@@ -146,7 +221,7 @@ class TestDeepEPv2MaskedSlab(CustomTestCase):
|
|||||||
)
|
)
|
||||||
recv_x, recv_x_scale = sglang_per_token_group_quant_fp8(
|
recv_x, recv_x_scale = sglang_per_token_group_quant_fp8(
|
||||||
raw,
|
raw,
|
||||||
128,
|
group_size,
|
||||||
column_major_scales=True,
|
column_major_scales=True,
|
||||||
scale_tma_aligned=True,
|
scale_tma_aligned=True,
|
||||||
scale_ue8m0=True,
|
scale_ue8m0=True,
|
||||||
@@ -157,9 +232,14 @@ class TestDeepEPv2MaskedSlab(CustomTestCase):
|
|||||||
return recv_x, recv_x_scale, psum, starts, total, hidden
|
return recv_x, recv_x_scale, psum, starts, total, hidden
|
||||||
|
|
||||||
def test_fp8_packed_ue8m0_scale_from_production_quantizer(self):
|
def test_fp8_packed_ue8m0_scale_from_production_quantizer(self):
|
||||||
|
for group_size in (32, 128):
|
||||||
|
with self.subTest(group_size=group_size):
|
||||||
|
self._check_packed_scale(group_size)
|
||||||
|
|
||||||
|
def _check_packed_scale(self, group_size):
|
||||||
counts = [3, 1, 6, 2]
|
counts = [3, 1, 6, 2]
|
||||||
recv_x, recv_x_scale, psum, starts, _, hidden = (
|
recv_x, recv_x_scale, psum, starts, _, hidden = (
|
||||||
self._production_packed_ue8m0_layout(counts)
|
self._production_packed_ue8m0_layout(counts, group_size)
|
||||||
)
|
)
|
||||||
E = len(counts)
|
E = len(counts)
|
||||||
masked_x, masked_x_scale, masked_m = expand_to_masked_slab(
|
masked_x, masked_x_scale, masked_m = expand_to_masked_slab(
|
||||||
@@ -175,10 +255,15 @@ class TestDeepEPv2MaskedSlab(CustomTestCase):
|
|||||||
torch.testing.assert_close(masked_x_scale[e, j], recv_x_scale[s + j])
|
torch.testing.assert_close(masked_x_scale[e, j], recv_x_scale[s + j])
|
||||||
|
|
||||||
def test_expand_under_cuda_graph_capture(self):
|
def test_expand_under_cuda_graph_capture(self):
|
||||||
|
for group_size in (32, 128):
|
||||||
|
with self.subTest(group_size=group_size):
|
||||||
|
self._check_graph_capture(group_size)
|
||||||
|
|
||||||
|
def _check_graph_capture(self, group_size):
|
||||||
# Exercise replay with the production packed scale layout.
|
# Exercise replay with the production packed scale layout.
|
||||||
counts = [3, 1, 6, 2]
|
counts = [3, 1, 6, 2]
|
||||||
recv_x, recv_x_scale, psum, starts, _, _ = self._production_packed_ue8m0_layout(
|
recv_x, recv_x_scale, psum, starts, _, _ = self._production_packed_ue8m0_layout(
|
||||||
counts
|
counts, group_size
|
||||||
)
|
)
|
||||||
E = len(counts)
|
E = len(counts)
|
||||||
warm = torch.cuda.Stream()
|
warm = torch.cuda.Stream()
|
||||||
@@ -235,10 +320,34 @@ class TestDeepEPv2HandleLifecycle(CustomTestCase):
|
|||||||
|
|
||||||
impl._get_buffer = _boom
|
impl._get_buffer = _boom
|
||||||
with self.assertRaisesRegex(RuntimeError, "boom"):
|
with self.assertRaisesRegex(RuntimeError, "boom"):
|
||||||
impl.combine(None)
|
impl.combine(SimpleNamespace(routewise_layout=None))
|
||||||
self.assertIsNone(impl._handle)
|
self.assertIsNone(impl._handle)
|
||||||
self.assertFalse(impl._pad_empty_combine)
|
self.assertFalse(impl._pad_empty_combine)
|
||||||
|
|
||||||
|
def test_unfinalized_routes_are_rejected_and_release_handle(self):
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.base import (
|
||||||
|
CombineInputChecker,
|
||||||
|
RoutewiseLayout,
|
||||||
|
)
|
||||||
|
from sglang.srt.layers.moe.token_dispatcher.deepep_v2 import (
|
||||||
|
DeepEPv2CombineInput,
|
||||||
|
)
|
||||||
|
|
||||||
|
impl = self._bare_impl()
|
||||||
|
impl._handle = object()
|
||||||
|
output = DeepEPv2CombineInput(
|
||||||
|
torch.empty(2, 8), torch.ones(2), RoutewiseLayout.EXPANDED
|
||||||
|
)
|
||||||
|
self.assertTrue(CombineInputChecker.needs_model_route_finalization(output))
|
||||||
|
with self.assertRaisesRegex(ValueError, "model route finalization"):
|
||||||
|
impl.combine(output)
|
||||||
|
self.assertIsNone(impl._handle)
|
||||||
|
self.assertFalse(
|
||||||
|
CombineInputChecker.needs_model_route_finalization(
|
||||||
|
DeepEPv2CombineInput(torch.empty(2, 8), None)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -52,6 +52,13 @@ class _FakeBuffer:
|
|||||||
num_recv = (x[0] if isinstance(x, tuple) else x).shape[0]
|
num_recv = (x[0] if isinstance(x, tuple) else x).shape[0]
|
||||||
topk_idx = kwargs["topk_idx"]
|
topk_idx = kwargs["topk_idx"]
|
||||||
topk_weights = kwargs["topk_weights"]
|
topk_weights = kwargs["topk_weights"]
|
||||||
|
if kwargs["do_expand"]:
|
||||||
|
if isinstance(x, tuple):
|
||||||
|
x = tuple(t.repeat_interleave(TOPK, dim=0) for t in x)
|
||||||
|
else:
|
||||||
|
x = x.repeat_interleave(TOPK, dim=0)
|
||||||
|
topk_idx = None
|
||||||
|
topk_weights = topk_weights.flatten()
|
||||||
event = SimpleNamespace(event=None, current_stream_wait=lambda: None)
|
event = SimpleNamespace(event=None, current_stream_wait=lambda: None)
|
||||||
return x, topk_idx, topk_weights, _FakeHandle(num_recv), event
|
return x, topk_idx, topk_weights, _FakeHandle(num_recv), event
|
||||||
|
|
||||||
@@ -93,7 +100,9 @@ class _DeepEPv2WireDtypeBase(CustomTestCase):
|
|||||||
for item in reversed(self._patches):
|
for item in reversed(self._patches):
|
||||||
item.stop()
|
item.stop()
|
||||||
|
|
||||||
def _dispatch(self, use_fp8_dispatch, num_tokens=8, is_extend_in_batch=True):
|
def _dispatch(
|
||||||
|
self, use_fp8_dispatch, num_tokens=8, is_extend_in_batch=True, group_size=128
|
||||||
|
):
|
||||||
dispatcher = deepep_v2.DeepEPv2Dispatcher(
|
dispatcher = deepep_v2.DeepEPv2Dispatcher(
|
||||||
group=_FakeGroup(),
|
group=_FakeGroup(),
|
||||||
router_topk=TOPK,
|
router_topk=TOPK,
|
||||||
@@ -102,6 +111,7 @@ class _DeepEPv2WireDtypeBase(CustomTestCase):
|
|||||||
hidden_size=HIDDEN,
|
hidden_size=HIDDEN,
|
||||||
params_dtype=torch.bfloat16,
|
params_dtype=torch.bfloat16,
|
||||||
use_fp8_dispatch=use_fp8_dispatch,
|
use_fp8_dispatch=use_fp8_dispatch,
|
||||||
|
activation_scale_block_size=group_size,
|
||||||
)
|
)
|
||||||
dispatcher._impl.num_max_dispatch_tokens_per_rank = NUM_MAX_TOKENS
|
dispatcher._impl.num_max_dispatch_tokens_per_rank = NUM_MAX_TOKENS
|
||||||
hidden_states = torch.randn((num_tokens, HIDDEN), dtype=torch.bfloat16)
|
hidden_states = torch.randn((num_tokens, HIDDEN), dtype=torch.bfloat16)
|
||||||
@@ -116,6 +126,18 @@ class _DeepEPv2WireDtypeBase(CustomTestCase):
|
|||||||
|
|
||||||
|
|
||||||
class TestDeepEPv2WireDtype(_DeepEPv2WireDtypeBase):
|
class TestDeepEPv2WireDtype(_DeepEPv2WireDtypeBase):
|
||||||
|
def test_mxfp8_scale_group_survives_both_dispatch_layouts(self):
|
||||||
|
for is_extend in (True, False):
|
||||||
|
with self.subTest(is_extend=is_extend):
|
||||||
|
_, out = self._dispatch(
|
||||||
|
use_fp8_dispatch=True,
|
||||||
|
is_extend_in_batch=is_extend,
|
||||||
|
group_size=32,
|
||||||
|
)
|
||||||
|
self.assertEqual(out.activation_scale_block_size, 32)
|
||||||
|
self.assertEqual(out.hidden_states_scale.shape[-1], HIDDEN // 32)
|
||||||
|
self.assertEqual(out.is_expanded, not is_extend)
|
||||||
|
|
||||||
def test_bf16_dispatch_sends_unquantized_activations(self):
|
def test_bf16_dispatch_sends_unquantized_activations(self):
|
||||||
hidden_states, out = self._dispatch(use_fp8_dispatch=False)
|
hidden_states, out = self._dispatch(use_fp8_dispatch=False)
|
||||||
self.assertIs(_FakeBuffer.last.dispatch_x, hidden_states)
|
self.assertIs(_FakeBuffer.last.dispatch_x, hidden_states)
|
||||||
|
|||||||
@@ -76,13 +76,18 @@ def _fp8_method(**overrides):
|
|||||||
return method
|
return method
|
||||||
|
|
||||||
|
|
||||||
def test_deepep_v2_quant_contract_accepts_blockwise_fp8(_moe_flags):
|
@pytest.mark.parametrize("use_mxfp8,block_size", [(False, [128, 128]), (True, [1, 32])])
|
||||||
|
def test_deepep_v2_quant_contract_accepts_blockwise_fp8(
|
||||||
|
_moe_flags, use_mxfp8, block_size
|
||||||
|
):
|
||||||
from sglang.srt.layers.moe.fused_moe_triton.layer import (
|
from sglang.srt.layers.moe.fused_moe_triton.layer import (
|
||||||
_validate_deepep_v2_quant_method,
|
_validate_deepep_v2_quant_method,
|
||||||
)
|
)
|
||||||
|
|
||||||
_moe_flags.a2a_backend = MoeA2ABackend.DEEPEP_V2
|
_moe_flags.a2a_backend = MoeA2ABackend.DEEPEP_V2
|
||||||
_validate_deepep_v2_quant_method(_fp8_method(weight_block_size=[128, 128]))
|
_validate_deepep_v2_quant_method(
|
||||||
|
_fp8_method(weight_block_size=block_size, use_mxfp8=use_mxfp8)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.parametrize(
|
@pytest.mark.parametrize(
|
||||||
@@ -90,7 +95,7 @@ def test_deepep_v2_quant_contract_accepts_blockwise_fp8(_moe_flags):
|
|||||||
[
|
[
|
||||||
({"activation_scheme": "static"}, "activation_scheme"),
|
({"activation_scheme": "static"}, "activation_scheme"),
|
||||||
({"weight_block_size": None}, "weight_block_size"),
|
({"weight_block_size": None}, "weight_block_size"),
|
||||||
({"weight_block_size": (1, 32), "use_mxfp8": True}, "MXFP8"),
|
({"weight_block_size": (128, 128), "use_mxfp8": True}, "MXFP8"),
|
||||||
({"is_fp4_expert": True}, "FP4 experts"),
|
({"is_fp4_expert": True}, "FP4 experts"),
|
||||||
],
|
],
|
||||||
)
|
)
|
||||||
@@ -149,5 +154,46 @@ def test_deepep_v2_runner_backstop(_moe_flags):
|
|||||||
assert MoeRunner(MoeRunnerBackend.DEEP_GEMM, MoeRunnerConfig()).runner_core
|
assert MoeRunner(MoeRunnerBackend.DEEP_GEMM, MoeRunnerConfig()).runner_core
|
||||||
|
|
||||||
|
|
||||||
|
def test_deepep_v2_registration_uses_primary_architecture_and_rejects_conflicts():
|
||||||
|
from sglang.srt.configs.moe_model_registry import (
|
||||||
|
model_requires_fp32_silu_mul,
|
||||||
|
model_supports_deepep_v2,
|
||||||
|
register_deepep_v2_model,
|
||||||
|
)
|
||||||
|
|
||||||
|
register_deepep_v2_model("TestRoutewiseMoe", silu_mul_keep_fp32=True)
|
||||||
|
config = SimpleNamespace(architectures=["TestRoutewiseMoe"])
|
||||||
|
assert model_supports_deepep_v2(config)
|
||||||
|
assert model_requires_fp32_silu_mul(config)
|
||||||
|
config.architectures = ["UnsupportedMoe", "TestRoutewiseMoe"]
|
||||||
|
assert not model_supports_deepep_v2(config)
|
||||||
|
assert not model_requires_fp32_silu_mul(config)
|
||||||
|
with pytest.raises(ValueError, match="Conflicting"):
|
||||||
|
register_deepep_v2_model("TestRoutewiseMoe", silu_mul_keep_fp32=False)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("block_size,width", [(32, 16), (128, 4)])
|
||||||
|
def test_mxfp8_recipes_keep_activation_and_weight_groups_separate(block_size, width):
|
||||||
|
# The packed format is a layout contract, independent of the host GPU.
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from sglang.srt.layers import deep_gemm_wrapper
|
||||||
|
from sglang.srt.layers.moe.moe_runner.deep_gemm import DeepGemmMoeQuantInfo
|
||||||
|
|
||||||
|
with patch.object(deep_gemm_wrapper, "DEEPGEMM_SCALE_UE8M0", True):
|
||||||
|
quant = DeepGemmMoeQuantInfo(
|
||||||
|
None, None, True, block_shape=[1, 32], use_mxfp8=True
|
||||||
|
)
|
||||||
|
assert quant.scale_recipes(
|
||||||
|
activation_block_size=block_size, hidden_size=2048, activation_scale_width=width
|
||||||
|
) == ((1, block_size), (1, 32))
|
||||||
|
with pytest.raises(AssertionError, match="activation scale mismatch"):
|
||||||
|
quant.scale_recipes(
|
||||||
|
activation_block_size=block_size,
|
||||||
|
hidden_size=2048,
|
||||||
|
activation_scale_width=width + 1,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
sys.exit(pytest.main([__file__, "-v"]))
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
|
|||||||
Reference in New Issue
Block a user