Support asymmetric compressed-tensors MoE (#27690)
This commit is contained in:
@@ -815,6 +815,7 @@ class FusedMoE(torch.nn.Module):
|
||||
"CompressedTensorsWNA16TritonMoE",
|
||||
]
|
||||
)
|
||||
and "zero" not in weight_name
|
||||
else loaded_weight
|
||||
)
|
||||
|
||||
@@ -1034,6 +1035,7 @@ class FusedMoE(torch.nn.Module):
|
||||
"CompressedTensorsWNA16TritonMoE",
|
||||
]
|
||||
)
|
||||
and "zero" not in weight_name
|
||||
else loaded_weight
|
||||
)
|
||||
|
||||
|
||||
+55
-2
@@ -22,6 +22,7 @@ from sglang.srt.layers.quantization.compressed_tensors.schemes import (
|
||||
from sglang.srt.layers.quantization.marlin_utils import (
|
||||
marlin_make_workspace,
|
||||
marlin_moe_permute_scales,
|
||||
moe_awq_to_marlin_zero_points,
|
||||
)
|
||||
from sglang.srt.layers.quantization.utils import replace_parameter
|
||||
from sglang.srt.utils import get_bool_env_var, is_cuda, is_hip, set_weight_attrs
|
||||
@@ -69,7 +70,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
||||
self.strategy = config.strategy
|
||||
self.group_size = config.group_size
|
||||
self.actorder = config.actorder
|
||||
assert config.symmetric, "Only symmetric quantization is supported for MoE"
|
||||
self.sym = config.symmetric
|
||||
|
||||
if not (
|
||||
self.quant_config.quant_format == CompressionFormat.pack_quantized.value
|
||||
@@ -129,7 +130,7 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
||||
|
||||
# In the case where we have actorder/g_idx,
|
||||
# we do not partition the w2 scales
|
||||
load_full_w2 = self.actorder and self.group_size != -1
|
||||
load_full_w2 = (self.actorder != "static") and self.group_size != -1
|
||||
|
||||
if load_full_w2:
|
||||
w2_scales_size = intermediate_size_per_partition * layer.moe_tp_size
|
||||
@@ -177,6 +178,32 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
||||
layer.register_parameter("w13_weight_shape", w13_weight_shape)
|
||||
set_weight_attrs(w13_weight_shape, extra_weight_attrs)
|
||||
|
||||
# add zero param
|
||||
if not self.sym:
|
||||
w13_qzeros = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
num_groups_w13,
|
||||
2 * intermediate_size_per_partition // self.packed_factor,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w13_weight_zero_point", w13_qzeros)
|
||||
set_weight_attrs(w13_qzeros, extra_weight_attrs)
|
||||
|
||||
w2_qzeros = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
num_groups_w2,
|
||||
hidden_size // self.packed_factor,
|
||||
dtype=torch.int32,
|
||||
),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.register_parameter("w2_weight_zero_point", w2_qzeros)
|
||||
set_weight_attrs(w2_qzeros, extra_weight_attrs)
|
||||
|
||||
w13_g_idx = torch.nn.Parameter(
|
||||
torch.empty(
|
||||
num_experts,
|
||||
@@ -236,6 +263,10 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
||||
layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape)
|
||||
layer._original_shapes["w13_weight_scale"] = tuple(w13_scale.shape)
|
||||
|
||||
if not self.sym:
|
||||
layer._original_shapes["w13_weight_zero_point"] = w13_qzeros.shape
|
||||
layer._original_shapes["w2_weight_zero_point"] = tuple(w2_qzeros.shape)
|
||||
|
||||
def process_weights_after_loading(self, layer: torch.nn.Module) -> None:
|
||||
|
||||
# Skip if the layer is already converted to Marlin format to prevent double-packing.
|
||||
@@ -339,6 +370,24 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
||||
)
|
||||
replace_tensor("w2_weight_scale", marlin_w2_scales)
|
||||
|
||||
# Repack zero
|
||||
if not self.sym:
|
||||
marlin_w13_zp = moe_awq_to_marlin_zero_points(
|
||||
layer.w13_weight_zero_point,
|
||||
size_k=layer.w13_weight_zero_point.shape[1],
|
||||
size_n=layer.w13_weight_zero_point.shape[2] * self.packed_factor,
|
||||
num_bits=self.num_bits,
|
||||
)
|
||||
replace_tensor("w13_weight_zero_point", marlin_w13_zp)
|
||||
|
||||
marlin_w2_zp = moe_awq_to_marlin_zero_points(
|
||||
layer.w2_weight_zero_point,
|
||||
size_k=layer.w2_weight_zero_point.shape[1],
|
||||
size_n=layer.w2_weight_zero_point.shape[2] * self.packed_factor,
|
||||
num_bits=self.num_bits,
|
||||
)
|
||||
replace_tensor("w2_weight_zero_point", marlin_w2_zp)
|
||||
|
||||
layer.workspace = marlin_make_workspace(layer.w13_weight_packed.device, 4)
|
||||
layer.is_marlin_converted = True
|
||||
|
||||
@@ -376,6 +425,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
||||
w13_g_idx=getattr(layer, "w13_weight_g_idx", None),
|
||||
w2_g_idx=getattr(layer, "w2_weight_g_idx", None),
|
||||
is_k_full=self.is_k_full,
|
||||
w13_qzeros=layer.w13_weight_zero_point if not self.sym else None,
|
||||
w2_qzeros=layer.w2_weight_zero_point if not self.sym else None,
|
||||
)
|
||||
|
||||
def apply_weights(
|
||||
@@ -422,6 +473,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
||||
g_idx2=layer.w2_weight_g_idx,
|
||||
sort_indices1=layer.w13_g_idx_sort_indices,
|
||||
sort_indices2=layer.w2_g_idx_sort_indices,
|
||||
w1_zeros=layer.w13_weight_zero_point if not self.sym else None,
|
||||
w2_zeros=layer.w2_weight_zero_point if not self.sym else None,
|
||||
num_bits=self.num_bits,
|
||||
is_k_full=self.is_k_full,
|
||||
routed_scaling_factor=self.moe_runner_config.routed_scaling_factor,
|
||||
|
||||
Reference in New Issue
Block a user