[Refactor] simplify multimodal data processing (#8107)

Signed-off-by: Xinyuan Tong <justinning0323@outlook.com>
This commit is contained in:
Xinyuan Tong
2025-07-20 21:43:09 -07:00
committed by GitHub
parent c9e8613c97
commit 8430bfe3e9
30 changed files with 297 additions and 421 deletions
+3 -3
View File
@@ -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
+2 -2
View File
@@ -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,
+7 -2
View File
@@ -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
),
)