[CPU] Fix mxfp4 padding size (#31334)
This commit is contained in:
@@ -252,6 +252,11 @@ def adjust_config_with_unaligned_cpu_tp(
|
|||||||
)
|
)
|
||||||
|
|
||||||
intermediate_padding_size = tp_size * get_moe_padding_size(weight_block_size)
|
intermediate_padding_size = tp_size * get_moe_padding_size(weight_block_size)
|
||||||
|
if model_config.quantization == "mxfp4":
|
||||||
|
# For mxfp4 quantization, 2 mxfp4 values are packed to 1 uint8,
|
||||||
|
# so we need to double the intermediate padding size to ensure
|
||||||
|
# the padded intermediate size is divisible by 2 for proper packing.
|
||||||
|
intermediate_padding_size *= 2
|
||||||
for moe_intermediate_attr in [
|
for moe_intermediate_attr in [
|
||||||
"moe_intermediate_size",
|
"moe_intermediate_size",
|
||||||
"intermediate_size",
|
"intermediate_size",
|
||||||
|
|||||||
@@ -857,29 +857,6 @@ class Mxfp4MoEMethod(FusedMoEMethodBase):
|
|||||||
layer.w2_weight_bias = Parameter(
|
layer.w2_weight_bias = Parameter(
|
||||||
layer.w2_weight_bias.float(), requires_grad=False
|
layer.w2_weight_bias.float(), requires_grad=False
|
||||||
)
|
)
|
||||||
return
|
|
||||||
# Fallback if the TP-sharded layer cannot be AMX-packed
|
|
||||||
from sglang.srt.layers.quantization.mxfp4_tensor import MXFP4QuantizeUtil
|
|
||||||
|
|
||||||
w13_weight = MXFP4QuantizeUtil.dequantize(
|
|
||||||
quantized_data=layer.w13_weight,
|
|
||||||
dtype=torch.bfloat16,
|
|
||||||
scale=layer.w13_weight_scale,
|
|
||||||
block_sizes=[32],
|
|
||||||
)
|
|
||||||
w2_weight = MXFP4QuantizeUtil.dequantize(
|
|
||||||
quantized_data=layer.w2_weight,
|
|
||||||
dtype=torch.bfloat16,
|
|
||||||
scale=layer.w2_weight_scale,
|
|
||||||
block_sizes=[32],
|
|
||||||
)
|
|
||||||
del layer.w13_weight
|
|
||||||
del layer.w2_weight
|
|
||||||
del layer.w13_weight_scale
|
|
||||||
del layer.w2_weight_scale
|
|
||||||
layer.w13_weight = Parameter(w13_weight, requires_grad=False)
|
|
||||||
layer.w2_weight = Parameter(w2_weight, requires_grad=False)
|
|
||||||
|
|
||||||
return
|
return
|
||||||
else:
|
else:
|
||||||
from triton_kernels.numerics_details.mxfp import upcast_from_mxfp
|
from triton_kernels.numerics_details.mxfp import upcast_from_mxfp
|
||||||
|
|||||||
Reference in New Issue
Block a user