[QUANT] Add GPTQModel Dynamic Quantization + lm_head Quantization (#3790)

Signed-off-by: ZX-ModelCloud <zx@modelcloud.ai>
Co-authored-by: ZX-ModelCloud <zx@modelcloud.ai>
This commit is contained in:
Qubitium-ModelCloud
2025-03-05 01:11:00 -08:00
committed by GitHub
co-authored by ZX-ModelCloud
parent 583d6af71b
commit 56a724eba3
56 changed files with 1988 additions and 282 deletions
+57 -11
View File
@@ -36,6 +36,7 @@ from sglang.srt.managers.schedule_batch import ImageInputs
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_loader.weight_utils import default_weight_loader
from sglang.srt.models.llama import LlamaDecoderLayer, LlamaMLP
from sglang.srt.utils import add_prefix
class ColumnParallelConv2dPatch(torch.nn.Module):
@@ -147,7 +148,12 @@ class MllamaPrecomputedPositionEmbedding(nn.Module):
class MllamaVisionMLP(nn.Module):
def __init__(self, config, quant_config: Optional[QuantizationConfig] = None):
def __init__(
self,
config,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.config = config
self.activation_fn = get_act_fn(config.hidden_act)
@@ -156,12 +162,14 @@ class MllamaVisionMLP(nn.Module):
config.intermediate_size,
bias=True,
quant_config=quant_config,
prefix=add_prefix("fc1", prefix),
)
self.fc2 = RowParallelLinear(
config.intermediate_size,
config.hidden_size,
bias=True,
quant_config=quant_config,
prefix=add_prefix("fc2", prefix),
)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
@@ -174,7 +182,10 @@ class MllamaVisionMLP(nn.Module):
class MllamaVisionEncoderLayer(nn.Module):
def __init__(
self, config: config_mllama.MllamaVisionConfig, is_gated: bool = False
self,
config: config_mllama.MllamaVisionConfig,
is_gated: bool = False,
prefix: str = "",
):
super().__init__()
@@ -193,8 +204,9 @@ class MllamaVisionEncoderLayer(nn.Module):
use_context_forward=False,
use_full_precision_softmax=False,
flatten_batch=False,
prefix=add_prefix("self_attn", prefix),
)
self.mlp = MllamaVisionMLP(config)
self.mlp = MllamaVisionMLP(config, prefix=add_prefix("mlp", prefix))
self.input_layernorm = nn.LayerNorm(self.hidden_size, eps=config.norm_eps)
self.post_attention_layernorm = nn.LayerNorm(
@@ -235,11 +247,17 @@ class MllamaVisionEncoder(nn.Module):
num_layers=32,
is_gated=False,
output_hidden_states=None,
prefix: str = "",
):
super().__init__()
self.config = config
self.layers = nn.ModuleList(
[MllamaVisionEncoderLayer(config, is_gated) for _ in range(num_layers)]
[
MllamaVisionEncoderLayer(
config, is_gated, prefix=add_prefix(f"layers.{i}", prefix)
)
for i in range(num_layers)
]
)
self.output_hidden_states = output_hidden_states or []
@@ -265,7 +283,7 @@ class MllamaVisionEncoder(nn.Module):
class MllamaVisionModel(nn.Module):
def __init__(self, config: config_mllama.MllamaVisionConfig):
def __init__(self, config: config_mllama.MllamaVisionConfig, prefix: str = ""):
super().__init__()
self.image_size = config.image_size
self.patch_size = config.patch_size
@@ -305,9 +323,13 @@ class MllamaVisionModel(nn.Module):
config.num_hidden_layers,
is_gated=False,
output_hidden_states=config.intermediate_layers_indices,
prefix=add_prefix("transformer", prefix),
)
self.global_transformer = MllamaVisionEncoder(
config, config.num_global_layers, is_gated=True
config,
config.num_global_layers,
is_gated=True,
prefix=add_prefix("global_transformer", prefix),
)
def apply_class_embedding(self, hidden_state: torch.Tensor) -> torch.Tensor:
@@ -464,6 +486,7 @@ class MllamaTextCrossAttention(nn.Module):
config: Optional[config_mllama.MllamaTextConfig] = None,
layer_id: Optional[int] = None,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.config = config
@@ -489,6 +512,7 @@ class MllamaTextCrossAttention(nn.Module):
self.num_key_value_heads,
bias=False,
quant_config=quant_config,
prefix=add_prefix("qkv_proj", prefix),
)
self.o_proj = RowParallelLinear(
self.num_heads * self.head_dim,
@@ -496,6 +520,7 @@ class MllamaTextCrossAttention(nn.Module):
bias=False,
input_is_parallel=True,
quant_config=quant_config,
prefix=add_prefix("o_proj", prefix),
)
# vllm.model_executor.layers.layernorm.RMSNorm has precision issue,
# use huggingface's instead
@@ -510,6 +535,7 @@ class MllamaTextCrossAttention(nn.Module):
self.num_local_key_value_heads,
layer_id=layer_id,
is_cross_attention=True,
prefix=add_prefix("attn", prefix),
)
def forward(
@@ -551,6 +577,7 @@ class MllamaCrossAttentionDecoderLayer(torch.nn.Module):
config: config_mllama.MllamaTextConfig,
layer_id: int,
quant_config: Optional[QuantizationConfig],
prefix: str = "",
) -> None:
super().__init__()
self.layer_id = layer_id
@@ -558,6 +585,7 @@ class MllamaCrossAttentionDecoderLayer(torch.nn.Module):
config=config,
layer_id=layer_id,
quant_config=quant_config,
prefix=add_prefix("cross_attn", prefix),
)
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
@@ -568,6 +596,7 @@ class MllamaCrossAttentionDecoderLayer(torch.nn.Module):
intermediate_size=config.intermediate_size,
hidden_act=config.hidden_act,
quant_config=quant_config,
prefix=add_prefix("mlp", prefix),
)
self.post_attention_layernorm = RMSNorm(
config.hidden_size, eps=config.rms_norm_eps
@@ -610,12 +639,15 @@ class MllamaTextModel(nn.Module):
self,
config: config_mllama.MllamaTextConfig,
quant_config: Optional[QuantizationConfig],
prefix: str = "",
):
super().__init__()
self.padding_id = config.pad_token_id
self.vocab_size = config.vocab_size
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size + 8, config.hidden_size
config.vocab_size + 8,
config.hidden_size,
prefix=add_prefix("embed_tokens", prefix),
)
self.cross_attention_layers = config.cross_attention_layers
@@ -624,14 +656,20 @@ class MllamaTextModel(nn.Module):
if layer_id in self.cross_attention_layers:
layers.append(
MllamaCrossAttentionDecoderLayer(
config, layer_id, quant_config=quant_config
config,
layer_id,
quant_config=quant_config,
prefix=add_prefix(f"layers.{layer_id}", prefix),
)
)
else:
# TODO: force LlamaDecoderLayer to config.attention_bias=False
layers.append(
LlamaDecoderLayer(
config, quant_config=quant_config, layer_id=layer_id
config,
quant_config=quant_config,
layer_id=layer_id,
prefix=add_prefix(f"layers.{layer_id}", prefix),
)
)
@@ -687,16 +725,20 @@ class MllamaForCausalLM(nn.Module):
self,
config: config_mllama.MllamaTextConfig,
quant_config: Optional[QuantizationConfig],
prefix: str = "",
):
super().__init__()
self.vocab_size = config.vocab_size
self.model = MllamaTextModel(config, quant_config)
self.model = MllamaTextModel(
config, quant_config, prefix=add_prefix("model", prefix)
)
self.lm_head = ParallelLMHead(
config.vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size,
padding_size=DEFAULT_VOCAB_PADDING_SIZE,
quant_config=quant_config,
prefix=add_prefix("lm_head", prefix),
)
def forward(
@@ -726,6 +768,7 @@ class MllamaForConditionalGeneration(nn.Module):
self,
config: config_mllama.MllamaConfig,
quant_config: Optional[QuantizationConfig] = None,
prefix: str = "",
):
super().__init__()
self.vocab_size = config.text_config.vocab_size
@@ -737,10 +780,13 @@ class MllamaForConditionalGeneration(nn.Module):
)
self.image_size = config.vision_config.image_size
self.vision_model = MllamaVisionModel(config.vision_config)
self.vision_model = MllamaVisionModel(
config.vision_config, prefix=add_prefix("vision_model", prefix)
)
self.language_model = MllamaForCausalLM(
config.text_config,
quant_config=quant_config,
prefix=add_prefix("language_model", prefix),
)
self.multi_modal_projector = nn.Linear(
config.vision_config.vision_output_dim,