[AMD] Add MoE weights and scales padding (#18684)

This commit is contained in:
mqhc2020
2026-03-20 14:55:09 -07:00
committed by GitHub
parent f97c09dac1
commit 9419453713
8 changed files with 131 additions and 36 deletions
@@ -6,7 +6,6 @@
from __future__ import annotations from __future__ import annotations
import functools import functools
import os
from typing import TYPE_CHECKING, List, Optional from typing import TYPE_CHECKING, List, Optional
import torch import torch
@@ -14,6 +13,7 @@ import torch.nn.functional as F
import triton.language as tl import triton.language as tl
from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig from sglang.srt.layers.moe.moe_runner import MoeRunnerConfig
from sglang.srt.layers.moe.utils import get_moe_padding_size
from sglang.srt.utils import ( from sglang.srt.utils import (
cpu_has_amx_support, cpu_has_amx_support,
get_bool_env_var, get_bool_env_var,
@@ -75,7 +75,7 @@ if not _is_cuda and not _is_hip and not _is_xpu:
# Fallback: vllm not available, will use native PyTorch implementations # Fallback: vllm not available, will use native PyTorch implementations
_has_vllm_ops = False _has_vllm_ops = False
padding_size = 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0 padding_size = get_moe_padding_size(_use_aiter)
@register_custom_op(mutates_args=["hidden_states"]) @register_custom_op(mutates_args=["hidden_states"])
@@ -1,7 +1,6 @@
from __future__ import annotations from __future__ import annotations
import functools import functools
import os
from collections import OrderedDict from collections import OrderedDict
from typing import Any, Dict, List, Optional from typing import Any, Dict, List, Optional
@@ -10,6 +9,7 @@ import triton
import triton.language as tl import triton.language as tl
from sglang.srt.batch_invariant_ops import is_batch_invariant_mode_enabled from sglang.srt.batch_invariant_ops import is_batch_invariant_mode_enabled
from sglang.srt.layers.moe.utils import get_moe_padding_size
from sglang.srt.layers.quantization.fp8_kernel import ( from sglang.srt.layers.quantization.fp8_kernel import (
per_token_group_quant_fp8, per_token_group_quant_fp8,
scaled_fp8_quant, scaled_fp8_quant,
@@ -49,7 +49,7 @@ elif _is_cpu and _is_cpu_amx_available:
elif _is_hip: elif _is_hip:
pass pass
padding_size = 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0 padding_size = get_moe_padding_size(_use_aiter)
def support_tensor_descriptor(): def support_tensor_descriptor():
@@ -440,8 +440,14 @@ class FusedMoE(torch.nn.Module):
# Use narrow_padded_param_and_loaded_weight for: # Use narrow_padded_param_and_loaded_weight for:
# 1. CPU (always) # 1. CPU (always)
# 2. GPU with flashinfer_trtllm padding (when intermediate_size is padded to 128) # 2. GPU with flashinfer_trtllm padding (when intermediate_size is padded to 128)
# 3. GPU with Aiter padding
# This handles the case where the loaded weights are smaller than the padded expert_data # This handles the case where the loaded weights are smaller than the padded expert_data
use_padded_loading = _is_cpu or self.use_flashinfer_trtllm_moe aiter_padded = (
_use_aiter
and hasattr(self, "w2_weight")
and getattr(self.w2_weight, "weight_padded", False)
)
use_padded_loading = _is_cpu or self.use_flashinfer_trtllm_moe or aiter_padded
if use_padded_loading: if use_padded_loading:
expert_data, loaded_weight = narrow_padded_param_and_loaded_weight( expert_data, loaded_weight = narrow_padded_param_and_loaded_weight(
expert_data, expert_data,
@@ -514,8 +520,14 @@ class FusedMoE(torch.nn.Module):
# Use narrow_padded_param_and_loaded_weight for: # Use narrow_padded_param_and_loaded_weight for:
# 1. CPU (always) # 1. CPU (always)
# 2. GPU with flashinfer_trtllm padding (when intermediate_size is padded to 128) # 2. GPU with flashinfer_trtllm padding (when intermediate_size is padded to 128)
# 3. GPU with Aiter padding
# This handles the case where the loaded weights are smaller than the padded expert_data # This handles the case where the loaded weights are smaller than the padded expert_data
use_padded_loading = _is_cpu or self.use_flashinfer_trtllm_moe aiter_padded = (
_use_aiter
and hasattr(self, "w2_weight")
and getattr(self.w2_weight, "weight_padded", False)
)
use_padded_loading = _is_cpu or self.use_flashinfer_trtllm_moe or aiter_padded
if use_padded_loading: if use_padded_loading:
expert_data, loaded_weight = narrow_padded_param_and_loaded_weight( expert_data, loaded_weight = narrow_padded_param_and_loaded_weight(
expert_data, expert_data,
+30
View File
@@ -1,10 +1,13 @@
from __future__ import annotations from __future__ import annotations
import logging import logging
import os
from contextlib import contextmanager from contextlib import contextmanager
from enum import Enum, IntEnum from enum import Enum, IntEnum
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING, Optional
import torch
from sglang.srt.distributed.parallel_state import get_moe_expert_parallel_world_size from sglang.srt.distributed.parallel_state import get_moe_expert_parallel_world_size
from sglang.srt.layers.dp_attention import ( from sglang.srt.layers.dp_attention import (
get_attention_dp_size, get_attention_dp_size,
@@ -341,3 +344,30 @@ class RoutingMethodType(IntEnum):
TopK = (5,) TopK = (5,)
# Unspecified # Unspecified
Unspecified = 6 Unspecified = 6
def get_moe_padding_size(is_aiter_moe):
if is_aiter_moe:
return 128
else:
return 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0
def get_moe_weight_sizes(inter_dim, is_concat, is_packed, is_aiter_moe):
w13_up_dim = 2 * inter_dim if is_concat else inter_dim
w2_down_dim = inter_dim // 2 if is_packed else inter_dim
if is_aiter_moe:
padding_size = get_moe_padding_size(True)
align_aiter = lambda n: ((n + padding_size - 1) // padding_size) * padding_size
is_padded = (w2_down_dim % padding_size) > 0
if is_padded:
w2_down_dim = align_aiter(w2_down_dim)
# up proj + gate fusion : 2x
if is_concat:
w13_up_dim = w2_down_dim * 2
# packed
if hasattr(torch, "float4_e2m1fn_x2") and is_packed:
w13_up_dim *= 2
return (w13_up_dim, w2_down_dim, False if not is_aiter_moe else is_padded)
@@ -12,7 +12,10 @@ from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
FlashInferTrtllmFp8MoeQuantInfo, FlashInferTrtllmFp8MoeQuantInfo,
) )
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
from sglang.srt.layers.moe.utils import get_moe_runner_backend from sglang.srt.layers.moe.utils import (
get_moe_runner_backend,
get_moe_weight_sizes,
)
from sglang.srt.layers.quantization.compressed_tensors.schemes import ( from sglang.srt.layers.quantization.compressed_tensors.schemes import (
CompressedTensorsMoEScheme, CompressedTensorsMoEScheme,
) )
@@ -120,11 +123,22 @@ class CompressedTensorsW8A8Fp8MoE(CompressedTensorsMoEScheme):
f"weight quantization block_k = {block_k}." f"weight quantization block_k = {block_k}."
) )
w13_up_dim, w2_down_dim, weight_padded = get_moe_weight_sizes(
intermediate_size_per_partition,
is_aiter_moe=True,
is_concat=True,
is_packed=False,
)
extra_weight_attrs.update(
{"weight_padded": weight_padded},
)
# WEIGHTS # WEIGHTS
w13_weight = torch.nn.Parameter( w13_weight = torch.nn.Parameter(
torch.empty( torch.empty(
num_experts, num_experts,
2 * intermediate_size_per_partition, w13_up_dim,
hidden_size, hidden_size,
dtype=params_dtype, dtype=params_dtype,
), ),
@@ -137,7 +151,7 @@ class CompressedTensorsW8A8Fp8MoE(CompressedTensorsMoEScheme):
torch.empty( torch.empty(
num_experts, num_experts,
hidden_size, hidden_size,
intermediate_size_per_partition, w2_down_dim,
dtype=params_dtype, dtype=params_dtype,
), ),
requires_grad=False, requires_grad=False,
@@ -161,7 +175,7 @@ class CompressedTensorsW8A8Fp8MoE(CompressedTensorsMoEScheme):
w13_weight_scale = torch.nn.Parameter( w13_weight_scale = torch.nn.Parameter(
torch.ones( torch.ones(
num_experts, num_experts,
2 * intermediate_size_per_partition, w13_up_dim,
1, 1,
dtype=torch.float32, dtype=torch.float32,
), ),
+37 -20
View File
@@ -26,7 +26,12 @@ from sglang.srt.layers.moe.moe_runner.flashinfer_trtllm import (
FlashInferTrtllmFp8MoeQuantInfo, FlashInferTrtllmFp8MoeQuantInfo,
) )
from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo from sglang.srt.layers.moe.moe_runner.triton import TritonMoeQuantInfo
from sglang.srt.layers.moe.utils import RoutingMethodType, get_moe_runner_backend from sglang.srt.layers.moe.utils import (
RoutingMethodType,
get_moe_padding_size,
get_moe_runner_backend,
get_moe_weight_sizes,
)
from sglang.srt.layers.parameter import ( from sglang.srt.layers.parameter import (
BlockQuantScaleParameter, BlockQuantScaleParameter,
ModelWeightParameter, ModelWeightParameter,
@@ -778,27 +783,38 @@ class Fp8MoEMethod(FusedMoEMethodBase):
if self.quant_config.is_checkpoint_fp8_serialized: if self.quant_config.is_checkpoint_fp8_serialized:
params_dtype = torch.uint32 if _use_hip_int4 else torch.float8_e4m3fn params_dtype = torch.uint32 if _use_hip_int4 else torch.float8_e4m3fn
tp_size = get_tensor_model_parallel_world_size() tp_size = get_tensor_model_parallel_world_size()
w13_up_dim, w2_up_dim, weight_padded = get_moe_weight_sizes(
intermediate_size_per_partition,
is_aiter_moe=True,
is_concat=True,
is_packed=False,
)
if self.block_quant: if self.block_quant:
block_n, block_k = ( block_n, block_k = (
self.quant_config.weight_block_size[0], self.quant_config.weight_block_size[0],
self.quant_config.weight_block_size[1], self.quant_config.weight_block_size[1],
) )
# NOTE(HandH1998): To ensure proper alignment of the block-wise quantization scales, the output_size of the weights for both the gate and up layers must be divisible by block_n.
# Required by column parallel or enabling merged weights padding_size = get_moe_padding_size(_use_aiter)
if intermediate_size_per_partition % block_n != 0: if not (_use_aiter and padding_size == block_n == block_k):
raise ValueError( # NOTE(HandH1998): To ensure proper alignment of the block-wise quantization scales, the output_size of the weights for both the gate and up layers must be divisible by block_n.
f"The output_size of gate's and up's weight = " # Required by column parallel or enabling merged weights
f"{intermediate_size_per_partition} is not divisible by " if intermediate_size_per_partition % block_n != 0:
f"weight quantization block_n = {block_n}."
)
if tp_size > 1:
# Required by row parallel
if intermediate_size_per_partition % block_k != 0:
raise ValueError( raise ValueError(
f"The input_size of down's weight = " f"The output_size of gate's and up's weight = "
f"{intermediate_size_per_partition} is not divisible by " f"{intermediate_size_per_partition} is not divisible by "
f"weight quantization block_k = {block_k}." f"weight quantization block_n = {block_n}."
) )
if tp_size > 1:
# Required by row parallel
if intermediate_size_per_partition % block_k != 0:
raise ValueError(
f"The input_size of down's weight = "
f"{intermediate_size_per_partition} is not divisible by "
f"weight quantization block_k = {block_k}."
)
# WEIGHTS # WEIGHTS
if _is_hip and _use_hip_int4: if _is_hip and _use_hip_int4:
@@ -825,7 +841,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
w13_weight = torch.nn.Parameter( w13_weight = torch.nn.Parameter(
torch.empty( torch.empty(
num_experts, num_experts,
2 * intermediate_size_per_partition, w13_up_dim,
hidden_size, hidden_size,
dtype=params_dtype, dtype=params_dtype,
), ),
@@ -835,12 +851,16 @@ class Fp8MoEMethod(FusedMoEMethodBase):
torch.empty( torch.empty(
num_experts, num_experts,
hidden_size, hidden_size,
intermediate_size_per_partition, w2_up_dim,
dtype=params_dtype, dtype=params_dtype,
), ),
requires_grad=False, requires_grad=False,
) )
extra_weight_attrs.update(
{"weight_padded": weight_padded},
)
layer.register_parameter("w13_weight", w13_weight) layer.register_parameter("w13_weight", w13_weight)
set_weight_attrs(w13_weight, extra_weight_attrs) set_weight_attrs(w13_weight, extra_weight_attrs)
@@ -1401,10 +1421,7 @@ class Fp8MoEMethod(FusedMoEMethodBase):
layer.w2_weight_scale1[expert_id] *= layer.w2_weight_scale[expert_id] layer.w2_weight_scale1[expert_id] *= layer.w2_weight_scale[expert_id]
def process_weights_hip_scale_padding(self, layer: Module): def process_weights_hip_scale_padding(self, layer: Module):
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import ( padding_size = get_moe_padding_size(_use_aiter)
padding_size, # Avoid circular import
)
if _use_aiter: if _use_aiter:
layer.w13_weight = torch.nn.Parameter( layer.w13_weight = torch.nn.Parameter(
shuffle_weight(layer.w13_weight.data, (16, 16)), shuffle_weight(layer.w13_weight.data, (16, 16)),
@@ -8,6 +8,7 @@ from typing import TYPE_CHECKING, Any
import torch import torch
from sglang.srt.layers.moe import MoeRunnerConfig from sglang.srt.layers.moe import MoeRunnerConfig
from sglang.srt.layers.moe.utils import get_moe_weight_sizes
from sglang.srt.layers.quantization.quark.schemes import QuarkMoEScheme from sglang.srt.layers.quantization.quark.schemes import QuarkMoEScheme
from sglang.srt.utils import ( from sglang.srt.utils import (
get_bool_env_var, get_bool_env_var,
@@ -73,10 +74,20 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported from sglang.srt.layers.moe.fused_moe_triton import FusedMoeWeightScaleSupported
w13_up_dim, w2_down_dim, weight_padded = get_moe_weight_sizes(
intermediate_size_per_partition,
is_aiter_moe=True,
is_concat=True,
is_packed=True,
)
# Add the quantization method used (per tensor/grouped/channel) # Add the quantization method used (per tensor/grouped/channel)
# to ensure the weight scales are loaded in properly # to ensure the weight scales are loaded in properly
extra_weight_attrs.update( extra_weight_attrs.update(
{"quant_method": FusedMoeWeightScaleSupported.BLOCK.value} {
"quant_method": FusedMoeWeightScaleSupported.BLOCK.value,
"weight_padded": weight_padded,
},
) )
params_dtype = torch.uint8 params_dtype = torch.uint8
@@ -85,7 +96,7 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
w13_weight = torch.nn.Parameter( w13_weight = torch.nn.Parameter(
torch.empty( torch.empty(
num_experts, num_experts,
2 * intermediate_size_per_partition, w13_up_dim,
hidden_size // 2, hidden_size // 2,
dtype=params_dtype, dtype=params_dtype,
), ),
@@ -99,7 +110,7 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
torch.empty( torch.empty(
num_experts, num_experts,
hidden_size, hidden_size,
intermediate_size_per_partition // 2, w2_down_dim,
dtype=params_dtype, dtype=params_dtype,
), ),
requires_grad=False, requires_grad=False,
@@ -112,17 +123,24 @@ class QuarkW4A4MXFp4MoE(QuarkMoEScheme):
w13_weight_scale = torch.nn.Parameter( w13_weight_scale = torch.nn.Parameter(
torch.ones( torch.ones(
num_experts, num_experts,
2 * intermediate_size_per_partition, w13_up_dim,
hidden_size // OCP_MX_BLOCK_SIZE, hidden_size // OCP_MX_BLOCK_SIZE,
dtype=params_dtype, dtype=params_dtype,
), ),
requires_grad=False, requires_grad=False,
) )
W2_SCALE_DIVIDEND = w2_down_dim * 2
W2_SCALE_DIVISOR = intermediate_size_per_partition
scaling_up = lambda dividend, divisor: (dividend * W2_SCALE_DIVIDEND) // (
divisor * W2_SCALE_DIVISOR
)
w2_weight_scale = torch.nn.Parameter( w2_weight_scale = torch.nn.Parameter(
torch.ones( torch.ones(
num_experts, num_experts,
hidden_size, hidden_size,
intermediate_size_per_partition // OCP_MX_BLOCK_SIZE, scaling_up(intermediate_size_per_partition, OCP_MX_BLOCK_SIZE),
dtype=params_dtype, dtype=params_dtype,
), ),
requires_grad=False, requires_grad=False,
@@ -162,6 +162,7 @@ from sglang.srt.utils import (
empty_context, empty_context,
enable_show_time_cost, enable_show_time_cost,
get_available_gpu_memory, get_available_gpu_memory,
get_bool_env_var,
get_cpu_ids_by_node, get_cpu_ids_by_node,
init_custom_process_group, init_custom_process_group,
is_hip, is_hip,
@@ -198,6 +199,7 @@ _is_hip = is_hip()
_is_npu = is_npu() _is_npu = is_npu()
_is_cpu_amx_available = cpu_has_amx_support() _is_cpu_amx_available = cpu_has_amx_support()
_is_cpu_arm64 = is_host_cpu_arm64() _is_cpu_arm64 = is_host_cpu_arm64()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
if _is_npu: if _is_npu:
from sglang.srt.hardware_backend.npu.utils import init_npu_backend from sglang.srt.hardware_backend.npu.utils import init_npu_backend
@@ -799,7 +801,9 @@ class ModelRunner(ModelRunnerKVCacheMixin):
f"moe_intermediate_size {moe_intermediate_size} must be divisible by moe_tp_size ({moe_tp_size}) which is tp_size ({self.tp_size}) divided by moe_ep_size ({self.moe_ep_size})." f"moe_intermediate_size {moe_intermediate_size} must be divisible by moe_tp_size ({moe_tp_size}) which is tp_size ({self.tp_size}) divided by moe_ep_size ({self.moe_ep_size})."
) )
if (moe_intermediate_size // moe_tp_size) % weight_block_size_n != 0: if (
moe_intermediate_size // moe_tp_size
) % weight_block_size_n != 0 and not _use_aiter:
raise ValueError( raise ValueError(
f"For quantized MoE models, please make sure ({moe_intermediate_size=} / {moe_tp_size=}) % {weight_block_size_n=} == 0 " f"For quantized MoE models, please make sure ({moe_intermediate_size=} / {moe_tp_size=}) % {weight_block_size_n=} == 0 "
f"where moe_tp_size is equal to tp_size ({self.tp_size}) divided by ep_size ({self.moe_ep_size}). " f"where moe_tp_size is equal to tp_size ({self.tp_size}) divided by ep_size ({self.moe_ep_size}). "