Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 28 additions & 2 deletions src/mcore_bridge/model/mm_gpts/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,24 @@ def _hf_get_inputs_embeds(inputs_embeds, inputs, visual, hf_config):
pixel_values_videos = inputs.get('pixel_values_videos')
image_grid_thw = inputs.get('image_grid_thw')
video_grid_thw = inputs.get('video_grid_thw')
token_types = inputs.get('mm_token_type_ids')
if token_types is not None:
if token_types.shape != input_ids.shape:
raise ValueError('mm_token_type_ids must match input_ids for vision embedding.')
token_types = token_types.to(device=input_ids.device)
torch._assert_async(((token_types == 0) | (token_types == 1) | (token_types == 2)).all(),
'Unsupported vision modality type.')
for modality, name in ((1, 'image_token_id'), (2, 'video_token_id')):
token_id = getattr(hf_config, name, None)
if token_id is None:
torch._assert_async((token_types != modality).all(), 'Unsupported vision modality.')
else:
torch._assert_async(((token_types != modality) | (input_ids == token_id)).all(),
'Vision modality type disagrees with placeholder token id.')
if pixel_values is None:
torch._assert_async((token_types != 1).all(), 'Image placeholders require pixel_values.')
if pixel_values_videos is None:
torch._assert_async((token_types != 2).all(), 'Video placeholders require pixel_values_videos.')
dtype = visual.dtype
vision_config = HuggingFaceVit._get_vision_config(hf_config)
if pixel_values is None and pixel_values_videos is None: # plain-text
Expand Down Expand Up @@ -128,13 +146,21 @@ def _hf_get_inputs_embeds(inputs_embeds, inputs, visual, hf_config):
video_embeds = mixed_embeds[image_tokens:]

if image_embeds is not None:
image_mask = (input_ids == hf_config.image_token_id).unsqueeze(-1).expand_as(inputs_embeds)
# Generated special tokens retain their text embedding. Only
# input placeholders carry a nonzero modality type.
image_positions = input_ids == hf_config.image_token_id if token_types is None else token_types == 1
torch._assert_async(image_positions.sum() == image_embeds.shape[0],
'Image placeholder and embedding counts differ.')
image_mask = image_positions.unsqueeze(-1).expand_as(inputs_embeds)
image_embeds = image_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
image_mask = image_mask.to(inputs_embeds.device)
inputs_embeds = inputs_embeds.masked_scatter(image_mask, image_embeds)

if video_embeds is not None:
video_mask = (input_ids == hf_config.video_token_id).unsqueeze(-1).expand_as(inputs_embeds)
video_positions = input_ids == hf_config.video_token_id if token_types is None else token_types == 2
torch._assert_async(video_positions.sum() == video_embeds.shape[0],
'Video placeholder and embedding counts differ.')
video_mask = video_positions.unsqueeze(-1).expand_as(inputs_embeds)
video_embeds = video_embeds.to(inputs_embeds.device, inputs_embeds.dtype)
video_mask = video_mask.to(inputs_embeds.device)
inputs_embeds = inputs_embeds.masked_scatter(video_mask, video_embeds)
Expand Down
68 changes: 68 additions & 0 deletions tests/test_vision_token_provenance.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,68 @@
# Copyright (c) ModelScope Contributors. All rights reserved.
import pytest
import torch
from types import SimpleNamespace

from mcore_bridge.model.mm_gpts.utils import HuggingFaceVit


class FakeVisual:
dtype = torch.float32

def __call__(self, pixels, grid_thw):
return pixels


@pytest.mark.parametrize('modality', ['image', 'video'])
@pytest.mark.parametrize('explicit_types', [False, True])
def test_input_vision_embedding_preserves_generated_special_token_gradient(modality, explicit_types):
config = SimpleNamespace(image_token_id=10, video_token_id=11, vision_config=SimpleNamespace(spatial_merge_size=1))
special = getattr(config, f'{modality}_token_id')
ids = torch.tensor([[1, special, special if explicit_types else 2]])
embeddings = torch.randn(1, 3, 4, requires_grad=True)
features = torch.randn(1, 4, requires_grad=True)
inputs = {'input_ids': ids, f'{modality}_grid_thw': torch.tensor([[1, 1, 1]])}
inputs['pixel_values' if modality == 'image' else 'pixel_values_videos'] = features
if explicit_types:
inputs['mm_token_type_ids'] = torch.tensor([[0, 1 if modality == 'image' else 2, 0]])
result = HuggingFaceVit._hf_get_inputs_embeds(embeddings, inputs, FakeVisual(), config)
torch.testing.assert_close(result[0, 1], features[0], atol=0, rtol=0)
torch.testing.assert_close(result[0, 2], embeddings[0, 2], atol=0, rtol=0)
result.sum().backward()
torch.testing.assert_close(embeddings.grad, torch.tensor([[[1.] * 4, [0.] * 4, [1.] * 4]]), atol=0, rtol=0)
torch.testing.assert_close(features.grad, torch.ones_like(features), atol=0, rtol=0)


@pytest.mark.parametrize('failure', ['shape', 'token_id', 'count', 'modality', 'missing_pixels'])
def test_invalid_vision_types_are_rejected(failure):
config = SimpleNamespace(image_token_id=10, video_token_id=11, vision_config=SimpleNamespace(spatial_merge_size=1))
inputs = {
'input_ids': torch.tensor([[1, 10, 2]]),
'mm_token_type_ids': torch.tensor([[0, 1, 0]]),
'pixel_values': torch.ones(1, 4),
'image_grid_thw': torch.tensor([[1, 1, 1]]),
}
if failure == 'shape':
inputs['mm_token_type_ids'] = torch.zeros(1, 2)
elif failure == 'token_id':
inputs['mm_token_type_ids'][0, 0] = 1
elif failure == 'count':
inputs['pixel_values'] = torch.ones(2, 4)
elif failure == 'missing_pixels':
inputs.pop('pixel_values')
else:
inputs['mm_token_type_ids'][0, 0] = 3
with pytest.raises((ValueError, RuntimeError)):
HuggingFaceVit._hf_get_inputs_embeds(torch.ones(1, 3, 4), inputs, FakeVisual(), config)


def test_image_only_config_does_not_require_video_token_id():
config = SimpleNamespace(image_token_id=10, vision_config=SimpleNamespace(spatial_merge_size=1))
inputs = {
'input_ids': torch.tensor([[10, 10]]),
'mm_token_type_ids': torch.tensor([[1, 0]]),
'pixel_values': torch.ones(1, 4),
'image_grid_thw': torch.tensor([[1, 1, 1]]),
}
result = HuggingFaceVit._hf_get_inputs_embeds(torch.zeros(1, 2, 4), inputs, FakeVisual(), config)
torch.testing.assert_close(result, torch.tensor([[[1.] * 4, [0.] * 4]]), atol=0, rtol=0)
Loading