Support asymmetric compressed-tensors MoE (#27690)
This commit is contained in:
@@ -815,6 +815,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
"CompressedTensorsWNA16TritonMoE",
|
"CompressedTensorsWNA16TritonMoE",
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
and "zero" not in weight_name
|
||||||
else loaded_weight
|
else loaded_weight
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -1034,6 +1035,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
"CompressedTensorsWNA16TritonMoE",
|
"CompressedTensorsWNA16TritonMoE",
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
and "zero" not in weight_name
|
||||||
else loaded_weight
|
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 (
|
from sglang.srt.layers.quantization.marlin_utils import (
|
||||||
marlin_make_workspace,
|
marlin_make_workspace,
|
||||||
marlin_moe_permute_scales,
|
marlin_moe_permute_scales,
|
||||||
|
moe_awq_to_marlin_zero_points,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.quantization.utils import replace_parameter
|
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
|
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.strategy = config.strategy
|
||||||
self.group_size = config.group_size
|
self.group_size = config.group_size
|
||||||
self.actorder = config.actorder
|
self.actorder = config.actorder
|
||||||
assert config.symmetric, "Only symmetric quantization is supported for MoE"
|
self.sym = config.symmetric
|
||||||
|
|
||||||
if not (
|
if not (
|
||||||
self.quant_config.quant_format == CompressionFormat.pack_quantized.value
|
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,
|
# In the case where we have actorder/g_idx,
|
||||||
# we do not partition the w2 scales
|
# 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:
|
if load_full_w2:
|
||||||
w2_scales_size = intermediate_size_per_partition * layer.moe_tp_size
|
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)
|
layer.register_parameter("w13_weight_shape", w13_weight_shape)
|
||||||
set_weight_attrs(w13_weight_shape, extra_weight_attrs)
|
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(
|
w13_g_idx = torch.nn.Parameter(
|
||||||
torch.empty(
|
torch.empty(
|
||||||
num_experts,
|
num_experts,
|
||||||
@@ -236,6 +263,10 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
|||||||
layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape)
|
layer._original_shapes["w2_weight_scale"] = tuple(w2_scale.shape)
|
||||||
layer._original_shapes["w13_weight_scale"] = tuple(w13_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:
|
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.
|
# 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)
|
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.workspace = marlin_make_workspace(layer.w13_weight_packed.device, 4)
|
||||||
layer.is_marlin_converted = True
|
layer.is_marlin_converted = True
|
||||||
|
|
||||||
@@ -376,6 +425,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
|||||||
w13_g_idx=getattr(layer, "w13_weight_g_idx", None),
|
w13_g_idx=getattr(layer, "w13_weight_g_idx", None),
|
||||||
w2_g_idx=getattr(layer, "w2_weight_g_idx", None),
|
w2_g_idx=getattr(layer, "w2_weight_g_idx", None),
|
||||||
is_k_full=self.is_k_full,
|
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(
|
def apply_weights(
|
||||||
@@ -422,6 +473,8 @@ class CompressedTensorsWNA16MoE(CompressedTensorsMoEScheme):
|
|||||||
g_idx2=layer.w2_weight_g_idx,
|
g_idx2=layer.w2_weight_g_idx,
|
||||||
sort_indices1=layer.w13_g_idx_sort_indices,
|
sort_indices1=layer.w13_g_idx_sort_indices,
|
||||||
sort_indices2=layer.w2_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,
|
num_bits=self.num_bits,
|
||||||
is_k_full=self.is_k_full,
|
is_k_full=self.is_k_full,
|
||||||
routed_scaling_factor=self.moe_runner_config.routed_scaling_factor,
|
routed_scaling_factor=self.moe_runner_config.routed_scaling_factor,
|
||||||
|
|||||||
Reference in New Issue
Block a user