[Refactor] simplify multimodal data processing (#8107)
Signed-off-by: Xinyuan Tong <justinning0323@outlook.com>
This commit is contained in:
@@ -260,7 +260,7 @@ class DeepseekVL2ForCausalLM(nn.Module):
|
||||
def get_image_feature(self, items: List[MultimodalDataItem]):
|
||||
|
||||
images_spatial_crop = torch.cat(
|
||||
[item.image_spatial_crop for item in items], dim=0
|
||||
[item.images_spatial_crop for item in items], dim=0
|
||||
)
|
||||
|
||||
assert images_spatial_crop.dim() == 3
|
||||
@@ -278,8 +278,8 @@ class DeepseekVL2ForCausalLM(nn.Module):
|
||||
_, hw, n_dim = images_embeds.shape
|
||||
h = w = int(hw**0.5)
|
||||
tile_index = 0
|
||||
for jdx in range(item.image_spatial_crop.shape[1]):
|
||||
num_width_tiles, num_height_tiles = item.image_spatial_crop[0, jdx]
|
||||
for jdx in range(item.images_spatial_crop.shape[1]):
|
||||
num_width_tiles, num_height_tiles = item.images_spatial_crop[0, jdx]
|
||||
if num_width_tiles == 0 or num_height_tiles == 0:
|
||||
break
|
||||
num_tiles_in_image = num_width_tiles * num_height_tiles
|
||||
|
||||
@@ -81,6 +81,7 @@ class Llama4ForConditionalGeneration(nn.Module):
|
||||
self.logits_processor = LogitsProcessor(
|
||||
config.text_config if hasattr(config, "text_config") else config
|
||||
)
|
||||
self.padding_pattern = MultiModalityDataPaddingPatternMultimodalTokens()
|
||||
|
||||
def _has_vision_weights(self, config) -> bool:
|
||||
"""Check if the model has vision components by examining the checkpoint."""
|
||||
@@ -135,8 +136,7 @@ class Llama4ForConditionalGeneration(nn.Module):
|
||||
return False
|
||||
|
||||
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
||||
pattern = MultiModalityDataPaddingPatternMultimodalTokens()
|
||||
return pattern.pad_input_tokens(input_ids, mm_inputs)
|
||||
return self.padding_pattern.pad_input_tokens(input_ids, mm_inputs)
|
||||
|
||||
def get_image_feature(
|
||||
self,
|
||||
|
||||
@@ -435,7 +435,12 @@ class Phi4MMForCausalLM(nn.Module):
|
||||
dtype = next(self.vision_encoder.parameters()).dtype
|
||||
pixel_values = torch.cat([item.feature for item in items], dim=0).type(dtype)
|
||||
image_attention_mask = torch.cat(
|
||||
[item.image_attention_mask for item in items], dim=0
|
||||
[
|
||||
item.image_attention_mask
|
||||
for item in items
|
||||
if hasattr(item, "image_attention_mask")
|
||||
],
|
||||
dim=0,
|
||||
)
|
||||
image_sizes = torch.cat([item.image_sizes for item in items], dim=0)
|
||||
image_embeds = self.vision_encoder(
|
||||
@@ -456,7 +461,7 @@ class Phi4MMForCausalLM(nn.Module):
|
||||
audio_features=item.feature.to(device).type(dtype),
|
||||
audio_attention_mask=(
|
||||
item.audio_attention_mask.to(device)
|
||||
if item.audio_attention_mask is not None
|
||||
if hasattr(item, "audio_attention_mask")
|
||||
else None
|
||||
),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user