[MM][style] rename inputs_embeds to input_embeds for consistency (#14240)
This commit is contained in:
@@ -598,20 +598,20 @@ def embed_mm_inputs(
|
|||||||
# filled with the hash values of the multimodal for the prefix matching in the radix attention.
|
# filled with the hash values of the multimodal for the prefix matching in the radix attention.
|
||||||
# There values are useless because their embeddings will be replaced by vision embeddings anyway.
|
# There values are useless because their embeddings will be replaced by vision embeddings anyway.
|
||||||
input_ids.clamp_(min=0, max=vocab_size - 1)
|
input_ids.clamp_(min=0, max=vocab_size - 1)
|
||||||
inputs_embeds = input_embedding(input_ids)
|
input_embeds = input_embedding(input_ids)
|
||||||
|
|
||||||
# deepstack embedding
|
# deepstack embedding
|
||||||
if use_deepstack:
|
if use_deepstack:
|
||||||
num_deepstack_embeddings = len(multimodal_model.deepstack_visual_indexes)
|
num_deepstack_embeddings = len(multimodal_model.deepstack_visual_indexes)
|
||||||
|
|
||||||
deepstack_embedding_shape = inputs_embeds.shape[:-1] + (
|
deepstack_embedding_shape = input_embeds.shape[:-1] + (
|
||||||
inputs_embeds.shape[-1] * num_deepstack_embeddings,
|
input_embeds.shape[-1] * num_deepstack_embeddings,
|
||||||
)
|
)
|
||||||
# a zero-filled embedding, with the same length of inputs_embeds, but different hidden_size
|
# a zero-filled embedding, with the same length of input_embeds, but different hidden_size
|
||||||
input_deepstack_embeds = torch.zeros(
|
input_deepstack_embeds = torch.zeros(
|
||||||
deepstack_embedding_shape,
|
deepstack_embedding_shape,
|
||||||
device=inputs_embeds.device,
|
device=input_embeds.device,
|
||||||
dtype=inputs_embeds.dtype,
|
dtype=input_embeds.dtype,
|
||||||
)
|
)
|
||||||
|
|
||||||
other_info["input_deepstack_embeds"] = input_deepstack_embeds
|
other_info["input_deepstack_embeds"] = input_deepstack_embeds
|
||||||
@@ -624,13 +624,13 @@ def embed_mm_inputs(
|
|||||||
continue
|
continue
|
||||||
# in-place update
|
# in-place update
|
||||||
indices = torch.where(mask.squeeze(dim=-1))[0]
|
indices = torch.where(mask.squeeze(dim=-1))[0]
|
||||||
inputs_embeds[indices] = embedding.to(inputs_embeds.device, inputs_embeds.dtype)
|
input_embeds[indices] = embedding.to(input_embeds.device, input_embeds.dtype)
|
||||||
if use_deepstack.get(modality, None):
|
if use_deepstack.get(modality, None):
|
||||||
input_deepstack_embeds[indices] = deepstack_embeddings[i].to(
|
input_deepstack_embeds[indices] = deepstack_embeddings[i].to(
|
||||||
inputs_embeds.device, inputs_embeds.dtype
|
input_embeds.device, input_embeds.dtype
|
||||||
)
|
)
|
||||||
|
|
||||||
return inputs_embeds, other_info
|
return input_embeds, other_info
|
||||||
|
|
||||||
|
|
||||||
def general_mm_embed_routine(
|
def general_mm_embed_routine(
|
||||||
@@ -692,7 +692,7 @@ def general_mm_embed_routine(
|
|||||||
for i, seq_len in enumerate(forward_batch.extend_seq_lens_cpu)
|
for i, seq_len in enumerate(forward_batch.extend_seq_lens_cpu)
|
||||||
if forward_batch.mm_inputs[i] is not None
|
if forward_batch.mm_inputs[i] is not None
|
||||||
]
|
]
|
||||||
inputs_embeds, other_info = embed_mm_inputs(
|
input_embeds, other_info = embed_mm_inputs(
|
||||||
mm_inputs_list=mm_inputs_list,
|
mm_inputs_list=mm_inputs_list,
|
||||||
extend_prefix_lens=extend_prefix_lens,
|
extend_prefix_lens=extend_prefix_lens,
|
||||||
extend_seq_lens=extend_seq_lens,
|
extend_seq_lens=extend_seq_lens,
|
||||||
@@ -712,18 +712,18 @@ def general_mm_embed_routine(
|
|||||||
# just being defensive here
|
# just being defensive here
|
||||||
forward_batch.mm_inputs = None
|
forward_batch.mm_inputs = None
|
||||||
else:
|
else:
|
||||||
inputs_embeds = embed_tokens(input_ids)
|
input_embeds = embed_tokens(input_ids)
|
||||||
# Copy to pre-allocated buffer if available (for CUDA graph address stability)
|
# Copy to pre-allocated buffer if available (for CUDA graph address stability)
|
||||||
if forward_batch.input_embeds is not None:
|
if forward_batch.input_embeds is not None:
|
||||||
forward_batch.input_embeds.copy_(inputs_embeds)
|
forward_batch.input_embeds.copy_(input_embeds)
|
||||||
inputs_embeds = forward_batch.input_embeds
|
input_embeds = forward_batch.input_embeds
|
||||||
else:
|
else:
|
||||||
inputs_embeds = None
|
input_embeds = None
|
||||||
|
|
||||||
hidden_states = language_model(
|
hidden_states = language_model(
|
||||||
input_ids=None,
|
input_ids=None,
|
||||||
forward_batch=forward_batch,
|
forward_batch=forward_batch,
|
||||||
input_embeds=inputs_embeds,
|
input_embeds=input_embeds,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
return hidden_states
|
return hidden_states
|
||||||
|
|||||||
Reference in New Issue
Block a user