model: qwen3-omni (thinker-only) (#10911)
Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
@@ -280,7 +280,6 @@ class MultiModalityDataPaddingPatternMultimodalTokens(MultiModalityDataPaddingPa
|
||||
input_ids_tensor[input_ids_tensor == token_id] = pad_value
|
||||
|
||||
ret_input_ids = input_ids_tensor.tolist()
|
||||
|
||||
return ret_input_ids
|
||||
|
||||
|
||||
@@ -507,7 +506,7 @@ def embed_mm_inputs(
|
||||
Modality, Callable[[List[MultimodalDataItem]], torch.Tensor]
|
||||
] = None,
|
||||
placeholder_tokens: dict[Modality, List[int]] = None,
|
||||
use_deepstack: bool = False,
|
||||
use_deepstack: Dict[Modality, bool] = {},
|
||||
) -> Optional[torch.Tensor]:
|
||||
"""
|
||||
Embed multimodal inputs and integrate them with text token embeddings.
|
||||
@@ -533,7 +532,9 @@ def embed_mm_inputs(
|
||||
for mm_inputs in mm_inputs_list:
|
||||
item_flatten_list += [item for item in mm_inputs.mm_items if item is not None]
|
||||
|
||||
embeddings, masks, deepstack_embeddings = [], [], []
|
||||
# deepstack_embeddings: per-modality
|
||||
modalities, embeddings, masks, deepstack_embeddings = [], [], [], []
|
||||
|
||||
# 2. Get multimodal embedding separately
|
||||
# Try get mm embedding if any
|
||||
for modality in Modality.all():
|
||||
@@ -549,7 +550,8 @@ def embed_mm_inputs(
|
||||
# "image", "video", etc
|
||||
modality_id = modality.name.lower()
|
||||
embedder = getattr(multimodal_model, f"get_{modality_id}_feature", None)
|
||||
if len(items) != 0 and embedder is not None:
|
||||
if len(items) != 0:
|
||||
assert embedder is not None, f"no embedding method found for {modality}"
|
||||
placeholder_tensor = torch.as_tensor(
|
||||
[item.pad_value for item in items],
|
||||
device=input_ids.device,
|
||||
@@ -580,11 +582,12 @@ def embed_mm_inputs(
|
||||
items_offset_list=items_offsets,
|
||||
)
|
||||
|
||||
if use_deepstack and embedding is not None:
|
||||
if use_deepstack.get(modality, None) and embedding is not None:
|
||||
embedding, deepstack_embedding = (
|
||||
multimodal_model.separate_deepstack_embeds(embedding)
|
||||
)
|
||||
deepstack_embeddings += [deepstack_embedding]
|
||||
modalities += [modality]
|
||||
embeddings += [embedding]
|
||||
masks += [mask]
|
||||
|
||||
@@ -597,17 +600,14 @@ def embed_mm_inputs(
|
||||
input_ids.clamp_(min=0, max=vocab_size - 1)
|
||||
inputs_embeds = input_embedding(input_ids)
|
||||
|
||||
# 4. scatter embeddings into input embedding
|
||||
|
||||
# deepstack embedding
|
||||
if use_deepstack:
|
||||
num_deepstack_embeddings = (
|
||||
len(multimodal_model.deepstack_visual_indexes) if use_deepstack else 0
|
||||
)
|
||||
num_deepstack_embeddings = len(multimodal_model.deepstack_visual_indexes)
|
||||
|
||||
deepstack_embedding_shape = inputs_embeds.shape[:-1] + (
|
||||
inputs_embeds.shape[-1] * num_deepstack_embeddings,
|
||||
)
|
||||
|
||||
# a zero-filled embedding, with the same length of inputs_embeds, but different hidden_size
|
||||
input_deepstack_embeds = torch.zeros(
|
||||
deepstack_embedding_shape,
|
||||
device=inputs_embeds.device,
|
||||
@@ -616,14 +616,16 @@ def embed_mm_inputs(
|
||||
|
||||
other_info["input_deepstack_embeds"] = input_deepstack_embeds
|
||||
|
||||
for i, embedding, mask in zip(range(len(embeddings)), embeddings, masks):
|
||||
# 4. scatter embeddings into input embedding
|
||||
for i, modality, embedding, mask in zip(
|
||||
range(len(embeddings)), modalities, embeddings, masks
|
||||
):
|
||||
if embedding is None or mask is None:
|
||||
continue
|
||||
# in-place update
|
||||
indices = torch.where(mask.squeeze(dim=-1))[0]
|
||||
inputs_embeds[indices] = embedding.to(inputs_embeds.device, inputs_embeds.dtype)
|
||||
|
||||
if use_deepstack:
|
||||
if use_deepstack.get(modality, None):
|
||||
input_deepstack_embeds[indices] = deepstack_embeddings[i].to(
|
||||
inputs_embeds.device, inputs_embeds.dtype
|
||||
)
|
||||
@@ -640,7 +642,7 @@ def general_mm_embed_routine(
|
||||
Modality, Callable[[List[MultimodalDataItem]], torch.Tensor]
|
||||
] = None,
|
||||
placeholder_tokens: Optional[dict[Modality, List[int]]] = None,
|
||||
use_deepstack: bool = False,
|
||||
use_deepstack: Dict[Modality, bool] = {},
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
@@ -652,7 +654,7 @@ def general_mm_embed_routine(
|
||||
language_model: Base language model to use
|
||||
data_embedding_funcs: A dictionary mapping from modality type to the corresponding embedding function.
|
||||
placeholder_tokens: Token IDs for multimodal placeholders
|
||||
use_deepstack: Whether to use deepstack embeddings
|
||||
use_deepstack: Whether to use deepstack embeddings for each modality, default False
|
||||
**kwargs: Additional arguments passed to language model
|
||||
|
||||
Returns:
|
||||
|
||||
Reference in New Issue
Block a user