From 57c778aeb470430f477988083c50754aabd08b66 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Thu, 17 Sep 2026 08:19:52 +0800 Subject: [PATCH 1/4] fix(template): truncate token type IDs with retained positions --- swift/template/base.py | 3 ++ tests/general/test_gemma3_template.py | 52 +++++++++++++++++++++++++++ 2 files changed, 55 insertions(+) diff --git a/swift/template/base.py b/swift/template/base.py index d227346032..382bd2c6ec 100644 --- a/swift/template/base.py +++ b/swift/template/base.py @@ -1497,6 +1497,9 @@ def _truncate(self, input_ids: List[int], labels: Optional[List[int]], encoded, loss_scale = torch.tensor(loss_scale)[protected].tolist() loss_scale[0] = 0 encoded['loss_scale'] = loss_scale + token_type_ids = encoded.get('token_type_ids') + if token_type_ids is not None: + encoded['token_type_ids'] = torch.tensor(token_type_ids)[protected].tolist() mm_token_type_ids = encoded.get('mm_token_type_ids') if mm_token_type_ids is not None: encoded['mm_token_type_ids'] = mm_token_type_ids[protected] diff --git a/tests/general/test_gemma3_template.py b/tests/general/test_gemma3_template.py index 0c7469afdc..714a69bcfa 100644 --- a/tests/general/test_gemma3_template.py +++ b/tests/general/test_gemma3_template.py @@ -1,7 +1,9 @@ +import torch import unittest from types import SimpleNamespace from unittest.mock import patch +from swift.template import TEMPLATE_MAPPING, TemplateType from swift.template.templates.gemma import Gemma3Template, Gemma3VisionTemplate @@ -17,3 +19,53 @@ def test_text_only_encode_has_token_type_ids(self): result = template._encode(inputs) self.assertEqual(result['token_type_ids'], [0, 0, 0]) + + def test_truncation_keeps_token_type_ids_aligned(self): + # BOI=99, image=90, EOI=98; text appears before and after each image. + single_image = [10, 11, 12, 13, 99, 90, 90, 98, 14, 15, 16, 17] + two_images = [10, 11, 99, 90, 90, 98, 12, 13, 99, 90, 90, 98, 14, 15] + cases = [ + (single_image, 'left', list(range(4, 12))), + (single_image, 'right', list(range(8))), + # Protected BOI tokens make these retention sets non-contiguous. + (two_images, 'left', [2, 8, 11, 12, 13]), + (two_images, 'right', [0, 1, 2, 3, 8]), + ] + for input_ids, strategy, kept in cases: + with self.subTest(strategy=strategy, kept=kept): + template = Gemma3VisionTemplate( + None, + TEMPLATE_MAPPING[TemplateType.gemma3_vision], + max_length=len(kept), + truncation_strategy=strategy) + template.placeholder_tokens = [99] + template.processor = SimpleNamespace(pad_token_id=0) + template.mode = 'train' + encoded = { + 'input_ids': input_ids, + 'labels': input_ids.copy(), + 'loss_scale': list(range(1, + len(input_ids) + 1)), + 'token_type_ids': [int(token == 90) for token in input_ids], + 'mm_token_type_ids': torch.arange(len(input_ids)), + } + with patch.object(template, '_preprocess_inputs'), patch.object( + template, '_encode', return_value=encoded): + result = template._encode_truncated(None) + + expected_ids = [input_ids[i] for i in kept] + expected = { + 'input_ids': expected_ids, + 'labels': [-100] + expected_ids[1:], + 'loss_scale': [0] + [i + 1 for i in kept[1:]], + 'token_type_ids': [int(token == 90) for token in expected_ids], + 'mm_token_type_ids': kept, + } + self.assertEqual(result['length'], len(kept)) + for key, value in expected.items(): + self.assertEqual(torch.as_tensor(result[key]).tolist(), value) + + batch = template._data_collator([result]) + for key, value in expected.items(): + self.assertEqual(tuple(batch[key].shape), (1, len(kept))) + self.assertEqual(batch[key].tolist(), [value]) From 20b19ad9b5e884965b802b344e8d552caa97a1b4 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Thu, 24 Sep 2026 13:51:39 +0800 Subject: [PATCH 2/4] fix(template): align PaliGemma and Cog multimodal fields Exercise actual template encoding and collation across image, audio, and video inputs. Correct the PaliGemma prompt boundary, extend Cog visual-token loss scales, and preserve mixed CogVLM batches while rejecting unsupported CogAgent cross-attention batches. --- swift/template/templates/gemma.py | 4 +- swift/template/templates/glm.py | 12 +- tests/general/test_token_type_templates.py | 264 +++++++++++++++++++++ 3 files changed, 276 insertions(+), 4 deletions(-) create mode 100644 tests/general/test_token_type_templates.py diff --git a/swift/template/templates/gemma.py b/swift/template/templates/gemma.py index 9ce1203bfc..34666dcb2f 100644 --- a/swift/template/templates/gemma.py +++ b/swift/template/templates/gemma.py @@ -5,7 +5,7 @@ from dataclasses import dataclass, field from typing import Any, Dict, List, Literal, Optional -from swift.utils import get_logger, upper_bound +from swift.utils import get_logger, lower_bound from ..base import Template from ..constant import LLMTemplateType, MLLMTemplateType from ..register import TemplateMeta, register_template @@ -48,7 +48,7 @@ def _encode(self, inputs: StdTemplateInputs) -> Dict[str, Any]: raw_image = inputs.images processor = self.processor if encoded['labels'] is not None: - n = upper_bound(0, len(encoded['labels']), lambda idx: encoded['labels'][idx] == -100) + n = lower_bound(0, len(encoded['labels']), lambda idx: encoded['labels'][idx] != -100) n2 = len(encoded['labels']) - n encoded['token_type_ids'] = [0] * n + [1] * n2 else: diff --git a/swift/template/templates/glm.py b/swift/template/templates/glm.py index 442309a997..b309048480 100644 --- a/swift/template/templates/glm.py +++ b/swift/template/templates/glm.py @@ -577,6 +577,9 @@ def _encode(self, inputs: StdTemplateInputs) -> Dict[str, Any]: encoded['input_ids'] = input_ids[:1] + [self.processor.pad_token_id] * image_token_len + input_ids[1:] if labels is not None: encoded['labels'] = labels[:1] + [-100] * image_token_len + labels[1:] + loss_scale = encoded.get('loss_scale') + if loss_scale is not None: + encoded['loss_scale'] = loss_scale[:1] + [0.] * image_token_len + loss_scale[1:] if len(image) > 0: encoded['images'] = [[img.to(dtype=self.model_info.torch_dtype)] for img in inputs2['images']] if 'cross_images' in inputs2: @@ -586,11 +589,13 @@ def _encode(self, inputs: StdTemplateInputs) -> Dict[str, Any]: return encoded def _data_collator(self, batch: List[Dict[str, Any]], *, padding_to: Optional[int] = None) -> Dict[str, Any]: + if any(b.get('cross_images') for b in batch) and not all(b.get('cross_images') for b in batch): + raise ValueError('CogAgent requires an image for every sample in a multimodal batch.') res = super()._data_collator(batch, padding_to=padding_to) keys = ['images', 'cross_images'] for key in keys: - if key in batch[0]: - res[key] = [b[key][0] for b in batch] + if any(b.get(key) for b in batch): + res[key] = [b[key][0] if b.get(key) else [] for b in batch] return res @@ -648,6 +653,9 @@ def _encode(self, inputs: StdTemplateInputs) -> Dict[str, Any]: encoded['input_ids'] = input_ids[:1] + [self.processor.pad_token_id] * video_token_len + input_ids[1:] if labels is not None: encoded['labels'] = labels[:1] + [-100] * video_token_len + labels[1:] + loss_scale = encoded.get('loss_scale') + if loss_scale is not None: + encoded['loss_scale'] = loss_scale[:1] + [0.] * video_token_len + loss_scale[1:] if len(video) > 0: dtype = model.dtype encoded['images'] = [[img.to(dtype=dtype)] for img in inputs2['images']] diff --git a/tests/general/test_token_type_templates.py b/tests/general/test_token_type_templates.py new file mode 100644 index 0000000000..6c85be04a0 --- /dev/null +++ b/tests/general/test_token_type_templates.py @@ -0,0 +1,264 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import numpy as np +import pytest +import torch +from copy import deepcopy +from types import SimpleNamespace +from unittest.mock import patch + +from swift.template import TEMPLATE_MAPPING +from swift.template.base import Template + + +class _Tokenizer: + pad_token_id = 0 + tokens = { + '': 90, + '': 99, + '': 98, + '<|image|>': 99, + '<|video|>': 97, + '': 90, + '': 95, + '': 96, + '<|reserved_special_token_0|>': 0, + } + sequences = { + 'image': [99, 90, 90, 89], + '\n\nimage': [99, 90, 90, 89], + 'audio': [98, 91, 91, 88], + 'video': [95, 90, 90, 96], + } + + def convert_tokens_to_ids(self, token): + return self.tokens[token] + + def encode(self, text, **kwargs): + return self.sequences[text].copy() + + def __call__(self, text, **kwargs): + return {'input_ids': self.encode(text)} + + +class _Processor: + pad_token_id = 0 + image_token_id = 90 + audio_token_id = 91 + image_token_ids = [90] + full_image_sequence = '\n\nimage' + full_audio_sequence = 'audio' + + def __init__(self, template_type): + self.tokenizer = _Tokenizer() + self.template_type = template_type + + def __call__(self, **kwargs): + return {'pixel_values': torch.ones(1, 3, 2, 2)} + + def image_processor(self, *args, **kwargs): + result = {'pixel_values': np.ones((1, 3, 2, 2), dtype=np.float32)} + if self.template_type == 'molmo2': + result.update( + pixel_values=torch.ones(1, 3, 2, 2), + image_grids=torch.tensor([[1, 2, 2]]), + image_token_pooling=torch.tensor([[0, 1]]), + image_num_crops=torch.tensor([1]), + ) + else: + result['num_crops'] = [1] + return result + + def feature_extractor(self, *args, **kwargs): + return { + 'input_features': np.ones((1, 3, 4), dtype=np.float32), + 'input_features_mask': np.ones((1, 3), dtype=np.bool_), + } + + def video_processor(self, **kwargs): + return { + 'pixel_values_videos': torch.ones(1, 3, 2, 2), + 'video_grids': torch.tensor([[1, 2, 2]]), + 'video_token_pooling': torch.tensor([[0, 1]]), + 'video_metadata': [SimpleNamespace(timestamps=[0.0])], + } + + def get_image_tokens(self, image_grid): + return ['image'] + + def get_video_string(self, video_grid, timestamps): + return 'video' + + +class _CogModel: + dtype = torch.float32 + + def __init__(self, cross_images): + self.cross_images = cross_images + + def build_conversation_input_ids(self, processor, *, images, **kwargs): + result = {'token_type_ids': torch.tensor([0, 1, 1, 0] if images else [0, 0])} + if images: + result['images'] = [torch.ones(3, 2, 2)] + if self.cross_images: + result['cross_images'] = [torch.ones(3, 4, 4)] + return result + + +def _inputs(media): + return SimpleNamespace( + images=['image'] if 'image' in media else [], + audios=[np.zeros(16)] if 'audio' in media else [], + videos=['video'] if 'video' in media else [], + messages=[{ + 'role': 'user', + 'content': 'query' + }], + to_history=lambda: { + 'query': 'query', + 'history': [] + }, + ) + + +@pytest.mark.parametrize('template_type,media', [ + ('paligemma', 'image'), + ('gemma3_vision', 'image'), + ('gemma3n', 'image'), + ('gemma3n', 'audio'), + ('gemma3n', 'image_audio'), + ('molmo2', 'image'), + ('molmo2', 'video'), + ('molmo2', 'image_video'), + ('cogvlm', 'image'), + ('cogvlm2', 'image'), + ('cogagent_chat', 'image'), + ('cogagent_vqa', 'image'), + ('cogvlm2_video', 'video'), +]) +@pytest.mark.parametrize('mode', ['train', 'train_with_loss_scale', 'transformers']) +@pytest.mark.parametrize('strategy', ['left', 'right']) +def test_model_token_types_through_encoding_truncation_and_collation(template_type, media, mode, strategy): + meta = TEMPLATE_MAPPING[template_type] + template = meta.template_cls(None, meta, truncation_strategy=strategy) + template.processor = _Processor(template_type) + template.model_info = SimpleNamespace(torch_dtype=torch.float32) + template.mode = 'transformers' if mode == 'transformers' else 'train' + template.placeholder_tokens = template.placeholder_tokens.copy() + template._init_placeholder_tokens() + template.boi_token_id = 99 + template.boa_token_id = 98 + is_cog = template_type.startswith('cog') + template.model = _CogModel(cross_images=template_type.startswith('cogagent')) + + placeholders = [] + expanded = [] + if not is_cog: + if 'image' in media: + placeholders += [90, 90] if template_type == 'paligemma' else [99] + expanded += [90, 90] if template_type == 'paligemma' else [99, 90, 90, 89] + if 'audio' in media: + placeholders += [98] + expanded += [98, 91, 91, 88] + if 'video' in media: + placeholders += [97] + expanded += [95, 90, 90, 96] + base_ids = [10, 11] + placeholders + [12, 13, 14, 15] + full_ids = [10, 0, 0, 11, 12, 13, 14, 15] if is_cog else [10, 11] + expanded + [12, 13, 14, 15] + base = { + 'input_ids': base_ids, + 'labels': [-100] * (len(base_ids) - 4) + [12, 13, 14, 15] if template.is_training else None, + 'loss_scale': [0.] * (len(base_ids) - 4) + [0.5, 1., 1.5, 2.] if mode == 'train_with_loss_scale' else None, + } + if template_type == 'paligemma': + full_types = [0] * (len(full_ids) - 4) + [1] * 4 if template.is_training else [0] * len(full_ids) + elif is_cog: + full_types = [0, 1, 1, 0, 0, 0, 0, 0] + else: + full_types = [3 if token == 91 else int(token == 90) for token in full_ids] + + # Stub text tokenization and media I/O; run each model's own encoder and collator. + with patch.object(template, '_preprocess_inputs'), patch.object( + Template, '_encode', side_effect=lambda inputs: deepcopy(base)), patch( + 'swift.template.templates.glm.load_batch', side_effect=lambda paths, loader: paths): + full = template._encode_truncated(_inputs(media)) + assert full['input_ids'] == full_ids + assert full['token_type_ids'] == full_types + if template.is_training: + assert full['labels'] == [token if token in {12, 13, 14, 15} else -100 for token in full_ids] + if mode == 'train_with_loss_scale': + response_scales = {12: 0.5, 13: 1., 14: 1.5, 15: 2.} + assert full['loss_scale'] == [response_scales.get(token, 0.) for token in full_ids] + + template.max_length = len(full_ids) - 2 + result = template._encode_truncated(_inputs(media)) + removed = {10, 11} if strategy == 'left' else {14, 15} + kept = [i for i, token in enumerate(full_ids) if token not in removed] + expected = {'input_ids': [full_ids[i] for i in kept], 'token_type_ids': [full_types[i] for i in kept]} + for key, first in (('labels', -100), ('loss_scale', 0)): + if full.get(key) is not None: + expected[key] = [full[key][i] for i in kept] + expected[key][0] = first + for key, value in expected.items(): + assert result[key] == value, key + assert result['length'] == len(kept) + + # A shorter text-only row must work before or after the multimodal row. + base = { + 'input_ids': [10, 11, 12], + 'labels': [-100, -100, 12] if template.is_training else None, + 'loss_scale': [0., 0., 0.5] if mode == 'train_with_loss_scale' else None + } + short = template._encode_truncated(_inputs('')) + short_types = [0, 0, 1] if template_type == 'paligemma' and template.is_training else [0, 0, 0] + assert short['token_type_ids'] == short_types + if template_type.startswith('cogagent'): + # CogAgent cross-attention requires an image for every row. + for rows in ([result, short], [short, result]): + with pytest.raises(ValueError, match='CogAgent requires an image for every sample'): + template._data_collator(deepcopy(rows)) + short = template._encode_truncated(_inputs('image')) + assert short['token_type_ids'] == [0, 1, 1, 0, 0] + + for rows in ([result, short], [short, result]): + batch = template._data_collator(deepcopy(rows)) + for key, pad_value in (('input_ids', 0), ('token_type_ids', 0), ('labels', -100), ('loss_scale', 0)): + if rows[0].get(key) is None: + assert key not in batch + continue + expected_rows = [] + for row in rows: + padding = [pad_value] * (len(kept) - len(row[key])) + expected_rows.append(row[key] + padding if template.is_training else padding + row[key]) + assert batch[key].shape == (2, len(kept)) + assert batch[key].tolist() == expected_rows, key + for key in ('pixel_values', 'pixel_values_videos', 'input_features', 'input_features_mask', 'image_grids', + 'video_grids', 'image_token_pooling', 'video_token_pooling', 'image_num_crops'): + if key in result: + torch.testing.assert_close(batch[key], result[key]) + if is_cog: + for key in ('images', 'cross_images'): + if key not in result: + continue + assert len(batch[key]) == 2 + for i, row in enumerate(rows): + assert len(batch[key][i]) == (1 if key in row else 0) + if key in row: + torch.testing.assert_close(batch[key][i][0], row[key][0][0]) + + +@pytest.mark.parametrize('labels,expected_types', [ + (None, [0, 0, 0]), + ([-100, -100, -100], [0, 0, 0]), + ([-100, -100, 12], [0, 0, 1]), + ([-100, 11, 12], [0, 1, 1]), + ([10, 11, 12], [1, 1, 1]), + ([], []), +]) +def test_paligemma_prompt_answer_boundary(labels, expected_types): + meta = TEMPLATE_MAPPING['paligemma'] + template = meta.template_cls(None, meta) + template.processor = _Processor('paligemma') + encoded = {'input_ids': [10, 11, 12][:len(expected_types)], 'labels': labels} + with patch.object(Template, '_encode', return_value=encoded): + result = template._encode(_inputs('')) + assert result['token_type_ids'] == expected_types From 30fdde557d1a5fa8d710a5efa84653089e10acf3 Mon Sep 17 00:00:00 2001 From: taking-lying-flat <1615405@qq.com> Date: Thu, 24 Sep 2026 14:15:26 +0800 Subject: [PATCH 3/4] test(template): cover all registrations and fix SailVL text inputs --- swift/template/templates/seed.py | 5 +- .../test_registered_template_truncation.py | 69 +++++++++++++++++++ tests/general/test_sailvl_template.py | 66 ++++++++++++++++++ tests/general/test_token_type_templates.py | 64 +++++++++++++++++ tests/test_align/test_template/test_video.py | 2 +- tests/test_align/test_template/test_vision.py | 2 +- 6 files changed, 203 insertions(+), 5 deletions(-) create mode 100644 tests/general/test_registered_template_truncation.py create mode 100644 tests/general/test_sailvl_template.py diff --git a/swift/template/templates/seed.py b/swift/template/templates/seed.py index 0babdcc8ed..296c6681e7 100644 --- a/swift/template/templates/seed.py +++ b/swift/template/templates/seed.py @@ -199,11 +199,11 @@ def _encode(self, inputs: StdTemplateInputs) -> Dict[str, Any]: input_ids = encoded['input_ids'] idx_list = findall(input_ids, -100) pixel_values = None + labels = encoded.get('labels') loss_scale = encoded.get('loss_scale', None) images = inputs.images processor = self.processor if images: - labels = encoded.get('labels') image_inputs = processor.image_processor(images) num_patches = image_inputs['num_patches_list'] pixel_values = image_inputs['pixel_values'] @@ -227,10 +227,10 @@ def _post_encode(self, model: nn.Module, inputs: Dict[str, Any]) -> Dict[str, An embedding = model.language_model.get_input_embeddings() device = embedding.weight.device input_ids = inputs['input_ids'] + inputs_embeds = embedding(input_ids).to(device=device) pixel_values = inputs.get('pixel_values') if pixel_values is not None: vit_embeds = model.extract_feature(pixel_values) - inputs_embeds = embedding(input_ids) B, N, C = inputs_embeds.shape inputs_embeds = inputs_embeds.reshape(B * N, C) @@ -242,7 +242,6 @@ def _post_encode(self, model: nn.Module, inputs: Dict[str, Any]) -> Dict[str, An inputs_embeds = inputs_embeds.reshape(B, N, C) elif is_deepspeed_enabled(): - inputs_embeds = embedding(input_ids).to(device=device) dummy_pixel_values = torch.zeros((1, 3, 32, 32), device=device, dtype=inputs_embeds.dtype) vit_embeds = model.extract_feature(dummy_pixel_values).to(device=device) inputs_embeds = inputs_embeds + vit_embeds.mean() * 0. diff --git a/tests/general/test_registered_template_truncation.py b/tests/general/test_registered_template_truncation.py new file mode 100644 index 0000000000..611fb6f293 --- /dev/null +++ b/tests/general/test_registered_template_truncation.py @@ -0,0 +1,69 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +"""Exercise the actual truncation method of every registered template. + +These are encoded-field contract tests, not tokenizer, processor or model tests. +Model-specific encoding and collation are covered separately. +""" +import pytest +import torch +from types import SimpleNamespace + +from swift.template import TEMPLATE_MAPPING +from swift.template.templates.moss import MossVLTemplate + + +@pytest.mark.parametrize('template_type', sorted(TEMPLATE_MAPPING)) +@pytest.mark.parametrize('strategy', ['left', 'right']) +def test_every_registered_template_truncates_aligned_fields(template_type, strategy): + meta = TEMPLATE_MAPPING[template_type] + template = meta.template_cls.__new__(meta.template_cls) + # SailVL reads these processor attributes in its constructor even when the + # processor initialization is deferred for an encoded-field test. + template.processor = SimpleNamespace(num_image_token=2, convert_tokens_to_ids=lambda token: 90) + meta.template_cls.__init__(template, None, meta, max_length=5, truncation_strategy=strategy) + + if isinstance(template, MossVLTemplate): + # MOSS-VL uses contiguous truncation and a cross-attention mask instead + # of token-type fields. Keep its complete vision span in both directions. + tokens = {'<|vision_start|>': 98, '<|image_pad|>': 99, '<|vision_end|>': 97} + template.processor = SimpleNamespace(unk_token_id=-1, convert_tokens_to_ids=tokens.get) + template.max_length = 6 + input_ids = [10, 11, 98, 99, 97, 12, 13, 14] + kept = list(range(2, 8)) if strategy == 'left' else list(range(6)) + mask = torch.arange(16).reshape(1, 1, 8, 2) + encoded = {'cross_attention_mask': mask.clone(), 'loss_scale': None} + actual_ids, actual_labels = template._truncate(input_ids, input_ids.copy(), encoded, strategy) + assert actual_ids == [input_ids[i] for i in kept] + assert actual_labels == [-100] + actual_ids[1:] + torch.testing.assert_close(encoded['cross_attention_mask'], mask[:, :, kept, :]) + return + + template.placeholder_tokens = [99, 90] + input_ids = [10, 99, 11, 90, 90, 12, 98, 13] + kept = [1, 3, 4, 6, 7] if strategy == 'left' else [0, 1, 2, 3, 4] + types = [0, 0, 0, 1, 1, 0, 0, 0] + for token_types in (None, types, torch.tensor(types, dtype=torch.int32), torch.tensor([types])): + encoded = { + 'loss_scale': [0., 0., 0.25, 0., 0., 0.5, 0.75, 1.], + 'token_type_ids': token_types, + 'mm_token_type_ids': torch.arange(8), + 'image_token_types': torch.tensor([-1, -1, -1, 0, 0, -1, -1, -1]), + } + expected_loss = [encoded['loss_scale'][i] for i in kept] + expected_loss[0] = 0 + expected_image_types = encoded['image_token_types'][kept] + actual_ids, actual_labels = template._truncate(input_ids, input_ids.copy(), encoded, strategy) + assert actual_ids == [input_ids[i] for i in kept] + assert actual_labels == [-100] + actual_ids[1:] + assert encoded['loss_scale'] == expected_loss + torch.testing.assert_close(encoded['mm_token_type_ids'], torch.tensor(kept)) + torch.testing.assert_close(encoded['image_token_types'], expected_image_types) + if token_types is None: + assert encoded['token_type_ids'] is None + elif isinstance(token_types, torch.Tensor): + expected = torch.tensor([types[i] for i in kept], dtype=token_types.dtype) + if token_types.ndim == 2: + expected = expected[None] + torch.testing.assert_close(encoded['token_type_ids'], expected) + else: + assert encoded['token_type_ids'] == [types[i] for i in kept] diff --git a/tests/general/test_sailvl_template.py b/tests/general/test_sailvl_template.py new file mode 100644 index 0000000000..95664de48d --- /dev/null +++ b/tests/general/test_sailvl_template.py @@ -0,0 +1,66 @@ +# Copyright (c) ModelScope Contributors. All rights reserved. +import pytest +import torch +from copy import deepcopy +from types import SimpleNamespace +from unittest.mock import patch + +from swift.template import TEMPLATE_MAPPING +from swift.template.base import Template +from swift.template.templates.seed import SailVLTemplate + + +def _template(): + template = SailVLTemplate.__new__(SailVLTemplate) + template.processor = SimpleNamespace(num_image_token=2, convert_tokens_to_ids=lambda token: 90) + template.__init__(None, TEMPLATE_MAPPING['sail_vl2']) + return template + + +@pytest.mark.parametrize('has_image', [False, True]) +@pytest.mark.parametrize('mode', ['train', 'train_with_loss_scale', 'transformers']) +def test_sailvl_text_and_image_encoding(has_image, mode): + template = _template() + base_ids = [10, -100, 11, 12] if has_image else [10, 11, 12] + base = { + 'input_ids': base_ids, + 'labels': [-100] * (len(base_ids) - 1) + [12] if mode != 'transformers' else None, + 'loss_scale': [0.] * (len(base_ids) - 1) + [0.5] if mode == 'train_with_loss_scale' else None, + } + pixels = torch.ones(2, 3, 2, 2) + template.processor.image_processor = lambda images: {'num_patches_list': [2], 'pixel_values': pixels} + template.processor.encode = lambda *args, **kwargs: [90] + with patch.object(Template, '_encode', return_value=deepcopy(base)): + encoded = template._encode(SimpleNamespace(images=['image'] if has_image else [])) + + expected = [10, 90, 90, 90, 90, 11, 12] if has_image else base_ids + assert encoded['input_ids'] == expected + assert encoded['labels'] == ([-100] * (len(expected) - 1) + [12] if mode != 'transformers' else None) + assert encoded['loss_scale'] == ([0.] * (len(expected) - 1) + [0.5] if mode == 'train_with_loss_scale' else None) + assert encoded['pixel_values'] is (pixels if has_image else None) + + +@pytest.mark.parametrize('has_image', [False, True]) +@pytest.mark.parametrize('deepspeed_enabled', [False, True]) +def test_sailvl_post_encode_preserves_text_and_connects_gradients(has_image, deepspeed_enabled): + template = _template() + embedding = torch.nn.Embedding(100, 4) + vision = torch.nn.Linear(1, 4) + model = SimpleNamespace( + language_model=SimpleNamespace(get_input_embeddings=lambda: embedding), + extract_feature=lambda pixels: vision(pixels.mean().reshape(1, 1)), + ) + input_ids = torch.tensor([[10, 90, 12] if has_image else [10, 11, 12]]) + pixels = torch.ones(1, 3, 2, 2) if has_image else None + expected = embedding(input_ids).detach().clone() + if has_image: + expected[:, 1] = model.extract_feature(pixels).detach() + with patch('swift.template.templates.seed.is_deepspeed_enabled', return_value=deepspeed_enabled): + encoded = template._post_encode(model, {'input_ids': input_ids, 'pixel_values': pixels}) + torch.testing.assert_close(encoded['inputs_embeds'], expected) + encoded['inputs_embeds'].sum().backward() + assert embedding.weight.grad is not None + if has_image or deepspeed_enabled: + assert vision.weight.grad is not None + else: + assert vision.weight.grad is None diff --git a/tests/general/test_token_type_templates.py b/tests/general/test_token_type_templates.py index 6c85be04a0..e53c6a38f2 100644 --- a/tests/general/test_token_type_templates.py +++ b/tests/general/test_token_type_templates.py @@ -10,6 +10,70 @@ from swift.template.base import Template +@pytest.mark.parametrize('strategy', ['left', 'right']) +@pytest.mark.parametrize('padding_side', ['left', 'right']) +@pytest.mark.parametrize('training', [False, True]) +def test_paligemma_truncated_batch_forward_and_backward(strategy, padding_side, training): + from transformers import PaliGemmaConfig, PaliGemmaForConditionalGeneration + + # A small random model exercises the real attention mask, vision projection + # and loss without downloading pretrained weights. Tokenization is stubbed. + config = PaliGemmaConfig( + vision_config=dict( + model_type='siglip_vision_model', + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=2, + image_size=2, + patch_size=1), + text_config=dict( + model_type='gemma', + vocab_size=100, + hidden_size=16, + intermediate_size=32, + num_hidden_layers=1, + num_attention_heads=2, + num_key_value_heads=1, + head_dim=8, + max_position_embeddings=64, + pad_token_id=0), + image_token_index=90, + vocab_size=100, + projection_dim=16, + hidden_size=16, + ) + model = PaliGemmaForConditionalGeneration(config) + model.train(training) + meta = TEMPLATE_MAPPING['paligemma'] + template = meta.template_cls(None, meta, max_length=8, truncation_strategy=strategy, padding_side=padding_side) + template.processor = _Processor('paligemma') + template.model_info = SimpleNamespace(torch_dtype=torch.float32) + template.mode = 'train' if training else 'transformers' + template.placeholder_tokens = [90] + batch = [] + for has_image in [True, False]: + ids = [10, 11, 90, 90, 90, 90, 12, 13, 14, 15] if has_image else [10, 12, 13] + labels = [-100] * 6 + [12, 13, 14, 15] if has_image else [-100, 12, 13] + base = {'input_ids': ids, 'labels': labels if training else None, 'loss_scale': None} + with patch.object(template, '_preprocess_inputs'), patch.object(Template, '_encode', return_value=base): + batch.append(template._encode_truncated(_inputs('image' if has_image else ''))) + collated = template._data_collator(batch) + output = model(**collated, use_cache=False) + assert output.logits.shape == (2, 8, 100) + assert torch.isfinite(output.logits).all() + if training: + assert torch.isfinite(output.loss) + output.loss.backward() + projector_parameters = [ + parameter for name, parameter in model.named_parameters() if 'multi_modal_projector' in name + ] + assert projector_parameters + for parameter in projector_parameters: + assert parameter.grad is not None + assert torch.isfinite(parameter.grad).all() + + class _Tokenizer: pad_token_id = 0 tokens = { diff --git a/tests/test_align/test_template/test_video.py b/tests/test_align/test_template/test_video.py index 282d42f0e6..3b6ebdadc2 100644 --- a/tests/test_align/test_template/test_video.py +++ b/tests/test_align/test_template/test_video.py @@ -55,7 +55,7 @@ def test_internvl2_5_mpo(): def test_xcomposer2_5(): - engine = TransformersEngine('Shanghai_AI_Laboratory/internlm-xcomposer2d5-ol-7b:base', torch.float16) + engine = TransformersEngine('Shanghai_AI_Laboratory/internlm-xcomposer2d5-ol-7b:base', torch_dtype=torch.float16) messages = [{'role': 'user', 'content': '