diff --git a/data/v1_multimodal_demo.jsonl b/data/v1_multimodal_demo.jsonl new file mode 100644 index 000000000..2094834f6 --- /dev/null +++ b/data/v1_multimodal_demo.jsonl @@ -0,0 +1,25 @@ +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家,从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家,从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家,从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]} +{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家,从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]} + diff --git a/data/v1_multimodal_demo.yaml b/data/v1_multimodal_demo.yaml new file mode 100644 index 000000000..74189bc1f --- /dev/null +++ b/data/v1_multimodal_demo.yaml @@ -0,0 +1,4 @@ +multimodal_demo: + path: data/v1_multimodal_demo.jsonl + source: local + diff --git a/examples/v1/train_full/train_multimodal.yaml b/examples/v1/train_full/train_multimodal.yaml new file mode 100644 index 000000000..07f8a980c --- /dev/null +++ b/examples/v1/train_full/train_multimodal.yaml @@ -0,0 +1,27 @@ +model: Qwen/Qwen3.5-0.8B +model_class: llm + + +kernel_config: + name: auto + +quant_config: null + +dist_config: + name: fsdp2 + dcp_path: null + +### data +train_dataset: data/v1_multimodal_demo.yaml + +### training +output_dir: outputs/test_multimodal +micro_batch_size: 1 +cutoff_len: 2048 +learning_rate: 1.0e-4 +max_steps: 5 + +### sample +sample_backend: hf +max_new_tokens: 128 + diff --git a/src/llamafactory/v1/core/base_trainer.py b/src/llamafactory/v1/core/base_trainer.py index a1c9720cd..1aba01e3b 100644 --- a/src/llamafactory/v1/core/base_trainer.py +++ b/src/llamafactory/v1/core/base_trainer.py @@ -43,7 +43,7 @@ from ..utils.callbacks import ( TrainerCallback, TrainerState, ) -from ..utils.helper import compute_valid_tokens +from ..utils.helper import compute_valid_tokens, is_tokenizer, model_uses_mrope from ..utils.types import BatchInput, HFModel, ModelOutput, Tensor, TorchDataset from .rendering import Renderer from .utils.batching import BatchGenerator @@ -75,6 +75,7 @@ class BaseTrainer: self.dp_size = DistributedInterface().get_world_size(Dim.DP) self.cp_size = DistributedInterface().get_world_size(Dim.CP) self.model_input_names = self.renderer.processor.model_input_names + self._uses_mrope = model_uses_mrope(self.model.config) self._create_batch_generator() # Calculate num_training_steps: max_steps takes priority if set @@ -89,6 +90,9 @@ class BaseTrainer: if self.args.enable_activation_checkpointing: self.model.gradient_checkpointing_enable({"use_reentrant": False}) + # Note: under FSDP2 bf16, encoder-tower nn.LayerNorms are made dtype-safe for the + # checkpoint recompute inside the FSDP2 engine (see fsdp2.py prepare_model), so the + # tower keeps activation checkpointing too. self._deepspeed_engine = None dist_name = self.args.dist_config.name if self.args.dist_config is not None else None @@ -184,7 +188,11 @@ class BaseTrainer: "dist_config is None but distributed training is enabled; falling back to DistributedDataParallel." ) device_ids = None if self.device.type == "cpu" else [self.device.index] - self.model = DDP(self.model, device_ids=device_ids) + # Multimodal models invoke the vision tower only when a step carries media; a + # globally media-less step leaves vision params unused, which trips DDP's default + # all-params-used assertion. (FSDP tolerates a uniform skip; DDP does not.) + find_unused = not is_tokenizer(self.renderer.processor) + self.model = DDP(self.model, device_ids=device_ids, find_unused_parameters=find_unused) else: from ..plugins.trainer_plugins.distributed.interface import DistributedPlugin @@ -224,6 +232,9 @@ class BaseTrainer: model_inputs = { k: v.to(self.device, non_blocking=True) for k, v in batch.items() if isinstance(v, torch.Tensor) } + # Let mRoPE models build their own multimodal 3D position ids (see _uses_mrope in __init__). + if self._uses_mrope: + model_inputs.pop("position_ids", None) labels = batch["labels"].to(self.device, non_blocking=True) outputs: ModelOutput = model(**model_inputs) logits = outputs.logits.float() diff --git a/src/llamafactory/v1/core/model_engine.py b/src/llamafactory/v1/core/model_engine.py index 209cc494b..67ad8a4f6 100644 --- a/src/llamafactory/v1/core/model_engine.py +++ b/src/llamafactory/v1/core/model_engine.py @@ -150,8 +150,19 @@ class ModelEngine: if self.args.model_class == ModelClass.LLM: from transformers import AutoModelForCausalLM, AutoModelForImageTextToText - if type(self.model_config) in AutoModelForImageTextToText._model_mapping.keys(): + # AutoModelForMultimodalLM (audio / other multimodal LMs, e.g. Qwen2-Audio) was added in + # a newer transformers; fall back gracefully when it is absent (e.g. 4.57.1). + try: + from transformers import AutoModelForMultimodalLM + except ImportError: + AutoModelForMultimodalLM = None + + cfg_type = type(self.model_config) + if cfg_type in AutoModelForImageTextToText._model_mapping.keys(): AutoClass = AutoModelForImageTextToText + elif AutoModelForMultimodalLM is not None and cfg_type in AutoModelForMultimodalLM._model_mapping.keys(): + # Audio / other multimodal LMs (e.g. Qwen2-Audio) live here, not in CausalLM. + AutoClass = AutoModelForMultimodalLM else: AutoClass = AutoModelForCausalLM @@ -187,6 +198,10 @@ class ModelEngine: init_mode = self.args.init_config.name if self.args.init_config is not None else "init_on_default" model._init_mode = init_mode + if hasattr(model, "thinker"): + model = model.thinker + model._init_mode = init_mode + if self.args.peft_config is None: if self.is_train: logger.info_rank0("Fine-tuning mode: full tuning") diff --git a/src/llamafactory/v1/core/rendering/format.py b/src/llamafactory/v1/core/rendering/format.py index 048e61217..81cc3323f 100644 --- a/src/llamafactory/v1/core/rendering/format.py +++ b/src/llamafactory/v1/core/rendering/format.py @@ -1,4 +1,4 @@ -# Copyright 2025 the LlamaFactory team. +# Copyright 2026 the LlamaFactory team. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -14,13 +14,15 @@ """Message <-> HF-template plumbing for rendering. -Pure, stateless helpers: convert v1 ``Message`` to HF chat-template format. No tokenization policy -decisions live here -- only mechanical conversion used by ``rendering.py``. +Pure, stateless helpers: convert v1 ``Message`` to HF chat-template format, extract/count media, and +guard media placeholder counts. No tokenization policy decisions live here -- only mechanical +conversion used by ``rendering.py``. """ import json -from ...utils.types import Message +from ...utils.helper import get_tokenizer +from ...utils.types import Message, Processor _FALLBACK_CHATML_JINJA = ( @@ -33,28 +35,59 @@ _FALLBACK_CHATML_JINJA = ( ) -def _to_hf_messages(messages: list[Message]) -> list[dict]: +def _to_hf_messages(messages: list[Message], is_multimodal: bool = False) -> list[dict]: """Convert v1 Message format to HF format for apply_chat_template.""" hf_messages = [] for message in messages: tool_calls: list[dict] = [] reasoning_content = "" - text = "" - for content in message["content"]: - if content["type"] == "text": - text += content["value"] - elif content["type"] == "reasoning": - reasoning_content += content["value"] - elif content["type"] == "tool_call": - try: - tc = json.loads(content["value"]) - except json.JSONDecodeError as e: - raise ValueError(f"tool_call value is not valid JSON: {content['value']!r}") from e - if not isinstance(tc, dict) or "name" not in tc or "arguments" not in tc: - raise ValueError(f"tool_call must be a JSON object with 'name' and 'arguments' keys, got {tc!r}") - tool_calls.append({"type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}}) - hf_msg = {"role": message["role"], "content": text} + if is_multimodal: + hf_content = [] + for content in message["content"]: + if content["type"] == "text": + hf_content.append({"type": "text", "text": content["value"]}) + elif content["type"] == "reasoning": + reasoning_content += content["value"] + elif content["type"] == "tool_call": + try: + tc = json.loads(content["value"]) + except json.JSONDecodeError as e: + raise ValueError(f"tool_call value is not valid JSON: {content['value']!r}") from e + if not isinstance(tc, dict) or "name" not in tc or "arguments" not in tc: + raise ValueError( + f"tool_call must be a JSON object with 'name' and 'arguments' keys, got {tc!r}" + ) + tool_calls.append( + {"type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}} + ) + elif content["type"] == "image_url": + hf_content.append({"type": "image", "image": content["value"]}) + elif content["type"] == "video_url": + hf_content.append({"type": "video", "video": content["value"]}) + elif content["type"] == "audio_url": + hf_content.append({"type": "audio", "audio": content["value"]}) + hf_msg = {"role": message["role"], "content": hf_content} + else: + text = "" + for content in message["content"]: + if content["type"] == "text": + text += content["value"] + elif content["type"] == "reasoning": + reasoning_content += content["value"] + elif content["type"] == "tool_call": + try: + tc = json.loads(content["value"]) + except json.JSONDecodeError as e: + raise ValueError(f"tool_call value is not valid JSON: {content['value']!r}") from e + if not isinstance(tc, dict) or "name" not in tc or "arguments" not in tc: + raise ValueError( + f"tool_call must be a JSON object with 'name' and 'arguments' keys, got {tc!r}" + ) + tool_calls.append( + {"type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}} + ) + hf_msg = {"role": message["role"], "content": text} if tool_calls: hf_msg["tool_calls"] = tool_calls @@ -63,3 +96,75 @@ def _to_hf_messages(messages: list[Message]) -> list[dict]: hf_messages.append(hf_msg) return hf_messages + + +def _extract_media_from_messages(messages: list[Message]) -> tuple[list, list, list]: + """Extract image, video and audio paths/values from messages in order.""" + images, videos, audios = [], [], [] + for message in messages: + for content in message["content"]: + if content["type"] == "image_url": + images.append(content["value"]) + elif content["type"] == "video_url": + videos.append(content["value"]) + elif content["type"] == "audio_url": + audios.append(content["value"]) + return images, videos, audios + + +def _count_media_in_messages(messages: list[Message]) -> tuple[int, int, int]: + """Count total images, videos and audios in messages.""" + n_images, n_videos, n_audios = 0, 0, 0 + for message in messages: + for content in message["content"]: + if content["type"] == "image_url": + n_images += 1 + elif content["type"] == "video_url": + n_videos += 1 + elif content["type"] == "audio_url": + n_audios += 1 + return n_images, n_videos, n_audios + + +def _load_audios(values: list, sampling_rate: int) -> list: + """Load audio inputs into mono waveforms resampled to ``sampling_rate``.""" + import numpy as np + import torchaudio + + results = [] + for value in values: + if isinstance(value, np.ndarray): + results.append(value) + continue + + waveform, sr = torchaudio.load(value) + if waveform.shape[0] > 1: # downmix to mono + waveform = waveform.mean(dim=0, keepdim=True) + if sr != sampling_rate: + waveform = torchaudio.functional.resample(waveform, sr, sampling_rate) + results.append(waveform.squeeze(0).numpy()) + return results + + +def _check_placeholder_counts( + processor: "Processor", full_text: str, n_images: int, n_videos: int, n_audios: int = 0 +) -> None: + """Guard: every media placeholder in the rendered text must originate from a media block.""" + tokenizer = get_tokenizer(processor) + for attr, count, kind in ( + ("image_token_id", n_images, "image"), + ("video_token_id", n_videos, "video"), + ("audio_token_id", n_audios, "audio"), + ): + tid = getattr(processor, attr, None) + if tid is None: + tid = getattr(tokenizer, attr, None) + if tid is None: + continue + placeholder = tokenizer.convert_ids_to_tokens(tid) + seen = full_text.count(placeholder) + if seen != count: + raise ValueError( + f"{kind} placeholder count ({seen}) != number of {kind} blocks ({count}); " + "media must be provided via image_url/video_url content blocks." + ) diff --git a/src/llamafactory/v1/core/rendering/rendering.py b/src/llamafactory/v1/core/rendering/rendering.py index 87a62e87f..94aa128d1 100644 --- a/src/llamafactory/v1/core/rendering/rendering.py +++ b/src/llamafactory/v1/core/rendering/rendering.py @@ -1,4 +1,4 @@ -# Copyright 2025 the LlamaFactory team. +# Copyright 2026 the LlamaFactory team. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -19,24 +19,29 @@ sibling modules: - ``format`` -- v1<->HF message conversion - ``escape`` -- special-token escaping (prompt-injection hardening) -Assistant supervision is located WITHOUT a per-model marker table: a training sample is rendered -so that its last message is the supervised assistant turn, and that turn's token span is recovered -by a single prompt/full difference -- encode the prompt (everything up to and including the -assistant role header, via ``add_generation_prompt=True``) and the full sequence, then the tail of -the full sequence that the prompt does not cover is exactly this turn. Multi-turn conversations are -split into one sample per supervised turn (see ``process_samples``) so the supervised turn is always -the last one; this keeps the diff on the only boundary that is prefix-stable across chat templates -(appending the final assistant turn never restripts earlier turns), so models with reasoning-history -stripping (e.g. Qwen3 ````) are handled correctly without hard-coding role markers. +Note: ``position_ids`` are assigned by ``process_samples`` (1-based); multimodal (mrope) position +ids are expected to be recomputed by the model/trainer. """ + import json +import numpy as np +import torch + from ...utils.constants import IGNORE_INDEX -from ...utils.helper import get_tokenizer +from ...utils.helper import get_tokenizer, is_tokenizer from ...utils.types import Message, ModelInput, Processor, Sample +from ..utils.collation import _MULTIMODAL_PASSTHROUGH_KEYS from .escape import _escape_special, _escape_special_in_messages, _special_token_strings -from .format import _FALLBACK_CHATML_JINJA, _to_hf_messages +from .format import ( + _FALLBACK_CHATML_JINJA, + _check_placeholder_counts, + _count_media_in_messages, + _extract_media_from_messages, + _load_audios, + _to_hf_messages, +) def _render_messages( @@ -46,20 +51,23 @@ def _render_messages( is_generate: bool = False, **kwargs, ) -> ModelInput: - r"""Render messages using the model's own chat template. + r"""Render messages using the model's own chat template, locating supervision by a prompt/full diff. Note: ``position_ids`` are not produced here; ``process_samples`` assigns a 1-based range. """ tokenizer = get_tokenizer(processor) - if not getattr(tokenizer, "chat_template", None): - tokenizer.chat_template = _FALLBACK_CHATML_JINJA + is_multimodal = not is_tokenizer(processor) + + template_caller = processor if is_multimodal else tokenizer + if not getattr(template_caller, "chat_template", None): + template_caller.chat_template = _FALLBACK_CHATML_JINJA # 0. Neutralize special-token strings in user-controlled text (no-op for normal data). specials = _special_token_strings(tokenizer) special_ids = {tid for tid, t in tokenizer.added_tokens_decoder.items() if getattr(t, "special", False)} messages = _escape_special_in_messages(messages, specials, special_ids, tokenizer) - hf_messages = _to_hf_messages(messages) + hf_messages = _to_hf_messages(messages, is_multimodal=is_multimodal) tools_parsed = None if tools: @@ -70,36 +78,76 @@ def _render_messages( raise ValueError(f"tools is not valid JSON: {tools!r}") from e if not isinstance(tools_parsed, list): tools_parsed = [tools_parsed] - if not is_generate and hf_messages and hf_messages[-1].get("reasoning_content"): - kwargs["enable_thinking"] = True - def _encode(msgs: list[dict], add_generation_prompt: bool) -> list[int]: - text = tokenizer.apply_chat_template( - msgs, tokenize=False, add_generation_prompt=add_generation_prompt, tools=tools_parsed, **kwargs + if not is_generate and hf_messages and hf_messages[-1]["role"] == "assistant": + kwargs["enable_thinking"] = bool(hf_messages[-1].get("reasoning_content")) + + def _encode(hf_msgs: list[dict], src_msgs: list[Message], add_generation_prompt: bool): + """Render + tokenize, expanding media via the processor. Returns (input_ids, mm_outputs).""" + text = template_caller.apply_chat_template( + hf_msgs, tokenize=False, add_generation_prompt=add_generation_prompt, tools=tools_parsed, **kwargs ) - return tokenizer(text, add_special_tokens=False)["input_ids"] + if is_multimodal and _count_media_in_messages(src_msgs) != (0, 0, 0): + images, videos, audios = _extract_media_from_messages(src_msgs) + # Every placeholder must come from a media block (escaping broke any literal ones). + _check_placeholder_counts(processor, text, len(images), len(videos), len(audios)) + proc_kwargs = {"return_tensors": "pt"} + if images: + proc_kwargs["images"] = images + if videos: + proc_kwargs["videos"] = videos + if audios: + # Audio processors want decoded waveforms at the model's sampling rate, not paths. + proc_kwargs["audio"] = _load_audios(audios, processor.feature_extractor.sampling_rate) + mm_outputs = processor(text=text, **proc_kwargs) + return mm_outputs["input_ids"][0].tolist(), mm_outputs + return tokenizer(text, add_special_tokens=False)["input_ids"], None - # 1. Full sequence, used verbatim. - input_ids = _encode(hf_messages, add_generation_prompt=is_generate) + # 1. Full sequence (used verbatim), plus its multimodal feature outputs. + input_ids, outputs = _encode(hf_messages, messages, add_generation_prompt=is_generate) n = len(input_ids) + def _attach_multimodal(result: ModelInput) -> None: + if outputs is None: + return + for key in _MULTIMODAL_PASSTHROUGH_KEYS: + if key in outputs: + result[key] = outputs[key] + mm_type_ids = outputs["mm_token_type_ids"][0].tolist() if "mm_token_type_ids" in outputs else None + + for attr, marker in (("image_token_id", 1), ("video_token_id", 2), ("audio_token_id", 3)): + token_id = getattr(processor, attr, None) + if token_id is None: + token_id = getattr(tokenizer, attr, None) + if token_id is None or token_id not in input_ids: + continue + if mm_type_ids is not None and marker in mm_type_ids: + continue + if mm_type_ids is None: + mm_type_ids = [0] * len(input_ids) + mm_type_ids = [marker if tid == token_id else t for t, tid in zip(mm_type_ids, input_ids)] + + if mm_type_ids is not None: + result["mm_token_type_ids"] = mm_type_ids + if is_generate: # Generation prompt only -- nothing is supervised. - return ModelInput( + result = ModelInput( input_ids=input_ids, attention_mask=[1] * n, labels=[IGNORE_INDEX] * n, loss_weights=[0.0] * n, ) + _attach_multimodal(result) + return result - # 2. Locate the supervised (last) assistant turn by a prompt/full diff (no marker table). if not messages or messages[-1]["role"] != "assistant": raise ValueError( "training render expects the last message to be the supervised assistant turn; " "multi-turn conversations are split per turn in process_samples." ) - prompt_ids = _encode(hf_messages[:-1], add_generation_prompt=True) + prompt_ids, _ = _encode(hf_messages[:-1], messages[:-1], add_generation_prompt=True) if input_ids[: len(prompt_ids)] != prompt_ids: # The prompt must be a token-prefix of the full sequence for the diff to be valid. If a # template re-renders earlier turns when the final turn is appended, fail loud rather than @@ -117,16 +165,21 @@ def _render_messages( labels.append(tid if supervised else IGNORE_INDEX) loss_weights.append(weight) - return ModelInput( + result = ModelInput( input_ids=input_ids, attention_mask=[1] * n, labels=labels, loss_weights=loss_weights, ) + _attach_multimodal(result) + return result class Renderer: - def __init__(self, processor: Processor) -> None: + def __init__(self, processor: Processor, config=None): + # ``config`` is accepted for call-site compatibility (ModelEngine passes the model config) + # but is no longer needed: supervision is located by a prompt/full diff, not a per-model + # marker table, so the renderer is model-agnostic. self.processor = processor def render_messages( @@ -152,6 +205,61 @@ class Renderer: """ return _render_messages(self.processor, messages, tools, is_generate, **kwargs) + def get_dummy_media_fragment(self, modality: str) -> dict: + """Build (and cache) a minimal valid media fragment for ``modality`` ("image"|"video"|"audio").""" + if modality not in ("image", "video", "audio"): + raise ValueError(f"Unsupported dummy media modality: {modality!r} (expected image/video/audio).") + if is_tokenizer(self.processor): + raise RuntimeError("Cannot build a dummy media fragment for a text-only processor.") + + if not hasattr(self, "_dummy_fragments"): + self._dummy_fragments: dict[str, dict] = {} + if modality in self._dummy_fragments: + return self._dummy_fragments[modality] + + from PIL import Image as _PILImage + + if modality == "image": + media_block = {"type": "image_url", "value": _PILImage.new("RGB", (64, 64))} + target, presence_key = 1, "pixel_values" + elif modality == "video": + # A minimal clip: the temporal patch size is typically 2, so provide two frames. + media_block = {"type": "video_url", "value": np.zeros((2, 64, 64, 3), dtype=np.uint8)} + target, presence_key = 2, "pixel_values_videos" + else: + # A short synthetic waveform at the model's sampling rate; the feature extractor pads it. + sr = self.processor.feature_extractor.sampling_rate + media_block = {"type": "audio_url", "value": np.zeros(sr // 10, dtype=np.float32)} + target, presence_key = 3, "input_features" + + messages: list[Message] = [ + {"role": "user", "content": [media_block]}, + {"role": "assistant", "content": [{"type": "text", "value": "ok"}]}, + ] + rendered = self.render_messages(messages) + + mm_type_ids = rendered.get("mm_token_type_ids") + if not mm_type_ids or target not in mm_type_ids or presence_key not in rendered: + raise RuntimeError(f"Processor did not emit {modality} placeholder tokens for the dummy sample.") + + positions = [i for i, t in enumerate(mm_type_ids) if t == target] + # Include the surrounding start/end delimiters (vision_start/end or audio_bos/eos) so the + # fragment matches exactly what the template emits around real media. + lo = max(positions[0] - 1, 0) + hi = min(positions[-1] + 2, len(rendered["input_ids"])) + + fragment: dict = { + "input_ids": list(rendered["input_ids"][lo:hi]), + "mm_token_type_ids": list(mm_type_ids[lo:hi]), + } + + for key in _MULTIMODAL_PASSTHROUGH_KEYS: + if key in rendered: + fragment[key] = rendered[key] + + self._dummy_fragments[modality] = fragment + return fragment + def process_samples(self, samples: list[Sample]) -> list[ModelInput]: """Process samples to model input. @@ -189,6 +297,17 @@ class Renderer: model_input["position_ids"] = list(range(1, len(chosen_input["input_ids"]) + 1)) + list( range(1, len(rejected_input["input_ids"]) + 1) ) + + for key in _MULTIMODAL_PASSTHROUGH_KEYS: + tensors = [inp[key] for inp in (chosen_input, rejected_input) if key in inp] + if tensors: + model_input[key] = torch.cat(tensors, dim=0) + + if "mm_token_type_ids" in chosen_input or "mm_token_type_ids" in rejected_input: + chosen_mm = chosen_input.get("mm_token_type_ids", [0] * len(chosen_input["input_ids"])) + rejected_mm = rejected_input.get("mm_token_type_ids", [0] * len(rejected_input["input_ids"])) + model_input["mm_token_type_ids"] = chosen_mm + rejected_mm + rendered.append(model_input) else: raise ValueError("No valid messages or chosen_messages/rejected_messages found in sample.") diff --git a/src/llamafactory/v1/core/utils/batching.py b/src/llamafactory/v1/core/utils/batching.py index c27c02a65..6009bac94 100644 --- a/src/llamafactory/v1/core/utils/batching.py +++ b/src/llamafactory/v1/core/utils/batching.py @@ -1,4 +1,4 @@ -# Copyright 2025 the LlamaFactory team. +# Copyright 2026 the LlamaFactory team. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -31,13 +31,16 @@ from torch.utils.data import default_collate from torchdata.stateful_dataloader import StatefulDataLoader from torchdata.stateful_dataloader.sampler import StatefulDistributedSampler +from ...accelerator.helper import ReduceOp from ...accelerator.interface import Dim, DistributedInterface from ...config import BatchingStrategy from ...utils import logging -from ...utils.helper import pad_and_truncate +from ...utils.constants import IGNORE_INDEX +from ...utils.helper import is_tokenizer from ...utils.objects import StatefulBuffer -from ...utils.types import BatchInfo, BatchInput, ModelInput, TorchDataset +from ...utils.types import BatchInfo, BatchInput, ModelInput, Tensor, TorchDataset from ..rendering import Renderer +from .collation import _MULTIMODAL_PASSTHROUGH_KEYS, pad_and_truncate logger = logging.get_logger(__name__) @@ -45,7 +48,82 @@ logger = logging.get_logger(__name__) __all__ = ["BatchGenerator"] -def default_collate_fn(buffer: StatefulBuffer, batch_info: BatchInfo) -> list[BatchInput] | None: +# (modality, presence/feature key, grid key, mm_token_type_ids marker) for encoder-tower alignment. +# The presence key is what survives collation when the modality is present; the grid key is unused +# here (kept for parity with the collation specs). Audio carries no grid -- feature_attention_mask +# rides along as a passthrough feature. +_ALIGN_MODALITIES = ( + ("image", "pixel_values", "image_grid_thw", 1), + ("video", "pixel_values_videos", "video_grid_thw", 2), + ("audio", "input_features", "feature_attention_mask", 3), +) + + +def _collate_micro_batch(micro_batch: list[ModelInput], cutoff_len: int) -> BatchInput: + """Pad/truncate then collate one micro batch (text fields stacked, MM features dim-0 concat).""" + padded = pad_and_truncate(micro_batch, cutoff_len) + standard_samples = [{k: v for k, v in s.items() if k not in _MULTIMODAL_PASSTHROUGH_KEYS} for s in padded] + collated = default_collate(standard_samples) + for key in _MULTIMODAL_PASSTHROUGH_KEYS: + tensors = [s[key] for s in padded if key in s] + if tensors: + collated[key] = torch.cat(tensors, dim=0) + return collated + + +def _inject_dummy_into_collated(collated: BatchInput, fragment: dict, marker: int) -> None: + """Append a zero-loss dummy media fragment to an already-collated micro batch, in place. + + Operates *after* pad_and_truncate so it reflects post-truncation presence: an image whose + placeholder tokens were partially cut is deleted by ``_align_multimodal_on_truncation``, + turning that sample text-only -- which must be detected here (not before truncation) or the + vision-tower call count still desyncs across ranks. + + The dummy tokens are appended (extra columns) into row 0 only; other rows get padding there. + Causal attention keeps every real token's logits unchanged; the dummy carries IGNORE_INDEX + labels and zero loss weight, so it contributes nothing to the loss while forcing the (FSDP- + sharded) vision tower to run. + """ + bsz, seqlen = collated["input_ids"].shape + frag_ids = torch.tensor(fragment["input_ids"], dtype=collated["input_ids"].dtype) + frag_len = frag_ids.numel() + frag_mm = torch.tensor(fragment["mm_token_type_ids"], dtype=torch.long) + new_len = seqlen + frag_len + + def _grow(tensor: Tensor, pad_value, row0_tail=None) -> Tensor: + out = torch.full((bsz, new_len), pad_value, dtype=tensor.dtype) + out[:, :seqlen] = tensor + if row0_tail is not None: + out[0, seqlen:] = row0_tail.to(tensor.dtype) + return out + + collated["input_ids"] = _grow(collated["input_ids"], 0, frag_ids) + collated["attention_mask"] = _grow(collated["attention_mask"], 0) + collated["attention_mask"][0, seqlen:] = 1 + collated["labels"] = _grow(collated["labels"], IGNORE_INDEX) # dummy region stays ignored + collated["loss_weights"] = _grow(collated["loss_weights"], 0.0) + if "position_ids" in collated: + pos = _grow(collated["position_ids"], 0) + pos[0, seqlen:] = torch.arange(seqlen + 1, new_len + 1, dtype=pos.dtype) + collated["position_ids"] = pos + + mm = collated.get("mm_token_type_ids") + if mm is not None: + collated["mm_token_type_ids"] = _grow(mm, 0, frag_mm) + else: + mm = torch.zeros((bsz, new_len), dtype=torch.long) + mm[0, seqlen:] = frag_mm + collated["mm_token_type_ids"] = mm + + for key, value in fragment.items(): + if key in ("input_ids", "mm_token_type_ids"): + continue + collated[key] = torch.cat([collated[key], value], dim=0) if key in collated else value + + +def default_collate_fn( + buffer: StatefulBuffer, batch_info: BatchInfo, renderer: Renderer | None = None +) -> list[BatchInput] | None: micro_batch_size = batch_info["micro_batch_size"] num_micro_batch = batch_info["num_micro_batch"] cutoff_len = batch_info["cutoff_len"] @@ -54,10 +132,24 @@ def default_collate_fn(buffer: StatefulBuffer, batch_info: BatchInfo) -> list[Ba return None samples = buffer.get(batch_size) - batch = [] - for i in range(num_micro_batch): - micro_batch = samples[i * micro_batch_size : (i + 1) * micro_batch_size] - batch.append(default_collate(pad_and_truncate(micro_batch, cutoff_len))) + micro_batches = [samples[i * micro_batch_size : (i + 1) * micro_batch_size] for i in range(num_micro_batch)] + + # Collate first; presence is judged on the *post-truncation* result, since truncation can + # delete a partially-cut image and turn a sample text-only (see _inject_dummy_into_collated). + batch = [_collate_micro_batch(mb, cutoff_len) for mb in micro_batches] + + if renderer is not None and not is_tokenizer(renderer.processor): + present = torch.zeros((num_micro_batch, len(_ALIGN_MODALITIES)), dtype=torch.int64) + for i, collated in enumerate(batch): + for m, (_, pixel_key, _, _) in enumerate(_ALIGN_MODALITIES): + present[i, m] = int(pixel_key in collated) + + present = DistributedInterface().all_reduce(present, op=ReduceOp.MAX, dim=Dim.DP) + + for i, collated in enumerate(batch): + for m, (modality, pixel_key, _, marker) in enumerate(_ALIGN_MODALITIES): + if present[i, m] and pixel_key not in collated: + _inject_dummy_into_collated(collated, renderer.get_dummy_media_fragment(modality), marker) return batch @@ -227,8 +319,17 @@ class BatchGenerator(Iterator): def _generate_batch(self) -> list[BatchInput] | None: if self.batching_strategy == BatchingStrategy.NORMAL: - return default_collate_fn(self._buffer, self._batch_info) + return default_collate_fn(self._buffer, self._batch_info, self.renderer) else: + # Non-NORMAL strategies (dynamic / padding_free) collate ragged pixel tensors with a + # bare default_collate and have no vision-tower alignment, so multimodal data would + # crash or hang. Fail loud instead of silently mishandling it. + if any(k in s for s in self._buffer.samples for k in _MULTIMODAL_PASSTHROUGH_KEYS): + raise NotImplementedError( + f"batching_strategy={self.batching_strategy.value!r} does not support multimodal data; " + "use the NORMAL strategy for image/video training." + ) + from ...plugins.trainer_plugins.batching import BatchingPlugin return BatchingPlugin(self.batching_strategy).generate_batch(self._buffer, self._batch_info) diff --git a/src/llamafactory/v1/core/utils/collation.py b/src/llamafactory/v1/core/utils/collation.py new file mode 100644 index 000000000..6f33f85f6 --- /dev/null +++ b/src/llamafactory/v1/core/utils/collation.py @@ -0,0 +1,277 @@ +# Copyright 2026 the LlamaFactory team. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Batch collation utils: padding/truncation and multimodal-feature alignment. + +These operate on already-rendered ``ModelInput`` dicts (token lists + pixel tensors) and produce +padded ``BatchInput`` tensors. They are pure batching concerns -- independent of how a sample was +rendered -- and are consumed by the batch generators in ``core/utils/batching.py`` and +``plugins/trainer_plugins/batching.py``. Kept out of ``rendering.py`` so that file is only about +turning messages into a single tokenized sample. +""" + +import torch + +from ...utils.constants import IGNORE_INDEX +from ...utils.types import BatchInput, ModelInput, Tensor + + +# Multimodal feature keys the processor emits per sample. They are NOT padded/stacked like text +# fields: pixel/audio-feature tensors are ragged (variable patch / frame counts), so the collators +# concatenate them along dim 0 instead. Shared by rendering (which copies them verbatim from the +# processor) and the collators (which merge them across a micro batch). +_MULTIMODAL_PASSTHROUGH_KEYS = frozenset( + { + "pixel_values", + "image_grid_thw", + "pixel_values_videos", + "video_grid_thw", + "second_per_grid_ts", # Qwen2.5-VL name for the video temporal grid spacing + "video_second_per_grid", # Qwen2.5-Omni name for the same (fed to get_rope_index) + "input_features", + "feature_attention_mask", + } +) + + +def _pad_and_truncate(tensor: Tensor, max_seqlen: int, pad_value: int = 0) -> Tensor: + if tensor.shape[-1] >= max_seqlen: + return tensor[..., :max_seqlen] + + pad_shape = list(tensor.shape) + pad_shape[-1] = max_seqlen - tensor.shape[-1] + pad_tensor = torch.full(pad_shape, pad_value, dtype=tensor.dtype, device=tensor.device) + return torch.cat([tensor, pad_tensor], dim=-1) + + +def _align_grid_media( + sample: ModelInput, + mm_type_ids: list[int], + max_length: int, + *, + target: int, + grid_key: str, + pixel_key: str, +) -> list[int]: + """Trim and zero one modality's orphaned tokens for a single sample. + + Layout-agnostic: a media item's placeholder tokens may be a single contiguous run or split + into per-frame sub-runs; completeness is decided per token *position*, so both are handled + identically. + + Returns the (possibly updated) ``mm_token_type_ids`` so chained calls see earlier zeroing. + """ + if grid_key not in sample or pixel_key not in sample: + return mm_type_ids + + grid = sample[grid_key] + n_items = len(grid) + if n_items == 0: + return mm_type_ids + + positions = [i for i, t in enumerate(mm_type_ids) if t == target] + patches_per_item = [int(grid[i].prod()) for i in range(n_items)] + total_patches = sum(patches_per_item) + total_tokens = len(positions) + + # merge_size**2 = pixel patches per placeholder token, derived from the data. Bail out + # untouched if the sample is inconsistent. + if total_tokens == 0 or total_patches % total_tokens != 0: + return mm_type_ids + merge_sq = total_patches // total_tokens + tokens_per_item = [p // merge_sq for p in patches_per_item] + if sum(tokens_per_item) != total_tokens: + return mm_type_ids + + # Each item owns a contiguous slice of `positions`; it is complete iff its last + # placeholder token lands inside the kept window [0, max_length). + n_complete = 0 + cum = 0 + for n_i in tokens_per_item: + if positions[cum + n_i - 1] < max_length: + n_complete += 1 + cum += n_i + else: + break + + if n_complete >= n_items: + return mm_type_ids + + # Trim pixel features and grid to the complete prefix. + keep_patches = sum(patches_per_item[:n_complete]) + sample[pixel_key] = sample[pixel_key][:keep_patches] + sample[grid_key] = grid[:n_complete] + + # Zero out orphaned placeholder tokens that fall inside the kept window; tokens + # beyond max_length are removed by truncation anyway (positions are sorted). + input_ids = list(sample["input_ids"]) + mm_type_ids = list(mm_type_ids) + labels = list(sample["labels"]) if "labels" in sample else None + loss_weights = list(sample["loss_weights"]) if "loss_weights" in sample else None + + for pos in positions[cum:]: + if pos >= max_length: + break + input_ids[pos] = 0 + mm_type_ids[pos] = 0 + if labels is not None: + labels[pos] = IGNORE_INDEX + if loss_weights is not None: + loss_weights[pos] = 0.0 + + sample["input_ids"] = input_ids + sample["mm_token_type_ids"] = mm_type_ids + if labels is not None: + sample["labels"] = labels + if loss_weights is not None: + sample["loss_weights"] = loss_weights + return mm_type_ids + + +def _align_audio(sample: ModelInput, mm_type_ids: list[int], max_length: int, *, target: int = 3) -> list[int]: + """Trim and zero orphaned audio tokens for a single sample on truncation. + + Returns the (possibly updated) ``mm_token_type_ids``. + """ + if "input_features" not in sample or "feature_attention_mask" not in sample: + return mm_type_ids + + n_items = sample["input_features"].shape[0] + if n_items == 0: + return mm_type_ids + + positions = [i for i, t in enumerate(mm_type_ids) if t == target] + if not positions: + return mm_type_ids + + # Group the marked positions into maximal contiguous runs; each run is one audio's token span. + runs: list[tuple[int, int]] = [] + run_start = prev = positions[0] + for pos in positions[1:]: + if pos != prev + 1: + runs.append((run_start, prev)) + run_start = pos + prev = pos + runs.append((run_start, prev)) + + # Layout must match the feature rows one-to-one, else bail rather than corrupt the mapping. + if len(runs) != n_items: + return mm_type_ids + + # An audio is complete iff its last placeholder token lands inside the kept window. + n_complete = 0 + for _start, end in runs: + if end < max_length: + n_complete += 1 + else: + break + + if n_complete >= n_items: + return mm_type_ids + + # Trim feature rows to the complete prefix. + sample["input_features"] = sample["input_features"][:n_complete] + sample["feature_attention_mask"] = sample["feature_attention_mask"][:n_complete] + + # Zero out orphaned placeholder tokens that fall inside the kept window; tokens beyond + # max_length are removed by truncation anyway. + input_ids = list(sample["input_ids"]) + mm_type_ids = list(mm_type_ids) + labels = list(sample["labels"]) if "labels" in sample else None + loss_weights = list(sample["loss_weights"]) if "loss_weights" in sample else None + + for start, end in runs[n_complete:]: + for pos in range(start, end + 1): + if pos >= max_length: + break + input_ids[pos] = 0 + mm_type_ids[pos] = 0 + if labels is not None: + labels[pos] = IGNORE_INDEX + if loss_weights is not None: + loss_weights[pos] = 0.0 + + sample["input_ids"] = input_ids + sample["mm_token_type_ids"] = mm_type_ids + if labels is not None: + sample["labels"] = labels + if loss_weights is not None: + sample["loss_weights"] = loss_weights + return mm_type_ids + + +def _align_multimodal_on_truncation(sample: ModelInput, max_length: int) -> ModelInput: + """Remove orphaned multimodal data when the sequence will be truncated. + + When cutoff_len truncates input_ids, media whose placeholder tokens are partially cut lose + their token<->feature correspondence. Trims pixel_values/grid_thw (vision) and + input_features/feature_attention_mask (audio) to the complete items and zeros out orphaned + placeholder tokens so the model ignores them. + """ + mm_type_ids = sample.get("mm_token_type_ids") + if mm_type_ids is None: + return sample + + sample = dict(sample) + + mm_type_ids = _align_grid_media( + sample, mm_type_ids, max_length, target=1, grid_key="image_grid_thw", pixel_key="pixel_values" + ) + mm_type_ids = _align_grid_media( + sample, mm_type_ids, max_length, target=2, grid_key="video_grid_thw", pixel_key="pixel_values_videos" + ) + mm_type_ids = _align_audio(sample, mm_type_ids, max_length, target=3) + + # Remove empty multimodal fields entirely + if "image_grid_thw" in sample and len(sample["image_grid_thw"]) == 0: + del sample["pixel_values"] + del sample["image_grid_thw"] + if "video_grid_thw" in sample and len(sample["video_grid_thw"]) == 0: + del sample["pixel_values_videos"] + del sample["video_grid_thw"] + if "input_features" in sample and sample["input_features"].shape[0] == 0: + del sample["input_features"] + del sample["feature_attention_mask"] + + return sample + + +def pad_and_truncate(samples: list[ModelInput], max_seqlen: int) -> list[BatchInput]: + max_length = min(max(len(sample["input_ids"]) for sample in samples), max_seqlen) + padded_samples = [] + for sample in samples: + # Align multimodal fields before truncation: remove images/videos whose + # placeholder tokens would be partially cut, preventing pixel<->token mismatch. + if len(sample["input_ids"]) > max_length and any(k in sample for k in _MULTIMODAL_PASSTHROUGH_KEYS): + sample = _align_multimodal_on_truncation(sample, max_length) + + padded_sample = {} + for key, value in sample.items(): + if key in _MULTIMODAL_PASSTHROUGH_KEYS: + padded_sample[key] = value + continue + + if "label" in key: + pad_value = IGNORE_INDEX + else: + pad_value = 0 + + if not isinstance(value, str): + padded_sample[key] = _pad_and_truncate(torch.tensor(value), max_length, pad_value) + else: + padded_sample[key] = value + + padded_samples.append(padded_sample) + + return padded_samples diff --git a/src/llamafactory/v1/plugins/data_plugins/converter.py b/src/llamafactory/v1/plugins/data_plugins/converter.py index 7075fe5dc..556f83ca1 100644 --- a/src/llamafactory/v1/plugins/data_plugins/converter.py +++ b/src/llamafactory/v1/plugins/data_plugins/converter.py @@ -14,11 +14,13 @@ import json +import re from typing import Any, Literal, NotRequired, TypedDict from ...utils import logging +from ...utils.constants import AUDIO_PLACEHOLDER, IMAGE_PLACEHOLDER, VIDEO_PLACEHOLDER from ...utils.plugin import BasePlugin -from ...utils.types import DPOSample, Sample, SFTSample, ToolCall +from ...utils.types import Content, DPOSample, Sample, SFTSample, ToolCall logger = logging.get_logger(__name__) @@ -29,6 +31,9 @@ class AlpacaSample(TypedDict, total=False): instruction: str input: NotRequired[str] output: str + images: NotRequired[list[str] | str] + videos: NotRequired[list[str] | str] + audios: NotRequired[list[str] | str] SharegptMessage = TypedDict( @@ -40,6 +45,9 @@ SharegptMessage = TypedDict( class SharegptSample(TypedDict, total=False): conversations: list[SharegptMessage] tools: NotRequired[str] + images: NotRequired[list[str] | str] + videos: NotRequired[list[str] | str] + audios: NotRequired[list[str] | str] class OpenaiMessage(TypedDict, total=False): @@ -54,6 +62,65 @@ class OpenaiSample(TypedDict, total=False): class PairSample(TypedDict, total=False): chosen: list[OpenaiMessage] rejected: list[OpenaiMessage] + images: NotRequired[list[str] | str] + videos: NotRequired[list[str] | str] + audios: NotRequired[list[str] | str] + + +# Inline media tag -> v1 content block type, and the raw-sample column holding the paths. +_MEDIA_SPECS: tuple[tuple[str, str, str], ...] = ( + (IMAGE_PLACEHOLDER, "image_url", "images"), + (VIDEO_PLACEHOLDER, "video_url", "videos"), + (AUDIO_PLACEHOLDER, "audio_url", "audios"), +) +_TAG_TO_BLOCK = {tag: block_type for tag, block_type, _col in _MEDIA_SPECS} +_TAG_PATTERN = re.compile("(" + "|".join(re.escape(tag) for tag, _b, _c in _MEDIA_SPECS) + ")") + + +def _as_media_list(value: Any) -> list: + """Normalize a media column value into a list of paths/URLs (None -> [], scalar -> [scalar]).""" + if value is None: + return [] + if isinstance(value, (list, tuple)): + return list(value) + return [value] + + +def _build_media_iters(raw_sample: dict[str, Any]) -> dict[str, Any]: + """Build per-modality path iterators from a raw sample's media columns.""" + return {tag: iter(_as_media_list(raw_sample.get(col))) for tag, _block_type, col in _MEDIA_SPECS} + + +def _to_content_blocks(text: str, media_iters: dict[str, Any]) -> list[Content]: + """Split ``text`` on inline media placeholders, interleaving media-url content blocks. + + Each placeholder consumes the next path from its modality iterator (in document order). Plain + text with no placeholders yields a single text block (byte-identical to the legacy behavior). + Raises on an unmatched placeholder (more tags than media files). + """ + if not _TAG_PATTERN.search(text): + return [{"type": "text", "value": text}] + + blocks: list[Content] = [] + for segment in _TAG_PATTERN.split(text): + block_type = _TAG_TO_BLOCK.get(segment) + if block_type is not None: + try: + path = next(media_iters[segment]) + except StopIteration: + raise ValueError(f"More {segment} tags than provided media files.") from None + blocks.append({"type": block_type, "value": path}) + elif segment: + blocks.append({"type": "text", "value": segment}) + return blocks + + +def _assert_media_consumed(media_iters: dict[str, Any]) -> None: + """Ensure every media file was referenced by a tag (fewer tags than media -> error).""" + for tag, media_iter in media_iters.items(): + unused = len(list(media_iter)) + if unused: + raise ValueError(f"Fewer {tag} tags than provided media files ({unused} unused).") class DataConverterPlugin(BasePlugin): @@ -76,6 +143,7 @@ def alpaca_converter(raw_sample: AlpacaSample) -> SFTSample: SFTSample: SFT sample. """ messages = [] + media_iters = _build_media_iters(raw_sample) if "system" in raw_sample: messages.append( {"role": "system", "content": [{"type": "text", "value": raw_sample["system"]}], "loss_weight": 0.0} @@ -85,9 +153,9 @@ def alpaca_converter(raw_sample: AlpacaSample) -> SFTSample: messages.append( { "role": "user", - "content": [ - {"type": "text", "value": raw_sample.get("instruction", "") + raw_sample.get("input", "")} - ], + "content": _to_content_blocks( + raw_sample.get("instruction", "") + raw_sample.get("input", ""), media_iters + ), "loss_weight": 0.0, } ) @@ -97,6 +165,7 @@ def alpaca_converter(raw_sample: AlpacaSample) -> SFTSample: {"role": "assistant", "content": [{"type": "text", "value": raw_sample["output"]}], "loss_weight": 1.0} ) + _assert_media_consumed(media_iters) return {"messages": messages} @@ -121,6 +190,7 @@ def sharegpt_converter(raw_sample: SharegptSample) -> SFTSample: } sample = {} messages = [] + media_iters = _build_media_iters(raw_sample) for message in raw_sample.get("conversations", []): tag = message["from"] if tag not in tag_mapping: @@ -146,11 +216,12 @@ def sharegpt_converter(raw_sample: SharegptSample) -> SFTSample: messages.append( { "role": tag_mapping[tag], - "content": [{"type": "text", "value": message["value"]}], + "content": _to_content_blocks(message["value"], media_iters), "loss_weight": 1.0 if tag == "gpt" else 0.0, } ) + _assert_media_consumed(media_iters) sample["messages"] = messages tools = raw_sample.get("tools") @@ -178,6 +249,8 @@ def pair_converter(raw_sample: PairSample) -> DPOSample: """ def process_message(raw_messages: list[OpenaiMessage]): + # chosen and rejected share the sample's media; each side consumes its own iterators. + media_iters = _build_media_iters(raw_sample) messages = [] for message in raw_messages: if message["role"] == "tool": @@ -201,11 +274,12 @@ def pair_converter(raw_sample: PairSample) -> DPOSample: messages.append( { "role": message["role"], - "content": [{"type": "text", "value": message["content"]}], + "content": _to_content_blocks(message["content"], media_iters), "loss_weight": 1.0 if message["role"] == "assistant" else 0.0, } ) + _assert_media_consumed(media_iters) return messages sample = {} @@ -221,3 +295,4 @@ def pair_converter(raw_sample: PairSample) -> DPOSample: logger.warning_rank0(f"Invalid tools format: {str(tools)}") return sample + diff --git a/src/llamafactory/v1/plugins/trainer_plugins/batching.py b/src/llamafactory/v1/plugins/trainer_plugins/batching.py index 4cbed315c..2425d7c19 100644 --- a/src/llamafactory/v1/plugins/trainer_plugins/batching.py +++ b/src/llamafactory/v1/plugins/trainer_plugins/batching.py @@ -1,4 +1,4 @@ -# Copyright 2025 the LlamaFactory team. +# Copyright 2026 the LlamaFactory team. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -20,8 +20,8 @@ from typing import Any import torch from torch.utils.data import default_collate +from ...core.utils.collation import pad_and_truncate from ...utils.constants import IGNORE_INDEX -from ...utils.helper import pad_and_truncate from ...utils.objects import StatefulBuffer from ...utils.plugin import BasePlugin, ensure_methods_implemented from ...utils.types import BatchInfo, BatchInput, DataLoader, ModelInput diff --git a/src/llamafactory/v1/plugins/trainer_plugins/distributed/deepspeed.py b/src/llamafactory/v1/plugins/trainer_plugins/distributed/deepspeed.py index 105478431..b7e72ddd1 100644 --- a/src/llamafactory/v1/plugins/trainer_plugins/distributed/deepspeed.py +++ b/src/llamafactory/v1/plugins/trainer_plugins/distributed/deepspeed.py @@ -34,6 +34,14 @@ from ...model_plugins.deepspeed_utils import infer_deepspeed_mixed_precision logger = get_logger(__name__) +# ZeRO-3 bucket sizes that accelerate derives from the model's hidden size +_ZERO3_BUCKET_FORMULAS = { + "reduce_bucket_size": lambda hidden: hidden * hidden, + "stage3_prefetch_bucket_size": lambda hidden: int(0.9 * hidden * hidden), + "stage3_param_persistence_threshold": lambda hidden: 10 * hidden, +} + + class DeepSpeedEngine: """DeepSpeed integration using accelerate's built-in capabilities. @@ -84,6 +92,7 @@ class DeepSpeedEngine: Internally calls deepspeed.initialize() and wraps the returned objects. """ + self._fill_zero3_bucket_sizes(model) if lr_scheduler is not None: model, optimizer, lr_scheduler = self.accelerator.prepare(model, optimizer, lr_scheduler) else: @@ -94,6 +103,23 @@ class DeepSpeedEngine: logger.info_rank0("Model, optimizer, and lr_scheduler prepared via accelerate") return model, optimizer, lr_scheduler + def _fill_zero3_bucket_sizes(self, model: HFModel) -> None: + """Fill ZeRO-3 ``auto`` bucket sizes that accelerate cannot infer for multimodal models.""" + zero_config = self.accelerator.state.deepspeed_plugin.deepspeed_config.get("zero_optimization", {}) + auto_keys = [key for key in _ZERO3_BUCKET_FORMULAS if zero_config.get(key) == "auto"] + if not auto_keys: + return + + config = model.config + text_config = config.get_text_config() if hasattr(config, "get_text_config") else config + hidden_size = getattr(text_config, "hidden_size", None) + if hidden_size is None: + return + + for key in auto_keys: + zero_config[key] = _ZERO3_BUCKET_FORMULAS[key](hidden_size) + logger.info_rank0(f"Resolved ZeRO-3 {auto_keys} from text-config hidden_size={hidden_size}.") + def backward(self, loss: torch.Tensor) -> None: """Backward pass using accelerate. @@ -108,7 +134,7 @@ class DeepSpeedEngine: """Get the global gradient norm from the DeepSpeed engine.""" engine_wrapper = getattr(self.accelerator, "deepspeed_engine_wrapped", None) if engine_wrapper is not None: - return engine_wrapper.engine.get_global_grad_norm() or 0.0 + return float(engine_wrapper.engine.get_global_grad_norm() or 0.0) return 0.0 diff --git a/src/llamafactory/v1/plugins/trainer_plugins/distributed/fsdp2.py b/src/llamafactory/v1/plugins/trainer_plugins/distributed/fsdp2.py index 4fbb5b61f..973c5054c 100644 --- a/src/llamafactory/v1/plugins/trainer_plugins/distributed/fsdp2.py +++ b/src/llamafactory/v1/plugins/trainer_plugins/distributed/fsdp2.py @@ -1,4 +1,4 @@ -# Copyright 2025 the LlamaFactory team. +# Copyright 2026 the LlamaFactory team. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -73,20 +73,50 @@ def _make_safetensor_loader(checkpoint_file: str, tensor_key: str): return _load_tensor -def get_transformer_layer_cls(model: HFModel) -> type[nn.Module] | None: +def _cast_norm_input_to_weight_dtype(module: nn.Module, args: tuple): + """forward-pre-hook: cast a norm layer's input to its weight dtype.""" + if not args: + return None + x = args[0] + weight = getattr(module, "weight", None) + if isinstance(x, torch.Tensor) and weight is not None and x.dtype != weight.dtype: + return (x.to(weight.dtype), *args[1:]) + return None + + +def _make_norms_dtype_safe(model: HFModel) -> int: + """Register the dtype-safe hook on every dtype-strict ``nn.LayerNorm`` in the model.""" + n = 0 + for module in model.modules(): + if isinstance(module, nn.LayerNorm): + module.register_forward_pre_hook(_cast_norm_input_to_weight_dtype) + n += 1 + return n + + +def get_transformer_layer_cls(model: HFModel) -> set[type[nn.Module]]: + classes: set[type[nn.Module]] = set() + for module in model.modules(): + for attr in ("layers", "blocks"): + seq = getattr(module, attr, None) + if isinstance(seq, nn.ModuleList) and len(seq) > 0: + classes.add(type(seq[0])) + if classes: + return classes + no_split_modules = getattr(model, "_no_split_modules", None) if no_split_modules: - if isinstance(no_split_modules, (list, tuple)): - for name, module in model.named_modules(): - for cls_name in no_split_modules: - if module.__class__.__name__ == cls_name: - return module.__class__ - if hasattr(model, "model") and hasattr(model.model, "layers"): - return type(model.model.layers[0]) - if hasattr(model, "layers"): - return type(model.layers[0]) + found: dict[str, type[nn.Module]] = {} + for _, module in model.named_modules(): + cls_name = module.__class__.__name__ + if cls_name in no_split_modules and cls_name not in found: + found[cls_name] = module.__class__ + if len(found) == len(no_split_modules): + break + if found: + return set(found.values()) - return None + return set() def save_model(model: HFModel, output_dir: str, processor: Processor) -> None: @@ -196,16 +226,15 @@ class FSDP2Engine: return model mp_policy = self.get_mp_policy() - layer_cls = get_transformer_layer_cls(model) + transformer_layer_cls_to_wrap = get_transformer_layer_cls(model) - if layer_cls is None: + if not transformer_layer_cls_to_wrap: logger.warning( "Could not identify Transformer Layer class, applying FSDP to the whole model structure only." ) - transformer_layer_cls_to_wrap = set() else: - logger.info(f"Applying per-layer FSDP to {layer_cls.__name__}") - transformer_layer_cls_to_wrap = {layer_cls} + names = ", ".join(cls.__name__ for cls in transformer_layer_cls_to_wrap) + logger.info(f"Applying per-layer FSDP to: {names}") if self.is_lora_module_wrap(model): lora_modules = [] @@ -259,6 +288,11 @@ class FSDP2Engine: model.get_input_embeddings().register_forward_hook(make_inputs_require_grad) + if self.mixed_precision == "bf16": + n_patched = _make_norms_dtype_safe(model) + if self.rank == 0 and n_patched: + logger.info(f"Made {n_patched} nn.LayerNorm(s) dtype-safe for bf16 checkpointing.") + fully_shard( model, mesh=self.fsdp_mesh, diff --git a/src/llamafactory/v1/utils/constants.py b/src/llamafactory/v1/utils/constants.py index 9ec68b44d..caf2715d0 100644 --- a/src/llamafactory/v1/utils/constants.py +++ b/src/llamafactory/v1/utils/constants.py @@ -1,4 +1,4 @@ -# Copyright 2025 the LlamaFactory team. +# Copyright 2026 the LlamaFactory team. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -12,4 +12,10 @@ # See the License for the specific language governing permissions and # limitations under the License. +import os + + IGNORE_INDEX = -100 +IMAGE_PLACEHOLDER = os.getenv("IMAGE_PLACEHOLDER", "") +VIDEO_PLACEHOLDER = os.getenv("VIDEO_PLACEHOLDER", "