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:
Harry Mellor
2026-07-09 00:14:33 +01:00
committed by GitHub
parent 26831949b4
commit 56da398dac
+26 -24
View File
@@ -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]