mirror of
https://github.com/vllm-project/vllm.git
synced 2026-08-16 02:38:12 +00:00
Refactor AWQ Marlin MoE onto modular WNA16 oracle (#42483)
Signed-off-by: Siddharth Bedekar <[email protected]> Signed-off-by: Siddharth Bedekar <[email protected]> Co-authored-by: Robert Shaw <[email protected]> Co-authored-by: OpenAI Codex <[email protected]> Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com> Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
This commit is contained in:
co-authored by
Robert Shaw
OpenAI Codex
gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
parent
f5d3dc7115
commit
78e7a7b9b0
@@ -1338,36 +1338,96 @@ def test_cpu_fused_moe_basic(
|
||||
torch.testing.assert_close(out, ref, atol=atol, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("m", [16, 32, 64])
|
||||
@pytest.mark.parametrize("n", [128])
|
||||
@pytest.mark.parametrize("k", [128])
|
||||
@pytest.mark.parametrize("e", [8, 12, 16, 32])
|
||||
@pytest.mark.parametrize("topk", [2, 4])
|
||||
@pytest.mark.parametrize("max_tokens_per_batch", [16, 32, 64])
|
||||
def _batched_fused_marlin_moe_cases() -> list[Any]:
|
||||
cases = [
|
||||
pytest.param(
|
||||
m,
|
||||
128,
|
||||
128,
|
||||
e,
|
||||
topk,
|
||||
max_tokens_per_batch,
|
||||
torch.bfloat16,
|
||||
scalar_types.float4_e2m1f,
|
||||
None,
|
||||
1e-3,
|
||||
id=(
|
||||
f"m{m}-n128-k128-e{e}-topk{topk}-max_tokens{max_tokens_per_batch}-mxfp4"
|
||||
),
|
||||
)
|
||||
for m in [16, 32, 64]
|
||||
for e in [8, 12, 16, 32]
|
||||
for topk in [2, 4]
|
||||
for max_tokens_per_batch in [16, 32, 64]
|
||||
]
|
||||
cases.append(
|
||||
pytest.param(
|
||||
32,
|
||||
128,
|
||||
128,
|
||||
8,
|
||||
2,
|
||||
64,
|
||||
torch.float16,
|
||||
scalar_types.uint4,
|
||||
scalar_types.int8,
|
||||
4e-2,
|
||||
id="awq-int8-activation-metadata",
|
||||
)
|
||||
)
|
||||
return cases
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("m,n,k,e,topk,max_tokens_per_batch,dtype,quant_dtype,input_type,atol"),
|
||||
_batched_fused_marlin_moe_cases(),
|
||||
)
|
||||
@pytest.mark.skipif(current_platform.is_rocm(), reason="Skip for rocm")
|
||||
def test_batched_fused_marlin_moe(
|
||||
m: int, n: int, k: int, e: int, topk: int, max_tokens_per_batch: int
|
||||
m: int,
|
||||
n: int,
|
||||
k: int,
|
||||
e: int,
|
||||
topk: int,
|
||||
max_tokens_per_batch: int,
|
||||
dtype: torch.dtype,
|
||||
quant_dtype: ScalarType,
|
||||
input_type: ScalarType | None,
|
||||
atol: float,
|
||||
):
|
||||
print(
|
||||
f"testing m={m}, n={n}, k={k}, e={e}, "
|
||||
f"topk={topk}, "
|
||||
f"max_tokens_per_batch={max_tokens_per_batch}"
|
||||
f"max_tokens_per_batch={max_tokens_per_batch}, "
|
||||
f"dtype={dtype}, quant_dtype={quant_dtype}, input_type={input_type}"
|
||||
)
|
||||
set_random_seed(0)
|
||||
|
||||
dtype = torch.bfloat16
|
||||
quant_dtype = scalar_types.float4_e2m1f
|
||||
group_size = 32
|
||||
if input_type == scalar_types.int8:
|
||||
input_dtype = torch.int8
|
||||
elif input_type == scalar_types.float8_e4m3fn:
|
||||
input_dtype = torch.float8_e4m3fn
|
||||
else:
|
||||
input_dtype = None
|
||||
|
||||
a = torch.randn((m, k), device="cuda", dtype=dtype) / 10
|
||||
w1 = torch.randn((e, 2 * n, k), device="cuda", dtype=dtype) / 20
|
||||
w2 = torch.randn((e, k, n), device="cuda", dtype=dtype) / 20
|
||||
|
||||
w1_data = MarlinMoEWeightData.make(
|
||||
w=w1, quant_type=quant_dtype, group_size=group_size, act_order=None
|
||||
w=w1,
|
||||
quant_type=quant_dtype,
|
||||
group_size=group_size,
|
||||
act_order=None,
|
||||
input_type=input_type,
|
||||
)
|
||||
w2_data = MarlinMoEWeightData.make(
|
||||
w=w2, quant_type=quant_dtype, group_size=group_size, act_order=None
|
||||
w=w2,
|
||||
quant_type=quant_dtype,
|
||||
group_size=group_size,
|
||||
act_order=None,
|
||||
input_type=input_type,
|
||||
)
|
||||
|
||||
score = torch.randn((m, e), device="cuda", dtype=dtype)
|
||||
@@ -1487,6 +1547,12 @@ def test_batched_fused_marlin_moe(
|
||||
"quant_type_id": quant_dtype.id,
|
||||
"is_k_full": True,
|
||||
}
|
||||
if input_dtype is not None:
|
||||
kwargs["input_dtype"] = input_dtype
|
||||
if w1_data.a_scales_factor is not None:
|
||||
kwargs["input_global_scale1"] = w1_data.a_scales_factor
|
||||
if w2_data.a_scales_factor is not None:
|
||||
kwargs["input_global_scale2"] = w2_data.a_scales_factor
|
||||
|
||||
# Reference
|
||||
fused_marlin_moe_kwargs = kwargs | {
|
||||
@@ -1502,7 +1568,7 @@ def test_batched_fused_marlin_moe(
|
||||
pytest.skip("Cannot represent data in Batched Format.")
|
||||
marlin_output = br.run(a, kwargs)
|
||||
|
||||
torch.testing.assert_close(marlin_output, ref_marlin_output, atol=1e-3, rtol=0)
|
||||
torch.testing.assert_close(marlin_output, ref_marlin_output, atol=atol, rtol=0)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("m,n,k", [(32, 1024, 1024)])
|
||||
|
||||
@@ -957,9 +957,14 @@ def awq_marlin_moe_quant_config(
|
||||
group_size: int,
|
||||
w1_bias: torch.Tensor | None = None,
|
||||
w2_bias: torch.Tensor | None = None,
|
||||
a1_gscale: torch.Tensor | None = None,
|
||||
a2_gscale: torch.Tensor | None = None,
|
||||
) -> FusedMoEQuantConfig:
|
||||
"""
|
||||
Construct a quant config for awq marlin quantization.
|
||||
|
||||
a1_gscale / a2_gscale are optional global scales applied to activation
|
||||
quantization scales when Marlin runs with 8-bit activations.
|
||||
"""
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import GroupShape
|
||||
|
||||
@@ -977,8 +982,8 @@ def awq_marlin_moe_quant_config(
|
||||
raise ValueError(f"Unsupported weight_bits: {weight_bits}")
|
||||
|
||||
return FusedMoEQuantConfig(
|
||||
_a1=FusedMoEQuantDesc(dtype=None, shape=a_shape),
|
||||
_a2=FusedMoEQuantDesc(dtype=None, shape=a_shape),
|
||||
_a1=FusedMoEQuantDesc(dtype=None, shape=a_shape, alpha_or_gscale=a1_gscale),
|
||||
_a2=FusedMoEQuantDesc(dtype=None, shape=a_shape, alpha_or_gscale=a2_gscale),
|
||||
_w1=FusedMoEQuantDesc(weight_dtype, w_shape, w1_scale, None, w1_zp, w1_bias),
|
||||
_w2=FusedMoEQuantDesc(weight_dtype, w_shape, w2_scale, None, w2_zp, w2_bias),
|
||||
)
|
||||
|
||||
@@ -401,6 +401,8 @@ def batched_fused_marlin_moe(
|
||||
global_num_experts: int = -1,
|
||||
activation: MoEActivation = MoEActivation.SILU,
|
||||
expert_map: torch.Tensor | None = None,
|
||||
input_global_scale1: torch.Tensor | None = None,
|
||||
input_global_scale2: torch.Tensor | None = None,
|
||||
global_scale1: torch.Tensor | None = None,
|
||||
global_scale2: torch.Tensor | None = None,
|
||||
g_idx1: torch.Tensor | None = None,
|
||||
@@ -414,6 +416,7 @@ def batched_fused_marlin_moe(
|
||||
intermediate_cache2: torch.Tensor | None = None,
|
||||
is_k_full: bool = True,
|
||||
output: torch.Tensor | None = None,
|
||||
input_dtype: torch.dtype | None = None,
|
||||
inplace: bool = False,
|
||||
clamp_limit: float | None = None,
|
||||
) -> torch.Tensor:
|
||||
@@ -489,7 +492,15 @@ def batched_fused_marlin_moe(
|
||||
topk = 1
|
||||
|
||||
# TODO(varun) : Choose a decent block size like in fused_marlin_moe
|
||||
# Tune block_size_m based on expert capacity to reduce padding overhead.
|
||||
block_size_m = 64
|
||||
for b_m in [8, 16, 32, 48, 64]:
|
||||
if BATCH_TOKENS_MAX / b_m < 0.9:
|
||||
block_size_m = b_m
|
||||
break
|
||||
|
||||
if input_dtype is not None and input_dtype.itemsize == 1:
|
||||
block_size_m = max(block_size_m, 16)
|
||||
|
||||
sorted_token_ids, expert_ids, num_tokens_post_padded = batched_moe_align_block_size(
|
||||
max_tokens_per_batch=BATCH_TOKENS_MAX,
|
||||
@@ -525,6 +536,8 @@ def batched_fused_marlin_moe(
|
||||
sorted_token_ids=sorted_token_ids,
|
||||
expert_ids=expert_ids,
|
||||
num_tokens_post_padded=num_tokens_post_padded,
|
||||
input_global_scale1=input_global_scale1,
|
||||
input_global_scale2=input_global_scale2,
|
||||
global_scale1=global_scale1,
|
||||
global_scale2=global_scale2,
|
||||
g_idx1=g_idx1,
|
||||
@@ -537,6 +550,7 @@ def batched_fused_marlin_moe(
|
||||
intermediate_cache13=intermediate_cache13,
|
||||
intermediate_cache2=intermediate_cache2,
|
||||
output=output.view(-1, K) if output is not None else output,
|
||||
input_dtype=input_dtype,
|
||||
is_k_full=is_k_full,
|
||||
clamp_limit=clamp_limit,
|
||||
)
|
||||
@@ -635,6 +649,8 @@ class MarlinExpertsBase(mk.FusedMoEExpertsModular):
|
||||
@property
|
||||
def quant_type_id(self) -> int:
|
||||
if self.quant_config.use_int4_w4a16:
|
||||
if self.w1_zp is not None or self.w2_zp is not None:
|
||||
return scalar_types.uint4.id
|
||||
return scalar_types.uint4b8.id
|
||||
elif self.quant_config.use_mxfp4_w4a16 or self.quant_config.use_nvfp4_w4a16:
|
||||
return scalar_types.float4_e2m1f.id
|
||||
@@ -754,6 +770,10 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase):
|
||||
topk_ids=topk_ids,
|
||||
global_scale1=self.g1_alphas,
|
||||
global_scale2=self.g2_alphas,
|
||||
input_global_scale1=self.a1_gscale,
|
||||
input_global_scale2=self.a2_gscale,
|
||||
w1_zeros=self.w1_zp,
|
||||
w2_zeros=self.w2_zp,
|
||||
quant_type_id=self.quant_type_id,
|
||||
apply_router_weight_on_input=apply_router_weight_on_input,
|
||||
global_num_experts=global_num_experts,
|
||||
@@ -853,6 +873,10 @@ class MarlinExperts(LoRAExpertsMixin, MarlinExpertsBase):
|
||||
topk_ids=topk_ids,
|
||||
global_scale1=self.g1_alphas,
|
||||
global_scale2=self.g2_alphas,
|
||||
input_global_scale1=self.a1_gscale,
|
||||
input_global_scale2=self.a2_gscale,
|
||||
w1_zeros=self.w1_zp,
|
||||
w2_zeros=self.w2_zp,
|
||||
quant_type_id=self.quant_type_id,
|
||||
apply_router_weight_on_input=apply_router_weight_on_input,
|
||||
global_num_experts=global_num_experts,
|
||||
@@ -967,6 +991,8 @@ class BatchedMarlinExperts(MarlinExpertsBase):
|
||||
global_num_experts=global_num_experts,
|
||||
activation=activation,
|
||||
expert_map=expert_map,
|
||||
input_global_scale1=self.a1_gscale,
|
||||
input_global_scale2=self.a2_gscale,
|
||||
output=output,
|
||||
intermediate_cache13=workspace13,
|
||||
intermediate_cache2=workspace2,
|
||||
@@ -974,6 +1000,9 @@ class BatchedMarlinExperts(MarlinExpertsBase):
|
||||
g_idx2=self.w2_g_idx,
|
||||
sort_indices1=self.w13_g_idx_sort_indices,
|
||||
sort_indices2=self.w2_g_idx_sort_indices,
|
||||
w1_zeros=self.w1_zp,
|
||||
w2_zeros=self.w2_zp,
|
||||
input_dtype=self.input_dtype,
|
||||
is_k_full=self.is_k_full,
|
||||
clamp_limit=self.gemm1_clamp_limit,
|
||||
)
|
||||
|
||||
@@ -21,6 +21,7 @@ from vllm.model_executor.layers.quantization.utils.marlin_utils import (
|
||||
marlin_act_int8_process_scales,
|
||||
marlin_moe_permute_scales,
|
||||
marlin_permute_bias,
|
||||
moe_awq_to_marlin_zero_points,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
QuantKey,
|
||||
@@ -28,6 +29,7 @@ from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from vllm.model_executor.layers.quantization.auto_gptq import AutoGPTQConfig
|
||||
from vllm.model_executor.layers.quantization.awq_marlin import AWQMarlinConfig
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -147,7 +149,6 @@ def make_wna16_moe_kernel(
|
||||
moe_quant_config: FusedMoEQuantConfig,
|
||||
moe_config: FusedMoEConfig,
|
||||
experts_cls: type[mk.FusedMoEExperts] | None,
|
||||
layer: torch.nn.Module,
|
||||
is_k_full: bool,
|
||||
w13_g_idx: torch.Tensor | None,
|
||||
w2_g_idx: torch.Tensor | None,
|
||||
@@ -220,6 +221,8 @@ def _process_weights_marlin(
|
||||
w2_scales: torch.Tensor,
|
||||
w13_g_idx: torch.Tensor,
|
||||
w2_g_idx: torch.Tensor,
|
||||
w13_qzeros: torch.Tensor | None = None,
|
||||
w2_qzeros: torch.Tensor | None = None,
|
||||
w13_bias: torch.Tensor | None = None,
|
||||
w2_bias: torch.Tensor | None = None,
|
||||
) -> tuple[
|
||||
@@ -231,6 +234,8 @@ def _process_weights_marlin(
|
||||
torch.Tensor, # w2_g_idx
|
||||
torch.Tensor, # w13_g_idx_sort_indices
|
||||
torch.Tensor, # w2_g_idx_sort_indices
|
||||
torch.Tensor | None, # w13_qzeros
|
||||
torch.Tensor | None, # w2_qzeros
|
||||
torch.Tensor | None, # w13_input_global_scale
|
||||
torch.Tensor | None, # w2_input_global_scale
|
||||
torch.Tensor | None, # w13_bias
|
||||
@@ -368,6 +373,151 @@ def _process_weights_marlin(
|
||||
w2_g_idx,
|
||||
w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices,
|
||||
w13_qzeros,
|
||||
w2_qzeros,
|
||||
w13_input_global_scale,
|
||||
w2_input_global_scale,
|
||||
w13_bias_out,
|
||||
w2_bias_out,
|
||||
)
|
||||
|
||||
|
||||
def _process_awq_weights_marlin(
|
||||
layer: torch.nn.Module,
|
||||
quant_config: "AWQMarlinConfig",
|
||||
input_dtype: torch.dtype | None,
|
||||
w13_qweight: torch.Tensor,
|
||||
w2_qweight: torch.Tensor,
|
||||
w13_scales: torch.Tensor,
|
||||
w2_scales: torch.Tensor,
|
||||
w13_qzeros: torch.Tensor,
|
||||
w2_qzeros: torch.Tensor,
|
||||
w13_bias: torch.Tensor | None = None,
|
||||
w2_bias: torch.Tensor | None = None,
|
||||
) -> tuple[
|
||||
torch.Tensor, # w13_qweight
|
||||
torch.Tensor, # w2_qweight
|
||||
torch.Tensor, # w13_scales
|
||||
torch.Tensor, # w2_scales
|
||||
torch.Tensor | None, # w13_g_idx
|
||||
torch.Tensor | None, # w2_g_idx
|
||||
torch.Tensor | None, # w13_g_idx_sort_indices
|
||||
torch.Tensor | None, # w2_g_idx_sort_indices
|
||||
torch.Tensor | None, # w13_qzeros
|
||||
torch.Tensor | None, # w2_qzeros
|
||||
torch.Tensor | None, # w13_input_global_scale
|
||||
torch.Tensor | None, # w2_input_global_scale
|
||||
torch.Tensor | None, # w13_bias
|
||||
torch.Tensor | None, # w2_bias
|
||||
]:
|
||||
"""AWQ-specific Marlin weight post-processing.
|
||||
|
||||
AWQ checkpoints use a different packing order than GPTQ, so they need
|
||||
AWQ-specific weight repacking and zero-point conversion before Marlin runs.
|
||||
"""
|
||||
num_experts = w13_qweight.shape[0]
|
||||
device = w13_qweight.device
|
||||
is_a_8bit = input_dtype is not None and input_dtype.itemsize == 1
|
||||
w13_input_global_scale: torch.Tensor | None = None
|
||||
w2_input_global_scale: torch.Tensor | None = None
|
||||
w13_bias_out: torch.Tensor | None = None
|
||||
w2_bias_out: torch.Tensor | None = None
|
||||
|
||||
if input_dtype == torch.float8_e4m3fn:
|
||||
ops.marlin_int4_fp8_preprocess(
|
||||
w13_qweight.view(-1, w13_qweight.size(2)),
|
||||
w13_qzeros.view(-1, w13_qzeros.size(2)),
|
||||
inplace=True,
|
||||
)
|
||||
ops.marlin_int4_fp8_preprocess(
|
||||
w2_qweight.view(-1, w2_qweight.size(2)),
|
||||
w2_qzeros.view(-1, w2_qzeros.size(2)),
|
||||
inplace=True,
|
||||
)
|
||||
w13_scales = w13_scales.data * 512
|
||||
w2_scales = w2_scales.data * 512
|
||||
|
||||
w13_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
)
|
||||
w2_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
)
|
||||
|
||||
marlin_w13_qweight = ops.awq_marlin_moe_repack(
|
||||
w13_qweight,
|
||||
w13_g_idx_sort_indices,
|
||||
size_k=w13_qweight.shape[1],
|
||||
size_n=w13_qweight.shape[2] * quant_config.pack_factor,
|
||||
num_bits=quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
marlin_w2_qweight = ops.awq_marlin_moe_repack(
|
||||
w2_qweight,
|
||||
w2_g_idx_sort_indices,
|
||||
size_k=w2_qweight.shape[1],
|
||||
size_n=w2_qweight.shape[2] * quant_config.pack_factor,
|
||||
num_bits=quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
|
||||
marlin_w13_scales = marlin_moe_permute_scales(
|
||||
s=w13_scales,
|
||||
size_k=layer.intermediate_size_per_partition,
|
||||
size_n=w13_scales.shape[2],
|
||||
group_size=quant_config.group_size,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
if input_dtype == torch.int8 and layer.num_groups_w13 > 1:
|
||||
marlin_w13_scales, w13_input_global_scale = marlin_act_int8_process_scales(
|
||||
marlin_w13_scales
|
||||
)
|
||||
|
||||
marlin_w2_scales = marlin_moe_permute_scales(
|
||||
s=w2_scales,
|
||||
size_k=layer.intermediate_size_per_partition,
|
||||
size_n=w2_scales.shape[2],
|
||||
group_size=quant_config.group_size,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
if input_dtype == torch.int8 and layer.num_groups_w2 > 1:
|
||||
marlin_w2_scales, w2_input_global_scale = marlin_act_int8_process_scales(
|
||||
marlin_w2_scales
|
||||
)
|
||||
|
||||
marlin_w13_qzeros = moe_awq_to_marlin_zero_points(
|
||||
w13_qzeros,
|
||||
size_k=w13_qzeros.shape[1],
|
||||
size_n=w13_qzeros.shape[2] * quant_config.pack_factor,
|
||||
num_bits=quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
marlin_w2_qzeros = moe_awq_to_marlin_zero_points(
|
||||
w2_qzeros,
|
||||
size_k=w2_qzeros.shape[1],
|
||||
size_n=w2_qzeros.shape[2] * quant_config.pack_factor,
|
||||
num_bits=quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
|
||||
if w13_bias is not None:
|
||||
w13_bias_out = marlin_permute_bias(w13_bias)
|
||||
if w2_bias is not None:
|
||||
w2_bias_out = marlin_permute_bias(w2_bias)
|
||||
|
||||
return (
|
||||
marlin_w13_qweight,
|
||||
marlin_w2_qweight,
|
||||
marlin_w13_scales,
|
||||
marlin_w2_scales,
|
||||
None,
|
||||
None,
|
||||
w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices,
|
||||
marlin_w13_qzeros,
|
||||
marlin_w2_qzeros,
|
||||
w13_input_global_scale,
|
||||
w2_input_global_scale,
|
||||
w13_bias_out,
|
||||
@@ -384,8 +534,10 @@ def convert_to_wna16_moe_kernel_format(
|
||||
w2: torch.Tensor,
|
||||
w13_scale: torch.Tensor,
|
||||
w2_scale: torch.Tensor,
|
||||
w13_g_idx: torch.Tensor,
|
||||
w2_g_idx: torch.Tensor,
|
||||
w13_g_idx: torch.Tensor | None = None,
|
||||
w2_g_idx: torch.Tensor | None = None,
|
||||
w13_qzeros: torch.Tensor | None = None,
|
||||
w2_qzeros: torch.Tensor | None = None,
|
||||
w13_bias: torch.Tensor | None = None,
|
||||
w2_bias: torch.Tensor | None = None,
|
||||
) -> tuple[
|
||||
@@ -397,6 +549,8 @@ def convert_to_wna16_moe_kernel_format(
|
||||
torch.Tensor | None, # w2_g_idx
|
||||
torch.Tensor | None, # w13_g_idx_sort_indices
|
||||
torch.Tensor | None, # w2_g_idx_sort_indices
|
||||
torch.Tensor | None, # w13_qzeros
|
||||
torch.Tensor | None, # w2_qzeros
|
||||
torch.Tensor | None, # w13_input_global_scale
|
||||
torch.Tensor | None, # w2_input_global_scale
|
||||
torch.Tensor | None, # w13_bias
|
||||
@@ -420,12 +574,35 @@ def convert_to_wna16_moe_kernel_format(
|
||||
from vllm.model_executor.layers.quantization.auto_gptq import (
|
||||
AutoGPTQConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.awq_marlin import (
|
||||
AWQMarlinConfig,
|
||||
)
|
||||
|
||||
if isinstance(quant_config, AWQMarlinConfig):
|
||||
if w13_qzeros is None or w2_qzeros is None:
|
||||
raise ValueError("AWQ Marlin MoE requires zero-point tensors.")
|
||||
return _process_awq_weights_marlin(
|
||||
layer,
|
||||
quant_config,
|
||||
input_dtype,
|
||||
w13,
|
||||
w2,
|
||||
w13_scale,
|
||||
w2_scale,
|
||||
w13_qzeros,
|
||||
w2_qzeros,
|
||||
w13_bias,
|
||||
w2_bias,
|
||||
)
|
||||
|
||||
if not isinstance(quant_config, AutoGPTQConfig):
|
||||
raise TypeError(
|
||||
"Marlin WNA16 MoE backend requires AutoGPTQConfig, got "
|
||||
"Marlin WNA16 MoE backend requires AutoGPTQConfig or "
|
||||
"AWQMarlinConfig, got "
|
||||
f"{type(quant_config).__name__}."
|
||||
)
|
||||
if w13_g_idx is None or w2_g_idx is None:
|
||||
raise ValueError("GPTQ Marlin MoE requires g_idx tensors.")
|
||||
return _process_weights_marlin(
|
||||
layer,
|
||||
quant_config,
|
||||
@@ -436,6 +613,8 @@ def convert_to_wna16_moe_kernel_format(
|
||||
w2_scale,
|
||||
w13_g_idx,
|
||||
w2_g_idx,
|
||||
w13_qzeros,
|
||||
w2_qzeros,
|
||||
w13_bias,
|
||||
w2_bias,
|
||||
)
|
||||
|
||||
@@ -659,6 +659,8 @@ class AutoGPTQMoEMethod(FusedMoEMethodBase):
|
||||
w2_g_idx,
|
||||
w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices,
|
||||
_w13_qzeros,
|
||||
_w2_qzeros,
|
||||
w13_input_global_scale,
|
||||
w2_input_global_scale,
|
||||
w13_bias,
|
||||
@@ -729,7 +731,6 @@ class AutoGPTQMoEMethod(FusedMoEMethodBase):
|
||||
moe_quant_config=self.moe_quant_config,
|
||||
moe_config=self.moe,
|
||||
experts_cls=self.experts_cls,
|
||||
layer=layer,
|
||||
is_k_full=self.is_k_full,
|
||||
w13_g_idx=layer.w13_g_idx,
|
||||
w2_g_idx=layer.w2_g_idx,
|
||||
|
||||
@@ -9,7 +9,6 @@ from torch.nn import Parameter
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
import vllm.model_executor.layers.fused_moe # noqa
|
||||
from vllm import _custom_ops as ops
|
||||
from vllm import envs
|
||||
from vllm.logger import init_logger
|
||||
from vllm.model_executor.kernels.linear import (
|
||||
@@ -27,7 +26,11 @@ from vllm.model_executor.layers.fused_moe.config import (
|
||||
FusedMoEConfig,
|
||||
FusedMoEQuantConfig,
|
||||
)
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import fused_marlin_moe
|
||||
from vllm.model_executor.layers.fused_moe.oracle.int_wna16 import (
|
||||
convert_to_wna16_moe_kernel_format,
|
||||
make_wna16_moe_kernel,
|
||||
select_wna16_moe_backend,
|
||||
)
|
||||
from vllm.model_executor.layers.linear import (
|
||||
LinearBase,
|
||||
LinearMethodBase,
|
||||
@@ -45,14 +48,13 @@ from vllm.model_executor.layers.quantization.utils.marlin_utils import (
|
||||
check_marlin_supports_layer,
|
||||
check_moe_marlin_supports_layer,
|
||||
get_marlin_input_dtype,
|
||||
marlin_act_int8_process_scales,
|
||||
marlin_make_workspace_new,
|
||||
marlin_moe_permute_scales,
|
||||
marlin_permute_bias,
|
||||
moe_awq_to_marlin_zero_points,
|
||||
verify_marlin_supported,
|
||||
)
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import is_layer_skipped
|
||||
from vllm.model_executor.layers.quantization.utils.quant_utils import (
|
||||
is_layer_skipped,
|
||||
kInt4Static,
|
||||
)
|
||||
from vllm.model_executor.layers.vocab_parallel_embedding import ParallelLMHead
|
||||
from vllm.model_executor.parameter import GroupQuantScaleParameter, PackedvLLMParameter
|
||||
from vllm.platforms import current_platform
|
||||
@@ -72,6 +74,19 @@ logger = init_logger(__name__)
|
||||
_REVERSE_AWQ_PACK_ORDER = [0, 4, 1, 5, 2, 6, 3, 7]
|
||||
|
||||
|
||||
def _replace_or_register_parameter(
|
||||
layer: torch.nn.Module,
|
||||
name: str,
|
||||
value: torch.Tensor | None,
|
||||
) -> None:
|
||||
if value is None:
|
||||
return
|
||||
if hasattr(layer, name):
|
||||
replace_parameter(layer, name, value)
|
||||
else:
|
||||
layer.register_parameter(name, Parameter(value, requires_grad=False))
|
||||
|
||||
|
||||
def _convert_awq_to_standard_format(
|
||||
layer: torch.nn.Module,
|
||||
w_q_name: str,
|
||||
@@ -505,6 +520,9 @@ class AWQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
self.quant_type = scalar_types.uint4
|
||||
self.input_dtype = None
|
||||
self.use_marlin = True
|
||||
self.wna16_moe_backend, self.experts_cls = select_wna16_moe_backend(
|
||||
moe, kInt4Static, quant_config.weight_bits
|
||||
)
|
||||
|
||||
def create_weights(
|
||||
self,
|
||||
@@ -608,52 +626,38 @@ class AWQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
layer.workspace = marlin_make_workspace_new(device, 4)
|
||||
|
||||
def process_weights_after_loading(self, layer: RoutedExperts) -> None:
|
||||
num_experts = layer.w13_qweight.shape[0]
|
||||
device = layer.w13_qweight.device
|
||||
is_a_8bit = self.input_dtype is not None and self.input_dtype.itemsize == 1
|
||||
|
||||
if self.input_dtype == torch.float8_e4m3fn:
|
||||
ops.marlin_int4_fp8_preprocess(
|
||||
layer.w13_qweight.view(-1, layer.w13_qweight.size(2)),
|
||||
layer.w13_qzeros.view(-1, layer.w13_qzeros.size(2)),
|
||||
inplace=True,
|
||||
)
|
||||
ops.marlin_int4_fp8_preprocess(
|
||||
layer.w2_qweight.view(-1, layer.w2_qweight.size(2)),
|
||||
layer.w2_qzeros.view(-1, layer.w2_qzeros.size(2)),
|
||||
inplace=True,
|
||||
)
|
||||
layer.w13_scales.data = layer.w13_scales.data * 512
|
||||
layer.w2_scales.data = layer.w2_scales.data * 512
|
||||
|
||||
layer.w13_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
)
|
||||
layer.w2_g_idx_sort_indices = torch.nn.Parameter(
|
||||
torch.empty((num_experts, 0), dtype=torch.int32, device=device),
|
||||
requires_grad=False,
|
||||
(
|
||||
w13,
|
||||
w2,
|
||||
w13_scale,
|
||||
w2_scale,
|
||||
w13_g_idx,
|
||||
w2_g_idx,
|
||||
w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices,
|
||||
w13_qzeros,
|
||||
w2_qzeros,
|
||||
w13_input_global_scale,
|
||||
w2_input_global_scale,
|
||||
w13_bias,
|
||||
w2_bias,
|
||||
) = convert_to_wna16_moe_kernel_format(
|
||||
backend=self.wna16_moe_backend,
|
||||
layer=layer,
|
||||
quant_config=self.quant_config,
|
||||
input_dtype=self.input_dtype,
|
||||
w13=layer.w13_qweight,
|
||||
w2=layer.w2_qweight,
|
||||
w13_scale=layer.w13_scales,
|
||||
w2_scale=layer.w2_scales,
|
||||
w13_qzeros=layer.w13_qzeros,
|
||||
w2_qzeros=layer.w2_qzeros,
|
||||
w13_bias=getattr(layer, "w13_bias", None),
|
||||
w2_bias=getattr(layer, "w2_bias", None),
|
||||
)
|
||||
|
||||
marlin_w13_qweight = ops.awq_marlin_moe_repack(
|
||||
layer.w13_qweight,
|
||||
layer.w13_g_idx_sort_indices,
|
||||
size_k=layer.w13_qweight.shape[1],
|
||||
size_n=layer.w13_qweight.shape[2] * self.quant_config.pack_factor,
|
||||
num_bits=self.quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
replace_parameter(layer, "w13_qweight", marlin_w13_qweight)
|
||||
|
||||
marlin_w2_qweight = ops.awq_marlin_moe_repack(
|
||||
layer.w2_qweight,
|
||||
layer.w2_g_idx_sort_indices,
|
||||
size_k=layer.w2_qweight.shape[1],
|
||||
size_n=layer.w2_qweight.shape[2] * self.quant_config.pack_factor,
|
||||
num_bits=self.quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
)
|
||||
replace_parameter(layer, "w2_qweight", marlin_w2_qweight)
|
||||
replace_parameter(layer, "w13_qweight", w13)
|
||||
replace_parameter(layer, "w2_qweight", w2)
|
||||
|
||||
# The modular kernel expects w13_weight and w2_weight,
|
||||
# but AWQ uses w13_qweight and w2_qweight
|
||||
@@ -662,70 +666,46 @@ class AWQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
# Alias for modular kernel
|
||||
layer.w2_weight = layer.w2_qweight
|
||||
|
||||
# Why does this take the intermediate size for size_k?
|
||||
marlin_w13_scales = marlin_moe_permute_scales(
|
||||
s=layer.w13_scales,
|
||||
size_k=layer.intermediate_size_per_partition,
|
||||
size_n=layer.w13_scales.shape[2],
|
||||
group_size=self.quant_config.group_size,
|
||||
is_a_8bit=is_a_8bit,
|
||||
replace_parameter(layer, "w13_scales", w13_scale)
|
||||
replace_parameter(layer, "w2_scales", w2_scale)
|
||||
_replace_or_register_parameter(
|
||||
layer, "w13_g_idx_sort_indices", w13_g_idx_sort_indices
|
||||
)
|
||||
if self.input_dtype == torch.int8 and layer.num_groups_w13 > 1:
|
||||
marlin_w13_scales, w13_input_global_scale = marlin_act_int8_process_scales(
|
||||
marlin_w13_scales
|
||||
)
|
||||
layer.register_parameter(
|
||||
"w13_input_global_scale",
|
||||
Parameter(w13_input_global_scale, requires_grad=False),
|
||||
)
|
||||
|
||||
replace_parameter(layer, "w13_scales", marlin_w13_scales)
|
||||
|
||||
marlin_w2_scales = marlin_moe_permute_scales(
|
||||
s=layer.w2_scales,
|
||||
size_k=layer.intermediate_size_per_partition,
|
||||
size_n=layer.w2_scales.shape[2],
|
||||
group_size=self.quant_config.group_size,
|
||||
is_a_8bit=is_a_8bit,
|
||||
_replace_or_register_parameter(
|
||||
layer, "w2_g_idx_sort_indices", w2_g_idx_sort_indices
|
||||
)
|
||||
if self.input_dtype == torch.int8 and layer.num_groups_w2 > 1:
|
||||
marlin_w2_scales, w2_input_global_scale = marlin_act_int8_process_scales(
|
||||
marlin_w2_scales
|
||||
)
|
||||
layer.register_parameter(
|
||||
"w2_input_global_scale",
|
||||
Parameter(w2_input_global_scale, requires_grad=False),
|
||||
)
|
||||
|
||||
replace_parameter(layer, "w2_scales", marlin_w2_scales)
|
||||
|
||||
marlin_w13_zp = moe_awq_to_marlin_zero_points(
|
||||
layer.w13_qzeros,
|
||||
size_k=layer.w13_qzeros.shape[1],
|
||||
size_n=layer.w13_qzeros.shape[2] * self.quant_config.pack_factor,
|
||||
num_bits=self.quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
_replace_or_register_parameter(layer, "w13_g_idx", w13_g_idx)
|
||||
_replace_or_register_parameter(layer, "w2_g_idx", w2_g_idx)
|
||||
_replace_or_register_parameter(layer, "w13_qzeros", w13_qzeros)
|
||||
_replace_or_register_parameter(layer, "w2_qzeros", w2_qzeros)
|
||||
_replace_or_register_parameter(
|
||||
layer, "w13_input_global_scale", w13_input_global_scale
|
||||
)
|
||||
replace_parameter(layer, "w13_qzeros", marlin_w13_zp)
|
||||
|
||||
marlin_w2_zp = moe_awq_to_marlin_zero_points(
|
||||
layer.w2_qzeros,
|
||||
size_k=layer.w2_qzeros.shape[1],
|
||||
size_n=layer.w2_qzeros.shape[2] * self.quant_config.pack_factor,
|
||||
num_bits=self.quant_config.weight_bits,
|
||||
is_a_8bit=is_a_8bit,
|
||||
_replace_or_register_parameter(
|
||||
layer, "w2_input_global_scale", w2_input_global_scale
|
||||
)
|
||||
replace_parameter(layer, "w2_qzeros", marlin_w2_zp)
|
||||
_replace_or_register_parameter(layer, "w13_bias", w13_bias)
|
||||
_replace_or_register_parameter(layer, "w2_bias", w2_bias)
|
||||
|
||||
if hasattr(layer, "w13_bias") and layer.w13_bias is not None:
|
||||
layer.w13_bias.data = marlin_permute_bias(layer.w13_bias)
|
||||
self._setup_kernel(layer)
|
||||
|
||||
if hasattr(layer, "w2_bias") and layer.w2_bias is not None:
|
||||
layer.w2_bias.data = marlin_permute_bias(layer.w2_bias)
|
||||
def _setup_kernel(self, layer: RoutedExperts) -> None:
|
||||
"""Build the FusedMoEKernel for this layer."""
|
||||
|
||||
def get_fused_moe_quant_config(
|
||||
self, layer: RoutedExperts
|
||||
) -> FusedMoEQuantConfig | None:
|
||||
self.moe_quant_config = self.get_fused_moe_quant_config(layer)
|
||||
self.moe_kernel = make_wna16_moe_kernel(
|
||||
moe_quant_config=self.moe_quant_config,
|
||||
moe_config=self.moe,
|
||||
experts_cls=self.experts_cls,
|
||||
is_k_full=self.is_k_full,
|
||||
w13_g_idx=getattr(layer, "w13_g_idx", None),
|
||||
w2_g_idx=getattr(layer, "w2_g_idx", None),
|
||||
w13_g_idx_sort_indices=layer.w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices=layer.w2_g_idx_sort_indices,
|
||||
routing_tables=layer._expert_routing_tables(),
|
||||
)
|
||||
|
||||
def get_fused_moe_quant_config(self, layer: RoutedExperts) -> FusedMoEQuantConfig:
|
||||
from vllm.model_executor.layers.fused_moe.config import (
|
||||
awq_marlin_moe_quant_config,
|
||||
)
|
||||
@@ -743,6 +723,8 @@ class AWQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
else None,
|
||||
w1_bias=getattr(layer, "w13_bias", None),
|
||||
w2_bias=getattr(layer, "w2_bias", None),
|
||||
a1_gscale=getattr(layer, "w13_input_global_scale", None),
|
||||
a2_gscale=getattr(layer, "w2_input_global_scale", None),
|
||||
)
|
||||
|
||||
def select_gemm_impl(
|
||||
@@ -750,67 +732,11 @@ class AWQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
prepare_finalize,
|
||||
layer: RoutedExperts,
|
||||
):
|
||||
"""
|
||||
Select the GEMM implementation for AWQ-Marlin MoE.
|
||||
Returns MarlinExperts configured for AWQ quantization.
|
||||
This is ONLY used when LoRA is enabled.
|
||||
Without LoRA, AWQ uses its own apply() method.
|
||||
"""
|
||||
# Only use modular kernels when LoRA is enabled
|
||||
# Without LoRA, AWQ's own apply() method works fine and is more efficient
|
||||
if not self.moe.is_lora_enabled:
|
||||
raise NotImplementedError(
|
||||
"AWQ-Marlin uses its own apply() method when LoRA is not enabled. "
|
||||
"Modular kernels are only used for LoRA support."
|
||||
)
|
||||
|
||||
from vllm.model_executor.layers.fused_moe import modular_kernel as mk
|
||||
from vllm.model_executor.layers.fused_moe.experts.marlin_moe import (
|
||||
BatchedMarlinExperts,
|
||||
MarlinExperts,
|
||||
raise ValueError(
|
||||
f"{self.__class__.__name__} uses the new modular kernel "
|
||||
"initialization logic. This function should not be called."
|
||||
)
|
||||
|
||||
# Ensure quant config is initialized
|
||||
assert self.moe_quant_config is not None, (
|
||||
"moe_quant_config must be initialized before select_gemm_impl"
|
||||
)
|
||||
|
||||
w13_g_idx = getattr(layer, "w13_g_idx", None)
|
||||
w2_g_idx = getattr(layer, "w2_g_idx", None)
|
||||
w13_g_idx_sort_indices = getattr(layer, "w13_g_idx_sort_indices", None)
|
||||
w2_g_idx_sort_indices = getattr(layer, "w2_g_idx_sort_indices", None)
|
||||
|
||||
# Check if using batched expert format (for Expert Parallelism)
|
||||
if (
|
||||
prepare_finalize.activation_format
|
||||
== mk.FusedMoEActivationFormat.BatchedExperts
|
||||
):
|
||||
# For batched format, use BatchedMarlinExperts
|
||||
max_num_tokens_per_rank = prepare_finalize.max_num_tokens_per_rank()
|
||||
assert max_num_tokens_per_rank is not None
|
||||
return BatchedMarlinExperts(
|
||||
max_num_tokens=max_num_tokens_per_rank,
|
||||
num_dispatchers=prepare_finalize.num_dispatchers(),
|
||||
moe_config=self.moe,
|
||||
quant_config=self.moe_quant_config,
|
||||
w13_g_idx=w13_g_idx,
|
||||
w2_g_idx=w2_g_idx,
|
||||
w13_g_idx_sort_indices=w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices=w2_g_idx_sort_indices,
|
||||
is_k_full=self.is_k_full,
|
||||
)
|
||||
else:
|
||||
# Standard Marlin experts for AWQ
|
||||
return MarlinExperts(
|
||||
moe_config=self.moe,
|
||||
quant_config=self.moe_quant_config,
|
||||
w13_g_idx=w13_g_idx,
|
||||
w2_g_idx=w2_g_idx,
|
||||
w13_g_idx_sort_indices=w13_g_idx_sort_indices,
|
||||
w2_g_idx_sort_indices=w2_g_idx_sort_indices,
|
||||
is_k_full=self.is_k_full,
|
||||
)
|
||||
|
||||
def apply(
|
||||
self,
|
||||
layer: RoutedExperts,
|
||||
@@ -820,25 +746,18 @@ class AWQMarlinMoEMethod(FusedMoEMethodBase):
|
||||
shared_experts: SharedExperts | None,
|
||||
shared_experts_input: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
return fused_marlin_moe(
|
||||
x,
|
||||
layer.w13_qweight,
|
||||
layer.w2_qweight,
|
||||
getattr(layer, "w13_bias", None),
|
||||
getattr(layer, "w2_bias", None),
|
||||
layer.w13_scales,
|
||||
layer.w2_scales,
|
||||
topk_weights,
|
||||
topk_ids,
|
||||
input_global_scale1=getattr(layer, "w13_input_global_scale", None),
|
||||
input_global_scale2=getattr(layer, "w2_input_global_scale", None),
|
||||
quant_type_id=self.quant_type.id,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
assert not self.is_monolithic
|
||||
assert self.moe_kernel is not None
|
||||
return self.moe_kernel.apply(
|
||||
hidden_states=x,
|
||||
w1=layer.w13_qweight,
|
||||
w2=layer.w2_qweight,
|
||||
topk_weights=topk_weights,
|
||||
topk_ids=topk_ids,
|
||||
activation=layer.activation,
|
||||
global_num_experts=layer.global_num_experts,
|
||||
apply_router_weight_on_input=layer.apply_router_weight_on_input,
|
||||
expert_map=layer.expert_map,
|
||||
w1_zeros=layer.w13_qzeros,
|
||||
w2_zeros=layer.w2_qzeros,
|
||||
workspace=layer.workspace,
|
||||
input_dtype=self.input_dtype,
|
||||
inplace=not self.moe.disable_inplace,
|
||||
shared_experts=shared_experts,
|
||||
shared_experts_input=shared_experts_input,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user