From 557aaf93b16d083fdec4f82a8251d47d47c76ccb Mon Sep 17 00:00:00 2001 From: root Date: Wed, 23 Sep 2026 11:15:01 +0800 Subject: [PATCH] fix: preserve generated vision token text embeddings --- src/mcore_bridge/model/mm_gpts/utils.py | 30 ++++++++++- tests/test_vision_token_provenance.py | 68 +++++++++++++++++++++++++ 2 files changed, 96 insertions(+), 2 deletions(-) create mode 100644 tests/test_vision_token_provenance.py diff --git a/src/mcore_bridge/model/mm_gpts/utils.py b/src/mcore_bridge/model/mm_gpts/utils.py index b610500a..28893765 100644 --- a/src/mcore_bridge/model/mm_gpts/utils.py +++ b/src/mcore_bridge/model/mm_gpts/utils.py @@ -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 @@ -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) diff --git a/tests/test_vision_token_provenance.py b/tests/test_vision_token_provenance.py new file mode 100644 index 00000000..6651d2ef --- /dev/null +++ b/tests/test_vision_token_provenance.py @@ -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)