diff --git a/python/sglang/srt/layers/linear.py b/python/sglang/srt/layers/linear.py index 66a1f00c5..39919eedd 100644 --- a/python/sglang/srt/layers/linear.py +++ b/python/sglang/srt/layers/linear.py @@ -31,7 +31,6 @@ from sglang.srt.layers.parameter import ( RowvLLMParameter, _ColumnvLLMParameter, ) -from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.utils import pad_or_narrow_weight from sglang.srt.utils import get_bool_env_var, is_cpu, is_hip, is_npu, set_weight_attrs @@ -165,6 +164,8 @@ class LinearBase(torch.nn.Module): self.params_dtype = params_dtype self.quant_config = quant_config if quant_config is None: + from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod + self.quant_method: Optional[QuantizeMethodBase] = UnquantizedLinearMethod() else: self.quant_method = quant_config.get_quant_method(self, prefix=prefix) diff --git a/python/sglang/srt/layers/quantization/auto_round.py b/python/sglang/srt/layers/quantization/auto_round.py index 30a0f8bb4..74c4f0231 100644 --- a/python/sglang/srt/layers/quantization/auto_round.py +++ b/python/sglang/srt/layers/quantization/auto_round.py @@ -13,8 +13,6 @@ from sglang.srt.layers.quantization.utils import get_scalar_types ScalarType, scalar_types = get_scalar_types() - -from sglang.srt.layers.linear import LinearBase, UnquantizedLinearMethod from sglang.srt.layers.quantization.base_config import QuantizationConfig @@ -217,11 +215,13 @@ class AutoRoundConfig(QuantizationConfig): return weight_bits < 16 def apply_awq_quant_layer(self, layer, prefix: str, backend: str = "auto"): + from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.quantization.marlin_utils import ( check_marlin_supported, check_moe_marlin_supports_layer, ) + from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead weight_bits, group_size, sym = self.get_layer_config(layer, prefix) @@ -299,11 +299,13 @@ class AutoRoundConfig(QuantizationConfig): return None def apply_gptq_quant_layer(self, layer, prefix: str, backend: str = "auto"): + from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.moe.fused_moe_triton import FusedMoE from sglang.srt.layers.quantization.marlin_utils import ( check_marlin_supported, check_moe_marlin_supports_layer, ) + from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead weight_bits, group_size, sym = self.get_layer_config(layer, prefix) diff --git a/python/sglang/srt/layers/quantization/petit.py b/python/sglang/srt/layers/quantization/petit.py index daac52ee2..37b7fbc54 100644 --- a/python/sglang/srt/layers/quantization/petit.py +++ b/python/sglang/srt/layers/quantization/petit.py @@ -8,7 +8,7 @@ import regex as re import torch from torch.nn.parameter import Parameter -from sglang.srt.layers.linear import LinearBase, UnquantizedLinearMethod +from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.parameter import ModelWeightParameter, PerTensorScaleParameter from sglang.srt.layers.quantization.base_config import ( LinearMethodBase, @@ -20,6 +20,7 @@ from sglang.srt.layers.quantization.petit_utils import ( prepare_nvfp4_layer_for_petit, verify_petit_nvfp4_supported, ) +from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.quantization.utils import is_layer_skipped from sglang.srt.utils import is_hip diff --git a/python/sglang/srt/layers/quantization/quark/quark.py b/python/sglang/srt/layers/quantization/quark/quark.py index 783f8ea4b..6abde9b4f 100644 --- a/python/sglang/srt/layers/quantization/quark/quark.py +++ b/python/sglang/srt/layers/quantization/quark/quark.py @@ -6,7 +6,7 @@ from typing import Any, List, Optional, cast import torch -from sglang.srt.layers.linear import LinearBase, UnquantizedLinearMethod +from sglang.srt.layers.linear import LinearBase from sglang.srt.layers.quantization.base_config import ( # noqa: E501 LinearMethodBase, QuantizationConfig, @@ -20,6 +20,7 @@ from sglang.srt.layers.quantization.quark.schemes import ( QuarkW8A8Fp8, ) from sglang.srt.layers.quantization.quark.utils import deep_compare, should_ignore_layer +from sglang.srt.layers.quantization.unquant import UnquantizedLinearMethod from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.utils import get_device_capability diff --git a/python/sglang/srt/layers/quantization/w4afp8.py b/python/sglang/srt/layers/quantization/w4afp8.py index 1dc592957..83869289e 100644 --- a/python/sglang/srt/layers/quantization/w4afp8.py +++ b/python/sglang/srt/layers/quantization/w4afp8.py @@ -7,7 +7,6 @@ import torch from torch.nn import Module from torch.nn.parameter import Parameter -from sglang.srt.layers.linear import UnquantizedLinearMethod from sglang.srt.layers.quantization.base_config import ( FusedMoEMethodBase, QuantizationConfig,