[Quantization] Fix GPTQ scheme attachment broken by LinearBase.scheme default (#34962)
Co-authored-by: Mohammad Angkad <mohammad.angkad@radixark.ai>
This commit is contained in:
co-authored by
Mohammad Angkad
parent
2ee0d38a85
commit
6ab4b99bc2
@@ -155,8 +155,13 @@ class LinearBase(torch.nn.Module):
|
|||||||
quant_config: Quantization configure.
|
quant_config: Quantization configure.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Set by quant methods that attach a per-layer scheme (e.g. Quark) inside
|
# Set by quant methods that attach a per-layer scheme, eagerly in
|
||||||
# get_quant_method(), which runs before create_weights() picks the loader.
|
# get_quant_method(), which runs before create_weights() picks the loader,
|
||||||
|
# or lazily inside create_weights() itself (GPTQ). The default is what lets
|
||||||
|
# callers probe with `is None`; a hasattr() probe answers "yes" once it
|
||||||
|
# exists. Schemes must stay plain objects -- nn.Module.__setattr__ files a
|
||||||
|
# Module value under self._modules, which this default then shadows on read.
|
||||||
|
# VocabParallelEmbedding and FusedMoE carry the same default.
|
||||||
scheme = None
|
scheme = None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
|
|||||||
@@ -228,6 +228,9 @@ class FusedMoE(torch.nn.Module):
|
|||||||
# backend resolution distinguish them from routed experts.
|
# backend resolution distinguish them from routed experts.
|
||||||
is_shared_fused_moe = False
|
is_shared_fused_moe = False
|
||||||
|
|
||||||
|
# Attached by quant methods for a quantized MoE layer; see LinearBase.scheme.
|
||||||
|
scheme = None
|
||||||
|
|
||||||
_skip_aiter_moe_shuffle: bool = False
|
_skip_aiter_moe_shuffle: bool = False
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
@@ -1082,7 +1085,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
# TODO (mgoin): check self.quant_method.quant_config.quant_format
|
# TODO (mgoin): check self.quant_method.quant_config.quant_format
|
||||||
# against known CompressionFormat enum values that have this quality
|
# against known CompressionFormat enum values that have this quality
|
||||||
method = self.quant_method
|
method = self.quant_method
|
||||||
if hasattr(self, "scheme"):
|
if self.scheme is not None:
|
||||||
method = self.scheme
|
method = self.scheme
|
||||||
if method.__class__.__name__ == "KTEPWrapperMethod":
|
if method.__class__.__name__ == "KTEPWrapperMethod":
|
||||||
method = method.gpu_method
|
method = method.gpu_method
|
||||||
@@ -1349,7 +1352,7 @@ class FusedMoE(torch.nn.Module):
|
|||||||
# TODO: check self.quant_method.quant_config.quant_format
|
# TODO: check self.quant_method.quant_config.quant_format
|
||||||
# against known CompressionFormat enum values that have this quality
|
# against known CompressionFormat enum values that have this quality
|
||||||
method = self.quant_method
|
method = self.quant_method
|
||||||
if hasattr(self, "scheme"):
|
if self.scheme is not None:
|
||||||
method = self.scheme
|
method = self.scheme
|
||||||
if isinstance(method, Fp8MoEMethod) and (
|
if isinstance(method, Fp8MoEMethod) and (
|
||||||
get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
get_moe_runner_backend().is_flashinfer_trtllm_routed()
|
||||||
|
|||||||
@@ -471,7 +471,7 @@ class GPTQLinearMethod(LinearMethodBase):
|
|||||||
params_dtype: torch.dtype,
|
params_dtype: torch.dtype,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
):
|
):
|
||||||
if not hasattr(layer, "scheme"):
|
if layer.scheme is None:
|
||||||
layer.scheme = self.quant_config.get_linear_scheme(layer)
|
layer.scheme = self.quant_config.get_linear_scheme(layer)
|
||||||
weight_loader = extra_weight_attrs.get("weight_loader")
|
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||||
layer.scheme.create_weights(
|
layer.scheme.create_weights(
|
||||||
@@ -511,7 +511,7 @@ class GPTQMoEMethod(FusedMoEMethodBase):
|
|||||||
params_dtype: torch.dtype,
|
params_dtype: torch.dtype,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
):
|
):
|
||||||
if not hasattr(layer, "scheme"):
|
if layer.scheme is None:
|
||||||
layer.scheme = self.quant_config.get_moe_scheme(layer)
|
layer.scheme = self.quant_config.get_moe_scheme(layer)
|
||||||
layer.scheme.create_weights(
|
layer.scheme.create_weights(
|
||||||
layer=layer,
|
layer=layer,
|
||||||
@@ -563,7 +563,7 @@ class GPTQMarlinLinearMethod(LinearMethodBase):
|
|||||||
params_dtype: torch.dtype,
|
params_dtype: torch.dtype,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
) -> None:
|
) -> None:
|
||||||
if not hasattr(layer, "scheme"):
|
if layer.scheme is None:
|
||||||
layer.scheme = self.quant_config.get_linear_scheme(layer)
|
layer.scheme = self.quant_config.get_linear_scheme(layer)
|
||||||
weight_loader = extra_weight_attrs.get("weight_loader")
|
weight_loader = extra_weight_attrs.get("weight_loader")
|
||||||
layer.scheme.create_weights(
|
layer.scheme.create_weights(
|
||||||
@@ -603,7 +603,7 @@ class GPTQMarlinMoEMethod(FusedMoEMethodBase):
|
|||||||
params_dtype: torch.dtype,
|
params_dtype: torch.dtype,
|
||||||
**extra_weight_attrs,
|
**extra_weight_attrs,
|
||||||
):
|
):
|
||||||
if not hasattr(layer, "scheme"):
|
if layer.scheme is None:
|
||||||
layer.scheme = self.quant_config.get_moe_scheme(layer)
|
layer.scheme = self.quant_config.get_moe_scheme(layer)
|
||||||
layer.scheme.create_weights(
|
layer.scheme.create_weights(
|
||||||
layer=layer,
|
layer=layer,
|
||||||
|
|||||||
@@ -224,6 +224,10 @@ class VocabParallelEmbedding(torch.nn.Module):
|
|||||||
prefix: full name of the layer in the state dict
|
prefix: full name of the layer in the state dict
|
||||||
""" # noqa: E501
|
""" # noqa: E501
|
||||||
|
|
||||||
|
# Attached by quant methods for a quantized ParallelLMHead; see
|
||||||
|
# LinearBase.scheme.
|
||||||
|
scheme = None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
num_embeddings: int,
|
num_embeddings: int,
|
||||||
|
|||||||
@@ -0,0 +1,60 @@
|
|||||||
|
"""GPTQ builds its per-layer scheme lazily in `create_weights`, so the layer has
|
||||||
|
to declare `scheme = None` for the `is None` probe to see it.
|
||||||
|
|
||||||
|
Regression: the probe used to be `hasattr(layer, "scheme")`, which degraded to
|
||||||
|
always-true once `LinearBase` grew that class default -- the scheme was never
|
||||||
|
built and every GPTQ model died with ``'NoneType' object has no attribute
|
||||||
|
'create_weights'``.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.linear import LinearBase, ReplicatedLinear
|
||||||
|
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
|
||||||
|
from sglang.srt.layers.quantization.gptq.gptq import GPTQConfig
|
||||||
|
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
from sglang.test.test_utils import CustomTestCase
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=5, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
_GPTQ_CHECKPOINT_CONFIG = {
|
||||||
|
"bits": 4,
|
||||||
|
"group_size": 128,
|
||||||
|
"desc_act": False,
|
||||||
|
"lm_head": False,
|
||||||
|
"dynamic": {},
|
||||||
|
"checkpoint_format": "gptq",
|
||||||
|
"true_sequential": True,
|
||||||
|
"static_groups": False,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class TestGPTQSchemeAttach(CustomTestCase):
|
||||||
|
def test_linear_layer_gets_a_scheme(self):
|
||||||
|
layer = ReplicatedLinear(
|
||||||
|
input_size=256,
|
||||||
|
output_size=128,
|
||||||
|
bias=False,
|
||||||
|
params_dtype=torch.float16,
|
||||||
|
quant_config=GPTQConfig.from_config(_GPTQ_CHECKPOINT_CONFIG),
|
||||||
|
prefix="model.layers.0.mlp.down_proj",
|
||||||
|
)
|
||||||
|
self.assertIsNotNone(layer.scheme)
|
||||||
|
self.assertTrue(hasattr(layer, "qweight"))
|
||||||
|
|
||||||
|
def test_scheme_default_is_declared_on_every_quantizable_layer_base(self):
|
||||||
|
"""`get_linear_quant_method` hands a linear method a `LinearBase` or a
|
||||||
|
quantized `ParallelLMHead`; `GPTQMarlinConfig` hands
|
||||||
|
`GPTQMarlinMoEMethod` a bare `FusedMoE`. The MoE attach has no e2e
|
||||||
|
coverage, so this is its only guard.
|
||||||
|
"""
|
||||||
|
self.assertIsNone(LinearBase.scheme)
|
||||||
|
self.assertIsNone(VocabParallelEmbedding.scheme)
|
||||||
|
self.assertIsNone(FusedMoE.scheme)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user