diff --git a/tests/test_distillation_trainer.py b/tests/test_distillation_trainer.py index 98ce8fe36eb..e82cb2f12a2 100644 --- a/tests/test_distillation_trainer.py +++ b/tests/test_distillation_trainer.py @@ -1284,10 +1284,11 @@ def test_liger_rejects_zero_logit_scale(self): @require_vision class TestDistillationTrainerVLM(TrlTestCase): @pytest.mark.parametrize( - "model_id", + "model_id,image_token", [ pytest.param( "trl-internal-testing/tiny-Gemma3ForConditionalGeneration", + "", marks=pytest.mark.skipif( Version(transformers.__version__) < Version("4.57.0"), reason="transformers<4.57 Gemma3 image processor can't batch variable-size images", @@ -1295,17 +1296,19 @@ class TestDistillationTrainerVLM(TrlTestCase): ), pytest.param( "trl-internal-testing/tiny-Gemma4ForConditionalGeneration", + "<|image|>", marks=pytest.mark.skipif( Version(transformers.__version__) < Version("5.5.0"), reason="Gemma4 models were introduced in transformers-5.5.0", ), ), - "trl-internal-testing/tiny-InternVLForConditionalGeneration", - "trl-internal-testing/tiny-LlavaNextForConditionalGeneration", - "trl-internal-testing/tiny-Qwen2_5_VLForConditionalGeneration", - "trl-internal-testing/tiny-Qwen2VLForConditionalGeneration", + ("trl-internal-testing/tiny-InternVLForConditionalGeneration", ""), + ("trl-internal-testing/tiny-LlavaNextForConditionalGeneration", ""), + ("trl-internal-testing/tiny-Qwen2_5_VLForConditionalGeneration", "<|image_pad|>"), + ("trl-internal-testing/tiny-Qwen2VLForConditionalGeneration", "<|image_pad|>"), pytest.param( "trl-internal-testing/tiny-Qwen3_5ForConditionalGeneration-NoThink", + "<|image_pad|>", marks=pytest.mark.skipif( Version(transformers.__version__) < Version("5.2.0"), reason="Qwen3.5 models were introduced in transformers-5.2.0", @@ -1313,6 +1316,7 @@ class TestDistillationTrainerVLM(TrlTestCase): ), pytest.param( "trl-internal-testing/tiny-Qwen3_5MoeForConditionalGeneration-3.6", + "<|image_pad|>", marks=pytest.mark.skipif( Version(transformers.__version__) < Version("5.2.0"), reason="Qwen3.5 models were introduced in transformers-5.2.0", @@ -1321,7 +1325,7 @@ class TestDistillationTrainerVLM(TrlTestCase): # "trl-internal-testing/tiny-SmolVLMForConditionalGeneration", seems not to support bf16 properly ], ) - def test_train_vlm(self, model_id): + def test_train_vlm(self, model_id, image_token): dataset = load_dataset("trl-internal-testing/zen-image", "conversational_prompt_only", split="train") training_args = DistillationConfig( @@ -1336,6 +1340,7 @@ def test_train_vlm(self, model_id): args=training_args, train_dataset=dataset, ) + assert trainer._image_token_id == trainer._tokenizer.convert_tokens_to_ids(image_token) trainer.train() diff --git a/tests/test_grpo_trainer.py b/tests/test_grpo_trainer.py index 834833b3dc3..12f52ce5ff7 100644 --- a/tests/test_grpo_trainer.py +++ b/tests/test_grpo_trainer.py @@ -3893,22 +3893,24 @@ def test_single_reward_model_with_single_processing_class(self): @require_vision class TestGRPOTrainerVLM(TrlTestCase): @pytest.mark.parametrize( - "model_id", + "model_id,image_token", [ - "trl-internal-testing/tiny-Gemma3ForConditionalGeneration", + ("trl-internal-testing/tiny-Gemma3ForConditionalGeneration", ""), pytest.param( "trl-internal-testing/tiny-Gemma4ForConditionalGeneration", + "<|image|>", marks=pytest.mark.skipif( Version(transformers.__version__) < Version("5.5.0"), reason="Gemma4 models were introduced in transformers-5.5.0", ), ), - "trl-internal-testing/tiny-InternVLForConditionalGeneration", - "trl-internal-testing/tiny-LlavaNextForConditionalGeneration", - "trl-internal-testing/tiny-Qwen2_5_VLForConditionalGeneration", - "trl-internal-testing/tiny-Qwen2VLForConditionalGeneration", + ("trl-internal-testing/tiny-InternVLForConditionalGeneration", ""), + ("trl-internal-testing/tiny-LlavaNextForConditionalGeneration", ""), + ("trl-internal-testing/tiny-Qwen2_5_VLForConditionalGeneration", "<|image_pad|>"), + ("trl-internal-testing/tiny-Qwen2VLForConditionalGeneration", "<|image_pad|>"), pytest.param( "trl-internal-testing/tiny-Qwen3_5ForConditionalGeneration-NoThink", + "<|image_pad|>", marks=pytest.mark.skipif( Version(transformers.__version__) < Version("5.2.0"), reason="Qwen3.5 models were introduced in transformers-5.2.0", @@ -3916,6 +3918,7 @@ class TestGRPOTrainerVLM(TrlTestCase): ), pytest.param( "trl-internal-testing/tiny-Qwen3_5MoeForConditionalGeneration-3.6", + "<|image_pad|>", marks=pytest.mark.skipif( Version(transformers.__version__) < Version("5.2.0"), reason="Qwen3.5 models were introduced in transformers-5.2.0", @@ -3924,7 +3927,7 @@ class TestGRPOTrainerVLM(TrlTestCase): # "trl-internal-testing/tiny-SmolVLMForConditionalGeneration", seems not to support bf16 properly ], ) - def test_train_vlm(self, model_id): + def test_train_vlm(self, model_id, image_token): dataset = load_dataset("trl-internal-testing/zen-image", "conversational_prompt_only", split="train") def reward_func(completions, **kwargs): @@ -3946,6 +3949,7 @@ def reward_func(completions, **kwargs): args=training_args, train_dataset=dataset, ) + assert trainer._image_token_id == trainer._tokenizer.convert_tokens_to_ids(image_token) previous_trainable_params = {n: param.clone() for n, param in trainer.model.named_parameters()} diff --git a/trl/trainer/distillation_trainer.py b/trl/trainer/distillation_trainer.py index 04a0803e315..c9c275ebe41 100644 --- a/trl/trainer/distillation_trainer.py +++ b/trl/trainer/distillation_trainer.py @@ -484,13 +484,13 @@ def __init__( # Resolve vision placeholder token IDs once. Used by the forward pass to rebuild mm_token_type_ids # when tool responses inject images into the completion (see _generate forward_kwargs block). - self._image_pad_token_id = None + self._image_token_id = None self._video_pad_token_id = None if self._is_vlm: - for candidate in ("<|image_pad|>", "<|image|>"): + for candidate in ("", "", "", "<|image|>", "<|image_pad|>"): tid = self._tokenizer.convert_tokens_to_ids(candidate) if tid != self._tokenizer.unk_token_id: - self._image_pad_token_id = tid + self._image_token_id = tid break tid = self._tokenizer.convert_tokens_to_ids("<|video_pad|>") if tid != self._tokenizer.unk_token_id: @@ -1111,8 +1111,8 @@ def _generate_single_turn(self, prompt_ids, images, multimodal_fields, has_tool_ # For VLM tool images: build token type IDs from the padded input IDs. if self._is_vlm and self.tools and has_tool_images: mm_ids = torch.zeros_like(padded_ids) - if self._image_pad_token_id is not None: - mm_ids[padded_ids == self._image_pad_token_id] = 1 + if self._image_token_id is not None: + mm_ids[padded_ids == self._image_token_id] = 1 if self._video_pad_token_id is not None: mm_ids[padded_ids == self._video_pad_token_id] = 2 @@ -1730,8 +1730,8 @@ def _generate_and_score_completions(self, inputs: list[dict[str, torch.Tensor | if self.tools and any(imgs for imgs in tool_images) and self._is_vlm: prompt_completion_ids = torch.cat([prompt_ids, completion_ids], dim=1) # (B, P+C) mm_ids = torch.zeros_like(prompt_completion_ids) - if self._image_pad_token_id is not None: - mm_ids[prompt_completion_ids == self._image_pad_token_id] = 1 + if self._image_token_id is not None: + mm_ids[prompt_completion_ids == self._image_token_id] = 1 if self._video_pad_token_id is not None: mm_ids[prompt_completion_ids == self._video_pad_token_id] = 2 diff --git a/trl/trainer/grpo_trainer.py b/trl/trainer/grpo_trainer.py index de126792d2f..a9557d01d5e 100644 --- a/trl/trainer/grpo_trainer.py +++ b/trl/trainer/grpo_trainer.py @@ -394,13 +394,13 @@ def __init__( # Resolve vision placeholder token IDs once. Used by the forward pass to rebuild mm_token_type_ids # when tool responses inject images into the completion (see _generate forward_kwargs block). - self._image_pad_token_id = None + self._image_token_id = None self._video_pad_token_id = None if self._is_vlm: - for candidate in ("<|image_pad|>", "<|image|>"): + for candidate in ("", "", "", "<|image|>", "<|image_pad|>"): tid = self._tokenizer.convert_tokens_to_ids(candidate) if tid != self._tokenizer.unk_token_id: - self._image_pad_token_id = tid + self._image_token_id = tid break tid = self._tokenizer.convert_tokens_to_ids("<|video_pad|>") if tid != self._tokenizer.unk_token_id: @@ -1881,8 +1881,8 @@ def _generate_single_turn(self, prompt_ids, images, multimodal_fields, has_tool_ # For VLM tool images: build token type IDs from the padded input IDs. if self._is_vlm and self.tools and has_tool_images: mm_ids = torch.zeros_like(padded_ids) - if self._image_pad_token_id is not None: - mm_ids[padded_ids == self._image_pad_token_id] = 1 + if self._image_token_id is not None: + mm_ids[padded_ids == self._image_token_id] = 1 if self._video_pad_token_id is not None: mm_ids[padded_ids == self._video_pad_token_id] = 2 @@ -2605,8 +2605,8 @@ def _generate_and_score_completions( # not just the prompt). if self.tools and any(imgs for imgs in tool_images) and self._is_vlm: mm_ids = torch.zeros_like(prompt_completion_ids) - if self._image_pad_token_id is not None: - mm_ids[prompt_completion_ids == self._image_pad_token_id] = 1 + if self._image_token_id is not None: + mm_ids[prompt_completion_ids == self._image_token_id] = 1 if self._video_pad_token_id is not None: mm_ids[prompt_completion_ids == self._video_pad_token_id] = 2