forked from Karylab-cklius/vllm
Fix embed scaling + CUDA graphs in Transformers modelling backend (#48010)
Signed-off-by: Harry Mellor <19981378+hmellor@users.noreply.github.com>
This commit is contained in:
@@ -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]
|
||||
|
||||
Reference in New Issue
Block a user