[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:
co-authored by
ZX-ModelCloud
parent
583d6af71b
commit
56a724eba3
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user