[FIX]Fix Step3-VL multi-image embedding and local patch splitting (#24970)
Co-authored-by: kousakawang <wanghanpei@bytedance.com>
This commit is contained in:
co-authored by
kousakawang
parent
873196f7fa
commit
8aaca72c21
@@ -1356,6 +1356,84 @@ def _slice_model_data(
|
|||||||
return sliced
|
return sliced
|
||||||
|
|
||||||
|
|
||||||
|
def _compute_patch_slices(model_specific_data: dict, num_items: int) -> tuple:
|
||||||
|
"""Compute per-item patch slice boundaries from 'num_patches' metadata.
|
||||||
|
|
||||||
|
Returns (patch_slices, total_num_patches) where patch_slices is a list of
|
||||||
|
(start, end) tuples for each item, or (None, None) if not applicable.
|
||||||
|
This function can be replaced or extended by model-specific plugins that
|
||||||
|
need custom patch-level splitting logic.
|
||||||
|
"""
|
||||||
|
num_patches = model_specific_data.get("num_patches")
|
||||||
|
if _get_length(num_patches) != num_items:
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
if isinstance(num_patches, torch.Tensor):
|
||||||
|
patch_counts = [int(x) for x in num_patches.flatten().cpu().tolist()]
|
||||||
|
elif isinstance(num_patches, np.ndarray):
|
||||||
|
patch_counts = [int(x) for x in num_patches.reshape(-1).tolist()]
|
||||||
|
else:
|
||||||
|
patch_counts = [
|
||||||
|
int(x.item()) if isinstance(x, torch.Tensor) else int(x)
|
||||||
|
for x in num_patches
|
||||||
|
]
|
||||||
|
|
||||||
|
if not all(count >= 0 for count in patch_counts):
|
||||||
|
return None, None
|
||||||
|
|
||||||
|
patch_slices = []
|
||||||
|
patch_start = 0
|
||||||
|
for count in patch_counts:
|
||||||
|
patch_end = patch_start + count
|
||||||
|
patch_slices.append((patch_start, patch_end))
|
||||||
|
patch_start = patch_end
|
||||||
|
return patch_slices, patch_start
|
||||||
|
|
||||||
|
|
||||||
|
# Keys whose dim-0 aligns with total patch count rather than num_items.
|
||||||
|
_PATCH_ALIGNED_KEYS = frozenset(("patch_pixel_values", "patch_newline_mask"))
|
||||||
|
|
||||||
|
|
||||||
|
def _split_model_data_for_item(
|
||||||
|
model_specific_data: dict,
|
||||||
|
index: int,
|
||||||
|
num_items: int,
|
||||||
|
patch_slices,
|
||||||
|
total_num_patches,
|
||||||
|
) -> dict:
|
||||||
|
"""Split model_specific_data for a single item during simple-split expansion.
|
||||||
|
|
||||||
|
This function encapsulates the per-item splitting logic for model-specific
|
||||||
|
data fields. It handles three categories:
|
||||||
|
1. Patch-aligned fields (dim-0 == total_num_patches): sliced by patch boundaries.
|
||||||
|
2. Item-aligned fields (dim-0 == num_items): sliced by item index.
|
||||||
|
3. Shared/scalar fields: copied as-is.
|
||||||
|
|
||||||
|
To support additional models, extend `_PATCH_ALIGNED_KEYS` or override this
|
||||||
|
function with a model-specific variant.
|
||||||
|
"""
|
||||||
|
new_data = {}
|
||||||
|
for k, v in model_specific_data.items():
|
||||||
|
if (
|
||||||
|
k in _PATCH_ALIGNED_KEYS
|
||||||
|
and patch_slices is not None
|
||||||
|
and _get_length(v) == total_num_patches
|
||||||
|
):
|
||||||
|
patch_start, patch_end = patch_slices[index]
|
||||||
|
new_data[k] = _slice_value(v, patch_start, patch_end)
|
||||||
|
elif isinstance(v, (list, tuple)) and len(v) == num_items:
|
||||||
|
new_data[k] = [v[index]]
|
||||||
|
elif (
|
||||||
|
isinstance(v, (torch.Tensor, np.ndarray))
|
||||||
|
and len(v.shape) > 0
|
||||||
|
and v.shape[0] == num_items
|
||||||
|
):
|
||||||
|
new_data[k] = v[index : index + 1]
|
||||||
|
else:
|
||||||
|
new_data[k] = v
|
||||||
|
return new_data
|
||||||
|
|
||||||
|
|
||||||
def _try_simple_split(item, num_items, expanded_mm_items):
|
def _try_simple_split(item, num_items, expanded_mm_items):
|
||||||
"""Try to split a bundled item by matching feature dim-0 to offset count.
|
"""Try to split a bundled item by matching feature dim-0 to offset count.
|
||||||
Returns True if split succeeded, False otherwise."""
|
Returns True if split succeeded, False otherwise."""
|
||||||
@@ -1373,6 +1451,10 @@ def _try_simple_split(item, num_items, expanded_mm_items):
|
|||||||
if feature_count != num_items:
|
if feature_count != num_items:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
patch_slices, total_num_patches = _compute_patch_slices(
|
||||||
|
item.model_specific_data, num_items
|
||||||
|
)
|
||||||
|
|
||||||
for i in range(num_items):
|
for i in range(num_items):
|
||||||
new_item = copy.copy(item)
|
new_item = copy.copy(item)
|
||||||
if item.feature is not None:
|
if item.feature is not None:
|
||||||
@@ -1386,19 +1468,9 @@ def _try_simple_split(item, num_items, expanded_mm_items):
|
|||||||
else:
|
else:
|
||||||
new_item.precomputed_embeddings = item.precomputed_embeddings[i : i + 1]
|
new_item.precomputed_embeddings = item.precomputed_embeddings[i : i + 1]
|
||||||
new_item.offsets = [item.offsets[i]]
|
new_item.offsets = [item.offsets[i]]
|
||||||
new_data = {}
|
new_item.model_specific_data = _split_model_data_for_item(
|
||||||
for k, v in item.model_specific_data.items():
|
item.model_specific_data, i, num_items, patch_slices, total_num_patches
|
||||||
if isinstance(v, (list, tuple)) and len(v) == num_items:
|
)
|
||||||
new_data[k] = [v[i]]
|
|
||||||
elif (
|
|
||||||
isinstance(v, (torch.Tensor, np.ndarray))
|
|
||||||
and len(v.shape) > 0
|
|
||||||
and v.shape[0] == num_items
|
|
||||||
):
|
|
||||||
new_data[k] = v[i : i + 1]
|
|
||||||
else:
|
|
||||||
new_data[k] = v
|
|
||||||
new_item.model_specific_data = new_data
|
|
||||||
new_item.hash = None
|
new_item.hash = None
|
||||||
expanded_mm_items.append(new_item)
|
expanded_mm_items.append(new_item)
|
||||||
return True
|
return True
|
||||||
|
|||||||
@@ -812,46 +812,93 @@ class Step3VLForConditionalGeneration(nn.Module):
|
|||||||
return image_features
|
return image_features
|
||||||
|
|
||||||
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
||||||
assert len(items) == 1 # We only have images.
|
# Phase 1: Collect thumbnails and patches separately (different resolutions).
|
||||||
|
all_thumbnails = []
|
||||||
|
all_patches = []
|
||||||
|
# Per-item metadata: (thumb_count, num_patches_list, patch_count)
|
||||||
|
item_metadata = []
|
||||||
|
|
||||||
item = items[0]
|
for item in items:
|
||||||
pixel_values = item.feature.type(self.vision_model.dtype)
|
pixel_values = item.feature.type(self.vision_model.dtype)
|
||||||
num_patches = item.model_specific_data.get("num_patches")
|
num_patches = item.model_specific_data.get("num_patches")
|
||||||
patch_pixel_values = item.model_specific_data.get("patch_pixel_values", None)
|
if num_patches is None:
|
||||||
if patch_pixel_values is not None:
|
raise ValueError("Step3-VL image item is missing num_patches.")
|
||||||
patch_pixel_values = patch_pixel_values.type(self.vision_model.dtype)
|
if isinstance(num_patches, torch.Tensor):
|
||||||
|
num_patches = [int(x) for x in num_patches.flatten().cpu().tolist()]
|
||||||
if patch_pixel_values is not None:
|
elif isinstance(num_patches, (list, tuple)):
|
||||||
patch_pixel_values = patch_pixel_values.to("cuda")
|
num_patches = [
|
||||||
|
int(x.item()) if isinstance(x, torch.Tensor) else int(x)
|
||||||
image_features = self._get_vision_model_output(pixel_values)
|
for x in num_patches
|
||||||
patch_image_features = (
|
|
||||||
self._get_vision_model_output(patch_pixel_values)
|
|
||||||
if patch_pixel_values is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
|
|
||||||
image_features = self._process_image_features(image_features)
|
|
||||||
patch_image_features = (
|
|
||||||
self._process_image_features(patch_image_features)
|
|
||||||
if patch_image_features is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
|
|
||||||
merged_image_features = []
|
|
||||||
cur_patch_idx = 0
|
|
||||||
for i, num_patch in enumerate(num_patches):
|
|
||||||
cur_feature = []
|
|
||||||
if num_patch > 0:
|
|
||||||
patch_slice = patch_image_features[
|
|
||||||
cur_patch_idx : cur_patch_idx + num_patch
|
|
||||||
]
|
]
|
||||||
cur_feature.append(patch_slice.view(-1, patch_slice.shape[-1]))
|
else:
|
||||||
cur_feature.append(image_features[i].view(-1, image_features.shape[-1]))
|
num_patches = [int(num_patches)]
|
||||||
cur_patch_idx += num_patch
|
|
||||||
merged_image_features.append(
|
patch_pixel_values = item.model_specific_data.get(
|
||||||
torch.cat(cur_feature) if len(cur_feature) > 1 else cur_feature[0]
|
"patch_pixel_values", None
|
||||||
)
|
)
|
||||||
|
if patch_pixel_values is not None and patch_pixel_values.shape[0] == 0:
|
||||||
|
patch_pixel_values = None
|
||||||
|
if patch_pixel_values is not None:
|
||||||
|
patch_pixel_values = patch_pixel_values.type(
|
||||||
|
self.vision_model.dtype
|
||||||
|
).to(self.device)
|
||||||
|
|
||||||
|
all_thumbnails.append(pixel_values)
|
||||||
|
thumb_count = pixel_values.shape[0]
|
||||||
|
patch_count = 0
|
||||||
|
if patch_pixel_values is not None:
|
||||||
|
all_patches.append(patch_pixel_values)
|
||||||
|
patch_count = patch_pixel_values.shape[0]
|
||||||
|
item_metadata.append((thumb_count, num_patches, patch_count))
|
||||||
|
|
||||||
|
# Phase 2: Batched ViT + projector forward (one pass per resolution).
|
||||||
|
all_thumbnails = torch.cat(all_thumbnails, dim=0)
|
||||||
|
all_thumb_features = self._process_image_features(
|
||||||
|
self._get_vision_model_output(all_thumbnails)
|
||||||
|
)
|
||||||
|
|
||||||
|
all_patch_features = None
|
||||||
|
if all_patches:
|
||||||
|
all_patches = torch.cat(all_patches, dim=0)
|
||||||
|
all_patch_features = self._process_image_features(
|
||||||
|
self._get_vision_model_output(all_patches)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Phase 3: Split results back and merge per-image features.
|
||||||
|
merged_image_features = []
|
||||||
|
thumb_offset = 0
|
||||||
|
patch_offset = 0
|
||||||
|
for thumb_count, num_patches_list, patch_count in item_metadata:
|
||||||
|
item_thumb_features = all_thumb_features[
|
||||||
|
thumb_offset : thumb_offset + thumb_count
|
||||||
|
]
|
||||||
|
thumb_offset += thumb_count
|
||||||
|
item_patch_features = (
|
||||||
|
all_patch_features[patch_offset : patch_offset + patch_count]
|
||||||
|
if patch_count > 0
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
patch_offset += patch_count
|
||||||
|
|
||||||
|
cur_patch_idx = 0
|
||||||
|
for i, num_patch in enumerate(num_patches_list):
|
||||||
|
cur_feature = []
|
||||||
|
if num_patch > 0:
|
||||||
|
if item_patch_features is None:
|
||||||
|
raise ValueError(
|
||||||
|
"Step3-VL image item has num_patches > 0 but no patch_pixel_values."
|
||||||
|
)
|
||||||
|
patch_slice = item_patch_features[
|
||||||
|
cur_patch_idx : cur_patch_idx + num_patch
|
||||||
|
]
|
||||||
|
cur_feature.append(patch_slice.view(-1, patch_slice.shape[-1]))
|
||||||
|
cur_feature.append(
|
||||||
|
item_thumb_features[i].view(-1, item_thumb_features.shape[-1])
|
||||||
|
)
|
||||||
|
cur_patch_idx += num_patch
|
||||||
|
merged_image_features.append(
|
||||||
|
torch.cat(cur_feature) if len(cur_feature) > 1 else cur_feature[0]
|
||||||
|
)
|
||||||
return self._flatten_embeddings(merged_image_features)
|
return self._flatten_embeddings(merged_image_features)
|
||||||
|
|
||||||
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
||||||
|
|||||||
@@ -481,43 +481,93 @@ class StepVLForConditionalGeneration(nn.Module):
|
|||||||
return image_features
|
return image_features
|
||||||
|
|
||||||
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
def get_image_feature(self, items: List[MultimodalDataItem]) -> torch.Tensor:
|
||||||
assert len(items) == 1 # We only have images.
|
# Phase 1: Collect thumbnails and patches separately (different resolutions).
|
||||||
|
all_thumbnails = []
|
||||||
|
all_patches = []
|
||||||
|
# Per-item metadata: (thumb_count, num_patches_list, patch_count)
|
||||||
|
item_metadata = []
|
||||||
|
|
||||||
item = items[0]
|
for item in items:
|
||||||
pixel_values = item.feature.type(self.vision_model.dtype)
|
pixel_values = item.feature.type(self.vision_model.dtype)
|
||||||
num_patches = item.model_specific_data.get("num_patches")
|
num_patches = item.model_specific_data.get("num_patches")
|
||||||
patch_pixel_values = item.model_specific_data.get("patch_pixel_values", None)
|
if num_patches is None:
|
||||||
if patch_pixel_values is not None:
|
raise ValueError("Step3-VL image item is missing num_patches.")
|
||||||
patch_pixel_values = patch_pixel_values.type(self.vision_model.dtype).to(
|
if isinstance(num_patches, torch.Tensor):
|
||||||
self.device
|
num_patches = [int(x) for x in num_patches.flatten().cpu().tolist()]
|
||||||
)
|
elif isinstance(num_patches, (list, tuple)):
|
||||||
|
num_patches = [
|
||||||
image_features = self._get_vision_model_output(pixel_values)
|
int(x.item()) if isinstance(x, torch.Tensor) else int(x)
|
||||||
patch_image_features = (
|
for x in num_patches
|
||||||
self._get_vision_model_output(patch_pixel_values)
|
|
||||||
if patch_pixel_values is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
image_features = self._process_image_features(image_features)
|
|
||||||
patch_image_features = (
|
|
||||||
self._process_image_features(patch_image_features)
|
|
||||||
if patch_image_features is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
merged_image_features = []
|
|
||||||
cur_patch_idx = 0
|
|
||||||
for i, num_patch in enumerate(num_patches):
|
|
||||||
cur_feature = []
|
|
||||||
if num_patch > 0:
|
|
||||||
patch_slice = patch_image_features[
|
|
||||||
cur_patch_idx : cur_patch_idx + num_patch
|
|
||||||
]
|
]
|
||||||
cur_feature.append(patch_slice.view(-1, patch_slice.shape[-1]))
|
else:
|
||||||
cur_feature.append(image_features[i].view(-1, image_features.shape[-1]))
|
num_patches = [int(num_patches)]
|
||||||
cur_patch_idx += num_patch
|
|
||||||
merged_image_features.append(
|
patch_pixel_values = item.model_specific_data.get(
|
||||||
torch.cat(cur_feature) if len(cur_feature) > 1 else cur_feature[0]
|
"patch_pixel_values", None
|
||||||
)
|
)
|
||||||
|
if patch_pixel_values is not None and patch_pixel_values.shape[0] == 0:
|
||||||
|
patch_pixel_values = None
|
||||||
|
if patch_pixel_values is not None:
|
||||||
|
patch_pixel_values = patch_pixel_values.type(
|
||||||
|
self.vision_model.dtype
|
||||||
|
).to(self.device)
|
||||||
|
|
||||||
|
all_thumbnails.append(pixel_values)
|
||||||
|
thumb_count = pixel_values.shape[0]
|
||||||
|
patch_count = 0
|
||||||
|
if patch_pixel_values is not None:
|
||||||
|
all_patches.append(patch_pixel_values)
|
||||||
|
patch_count = patch_pixel_values.shape[0]
|
||||||
|
item_metadata.append((thumb_count, num_patches, patch_count))
|
||||||
|
|
||||||
|
# Phase 2: Batched ViT + projector forward (one pass per resolution).
|
||||||
|
all_thumbnails = torch.cat(all_thumbnails, dim=0)
|
||||||
|
all_thumb_features = self._process_image_features(
|
||||||
|
self._get_vision_model_output(all_thumbnails)
|
||||||
|
)
|
||||||
|
|
||||||
|
all_patch_features = None
|
||||||
|
if all_patches:
|
||||||
|
all_patches = torch.cat(all_patches, dim=0)
|
||||||
|
all_patch_features = self._process_image_features(
|
||||||
|
self._get_vision_model_output(all_patches)
|
||||||
|
)
|
||||||
|
|
||||||
|
# Phase 3: Split results back and merge per-image features.
|
||||||
|
merged_image_features = []
|
||||||
|
thumb_offset = 0
|
||||||
|
patch_offset = 0
|
||||||
|
for thumb_count, num_patches_list, patch_count in item_metadata:
|
||||||
|
item_thumb_features = all_thumb_features[
|
||||||
|
thumb_offset : thumb_offset + thumb_count
|
||||||
|
]
|
||||||
|
thumb_offset += thumb_count
|
||||||
|
item_patch_features = (
|
||||||
|
all_patch_features[patch_offset : patch_offset + patch_count]
|
||||||
|
if patch_count > 0
|
||||||
|
else None
|
||||||
|
)
|
||||||
|
patch_offset += patch_count
|
||||||
|
|
||||||
|
cur_patch_idx = 0
|
||||||
|
for i, num_patch in enumerate(num_patches_list):
|
||||||
|
cur_feature = []
|
||||||
|
if num_patch > 0:
|
||||||
|
if item_patch_features is None:
|
||||||
|
raise ValueError(
|
||||||
|
"Step3-VL image item has num_patches > 0 but no patch_pixel_values."
|
||||||
|
)
|
||||||
|
patch_slice = item_patch_features[
|
||||||
|
cur_patch_idx : cur_patch_idx + num_patch
|
||||||
|
]
|
||||||
|
cur_feature.append(patch_slice.view(-1, patch_slice.shape[-1]))
|
||||||
|
cur_feature.append(
|
||||||
|
item_thumb_features[i].view(-1, item_thumb_features.shape[-1])
|
||||||
|
)
|
||||||
|
cur_patch_idx += num_patch
|
||||||
|
merged_image_features.append(
|
||||||
|
torch.cat(cur_feature) if len(cur_feature) > 1 else cur_feature[0]
|
||||||
|
)
|
||||||
return self._flatten_embeddings(merged_image_features)
|
return self._flatten_embeddings(merged_image_features)
|
||||||
|
|
||||||
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
def pad_input_ids(self, input_ids: List[int], mm_inputs: MultimodalInputs):
|
||||||
|
|||||||
Reference in New Issue
Block a user