diff --git a/vllm/model_executor/models/transformers/base.py b/vllm/model_executor/models/transformers/base.py index 510abe82ce7..879121145df 100644 --- a/vllm/model_executor/models/transformers/base.py +++ b/vllm/model_executor/models/transformers/base.py @@ -78,6 +78,17 @@ if TYPE_CHECKING: logger = init_logger(__name__) +class ScaledVocabParallelEmbedding(VocabParallelEmbedding): + """`VocabParallelEmbedding` that scales its output.""" + + def __init__(self, *args, embed_scale: float, **kwargs): + super().__init__(*args, **kwargs) + self.embed_scale = embed_scale + + def forward(self, input_: torch.Tensor) -> torch.Tensor: + return super().forward(input_) * self.embed_scale + + class Base( nn.Module, VllmModel, @@ -156,22 +167,26 @@ class Base( self.attention_instances = self.create_attention_instances() # Input embeddings - self.embed_scale = None input_embeddings = self.model.get_input_embeddings() if not isinstance(input_embeddings, PPMissingLayer): - # Some models scale embeddings inside the input embedding layer - self.embed_scale = getattr(input_embeddings, "embed_scale", None) names = ("embedding_size", "hidden_size") embedding_dim = getattr_iter(self.text_config, names, None) assert embedding_dim is not None - self.model.set_input_embeddings( - VocabParallelEmbedding( - self.text_config.vocab_size, - embedding_dim=embedding_dim, - org_num_embeddings=self.text_config.vocab_size, - quant_config=self.quant_config, - ) + embedding_kwargs = dict( + num_embeddings=self.text_config.vocab_size, + embedding_dim=embedding_dim, + org_num_embeddings=self.text_config.vocab_size, + quant_config=self.quant_config, ) + embed_scale = getattr(input_embeddings, "embed_scale", None) + if embed_scale is not None: + # Some models scale embeddings inside the input embedding layer + new_input_embeddings = ScaledVocabParallelEmbedding( + **embedding_kwargs, embed_scale=float(embed_scale) + ) + else: + new_input_embeddings = VocabParallelEmbedding(**embedding_kwargs) + self.model.set_input_embeddings(new_input_embeddings) # Initialize any parameters that have not had their modules replaced self.init_parameters(self.model) @@ -590,10 +605,7 @@ class Base( _init_parameters(module, dtype) def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: - inputs_embeds = self.model.get_input_embeddings()(input_ids) - if self.embed_scale is not None: - inputs_embeds *= self.embed_scale - return inputs_embeds + return self.model.get_input_embeddings()(input_ids) def forward( self, @@ -608,16 +620,6 @@ class Base( input_ids = None inputs_embeds = intermediate_tensors["hidden_states"] - # If the model scales embeddings inside the input embedding layer we must - # ensure they are scaled here since VocabParallelEmbedding will not do it - if ( - self.embed_scale is not None - and input_ids is not None - and inputs_embeds is None - ): - inputs_embeds = self.embed_input_ids(input_ids) - input_ids = None - # Add batch dimension before entering Transformers model if input_ids is not None and input_ids.ndim == 1: # [seq_len] -> [1, seq_len]