Skip to content
Open
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
17 changes: 11 additions & 6 deletions tests/test_distillation_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -1284,35 +1284,39 @@ 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",
"<image_soft_token>",
marks=pytest.mark.skipif(
Version(transformers.__version__) < Version("4.57.0"),
reason="transformers<4.57 Gemma3 image processor can't batch variable-size images",
),
),
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", "<IMG_CONTEXT>"),
("trl-internal-testing/tiny-LlavaNextForConditionalGeneration", "<image>"),
("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",
),
),
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",
Expand All @@ -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(
Expand All @@ -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()

Expand Down
18 changes: 11 additions & 7 deletions tests/test_grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -3893,29 +3893,32 @@ 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", "<image_soft_token>"),
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", "<IMG_CONTEXT>"),
("trl-internal-testing/tiny-LlavaNextForConditionalGeneration", "<image>"),
("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",
),
),
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",
Expand All @@ -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):
Expand All @@ -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()}

Expand Down
14 changes: 7 additions & 7 deletions trl/trainer/distillation_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ("<IMG_CONTEXT>", "<image_soft_token>", "<image>", "<|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:
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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

Expand Down
14 changes: 7 additions & 7 deletions trl/trainer/grpo_trainer.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 ("<IMG_CONTEXT>", "<image_soft_token>", "<image>", "<|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:
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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

Expand Down