ROCm: update AITER (#5816)
This commit is contained in:
@@ -45,7 +45,7 @@ if _is_cuda or _is_hip:
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
padding_size = 128 if bool(int(os.getenv("MOE_PADDING", "0"))) else 0
|
||||
padding_size = 128 if bool(int(os.getenv("SGLANG_MOE_PADDING", "0"))) else 0
|
||||
enable_moe_align_block_size_triton = bool(
|
||||
int(os.getenv("ENABLE_MOE_ALIGN_BLOCK_SIZE_TRITON", "0"))
|
||||
)
|
||||
@@ -1327,7 +1327,7 @@ def fused_experts_impl(
|
||||
if (
|
||||
not (use_fp8_w8a8 or use_int8_w8a8)
|
||||
or block_shape is not None
|
||||
or (_is_hip and get_bool_env_var("CK_MOE"))
|
||||
or (_is_hip and get_bool_env_var("SGLANG_AITER_MOE"))
|
||||
):
|
||||
padded_size = 0
|
||||
|
||||
|
||||
@@ -18,7 +18,7 @@ from sglang.srt.layers.quantization.base_config import (
|
||||
QuantizationConfig,
|
||||
QuantizeMethodBase,
|
||||
)
|
||||
from sglang.srt.utils import get_bool_env_var, is_hip, permute_weight, set_weight_attrs
|
||||
from sglang.srt.utils import get_bool_env_var, is_hip, set_weight_attrs
|
||||
|
||||
if torch.cuda.is_available():
|
||||
from sglang.srt.layers.moe.fused_moe_triton.fused_moe import fused_experts
|
||||
@@ -30,7 +30,9 @@ import logging
|
||||
_is_hip = is_hip()
|
||||
|
||||
if _is_hip:
|
||||
from aiter import ck_moe
|
||||
from aiter import ActivationType
|
||||
from aiter.fused_moe_bf16_asm import ck_moe_2stages
|
||||
from aiter.ops.shuffle import shuffle_weight
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -102,14 +104,14 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, CustomOp):
|
||||
set_weight_attrs(w2_weight, extra_weight_attrs)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
if _is_hip and get_bool_env_var("CK_MOE"):
|
||||
if _is_hip and get_bool_env_var("SGLANG_AITER_MOE"):
|
||||
layer.w13_weight = torch.nn.Parameter(
|
||||
permute_weight(layer.w13_weight.data),
|
||||
shuffle_weight(layer.w13_weight.data, (16, 16)),
|
||||
requires_grad=False,
|
||||
)
|
||||
torch.cuda.empty_cache()
|
||||
layer.w2_weight = torch.nn.Parameter(
|
||||
permute_weight(layer.w2_weight.data),
|
||||
shuffle_weight(layer.w2_weight.data, (16, 16)),
|
||||
requires_grad=False,
|
||||
)
|
||||
torch.cuda.empty_cache()
|
||||
@@ -182,21 +184,17 @@ class UnquantizedFusedMoEMethod(FusedMoEMethodBase, CustomOp):
|
||||
routed_scaling_factor=routed_scaling_factor,
|
||||
)
|
||||
|
||||
if _is_hip and get_bool_env_var("CK_MOE"):
|
||||
if _is_hip and get_bool_env_var("SGLANG_AITER_MOE"):
|
||||
assert not no_combine, "unsupported"
|
||||
return ck_moe(
|
||||
return ck_moe_2stages(
|
||||
x,
|
||||
layer.w13_weight,
|
||||
layer.w2_weight,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
32,
|
||||
None,
|
||||
activation,
|
||||
activation=(
|
||||
ActivationType.Silu if activation == "silu" else ActivationType.Gelu
|
||||
),
|
||||
)
|
||||
else:
|
||||
return fused_experts(
|
||||
@@ -527,7 +525,7 @@ class FusedMoE(torch.nn.Module):
|
||||
# Case input scale: input_scale loading is only supported for fp8
|
||||
if "input_scale" in weight_name:
|
||||
# INT4-FP8 (INT4 MoE Weight, FP8 Compute): Adjust input_scale for e4m3fnuz (AMD)
|
||||
if _is_hip and get_bool_env_var("USE_INT4_WEIGHT"):
|
||||
if _is_hip and get_bool_env_var("SGLANG_INT4_WEIGHT"):
|
||||
loaded_weight = loaded_weight * 2.0
|
||||
|
||||
# this is needed for compressed-tensors only
|
||||
@@ -569,7 +567,7 @@ class FusedMoE(torch.nn.Module):
|
||||
quant_method = getattr(param, "quant_method", None)
|
||||
if quant_method == FusedMoeWeightScaleSupported.CHANNEL.value:
|
||||
# INT4-FP8 (INT4 MoE Weight, FP8 Compute): Adjust INT4 column-wise scaling number to e4m3fnuz (AMD)
|
||||
if _is_hip and get_bool_env_var("USE_INT4_WEIGHT"):
|
||||
if _is_hip and get_bool_env_var("SGLANG_INT4_WEIGHT"):
|
||||
loaded_weight = loaded_weight * 0.5
|
||||
|
||||
self._load_per_channel_weight_scale(
|
||||
@@ -592,7 +590,7 @@ class FusedMoE(torch.nn.Module):
|
||||
)
|
||||
elif quant_method == FusedMoeWeightScaleSupported.TENSOR.value:
|
||||
# INT4-FP8 (INT4 MoE Weight, FP8 Compute): Adjust FP8 per-tensor scaling number for e4m3fnuz (AMD)
|
||||
if _is_hip and get_bool_env_var("USE_INT4_WEIGHT"):
|
||||
if _is_hip and get_bool_env_var("SGLANG_INT4_WEIGHT"):
|
||||
loaded_weight = loaded_weight * 2.0
|
||||
|
||||
self._load_per_tensor_weight_scale(
|
||||
|
||||
Reference in New Issue
Block a user