[v1] Support multimodal data training (#10656)

Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
浮梦
2026-07-31 18:54:13 +08:00
committed by GitHub
parent 3984675dd5
commit 62ae362455
18 changed files with 1304 additions and 122 deletions

View File

@@ -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日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]}

View File

@@ -0,0 +1,4 @@
multimodal_demo:
path: data/v1_multimodal_demo.jsonl
source: local

View File

@@ -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

View File

@@ -43,7 +43,7 @@ from ..utils.callbacks import (
TrainerCallback, TrainerCallback,
TrainerState, 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 ..utils.types import BatchInput, HFModel, ModelOutput, Tensor, TorchDataset
from .rendering import Renderer from .rendering import Renderer
from .utils.batching import BatchGenerator from .utils.batching import BatchGenerator
@@ -75,6 +75,7 @@ class BaseTrainer:
self.dp_size = DistributedInterface().get_world_size(Dim.DP) self.dp_size = DistributedInterface().get_world_size(Dim.DP)
self.cp_size = DistributedInterface().get_world_size(Dim.CP) self.cp_size = DistributedInterface().get_world_size(Dim.CP)
self.model_input_names = self.renderer.processor.model_input_names self.model_input_names = self.renderer.processor.model_input_names
self._uses_mrope = model_uses_mrope(self.model.config)
self._create_batch_generator() self._create_batch_generator()
# Calculate num_training_steps: max_steps takes priority if set # Calculate num_training_steps: max_steps takes priority if set
@@ -89,6 +90,9 @@ class BaseTrainer:
if self.args.enable_activation_checkpointing: if self.args.enable_activation_checkpointing:
self.model.gradient_checkpointing_enable({"use_reentrant": False}) 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 self._deepspeed_engine = None
dist_name = self.args.dist_config.name if self.args.dist_config is not None else 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." "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] 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: else:
from ..plugins.trainer_plugins.distributed.interface import DistributedPlugin from ..plugins.trainer_plugins.distributed.interface import DistributedPlugin
@@ -224,6 +232,9 @@ class BaseTrainer:
model_inputs = { model_inputs = {
k: v.to(self.device, non_blocking=True) for k, v in batch.items() if isinstance(v, torch.Tensor) 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) labels = batch["labels"].to(self.device, non_blocking=True)
outputs: ModelOutput = model(**model_inputs) outputs: ModelOutput = model(**model_inputs)
logits = outputs.logits.float() logits = outputs.logits.float()

View File

@@ -150,8 +150,19 @@ class ModelEngine:
if self.args.model_class == ModelClass.LLM: if self.args.model_class == ModelClass.LLM:
from transformers import AutoModelForCausalLM, AutoModelForImageTextToText 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 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: else:
AutoClass = AutoModelForCausalLM 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" init_mode = self.args.init_config.name if self.args.init_config is not None else "init_on_default"
model._init_mode = init_mode 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.args.peft_config is None:
if self.is_train: if self.is_train:
logger.info_rank0("Fine-tuning mode: full tuning") logger.info_rank0("Fine-tuning mode: full tuning")

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License. # you may not use this file except in compliance with the License.
@@ -14,13 +14,15 @@
"""Message <-> HF-template plumbing for rendering. """Message <-> HF-template plumbing for rendering.
Pure, stateless helpers: convert v1 ``Message`` to HF chat-template format. No tokenization policy Pure, stateless helpers: convert v1 ``Message`` to HF chat-template format, extract/count media, and
decisions live here -- only mechanical conversion used by ``rendering.py``. guard media placeholder counts. No tokenization policy decisions live here -- only mechanical
conversion used by ``rendering.py``.
""" """
import json import json
from ...utils.types import Message from ...utils.helper import get_tokenizer
from ...utils.types import Message, Processor
_FALLBACK_CHATML_JINJA = ( _FALLBACK_CHATML_JINJA = (
@@ -33,13 +35,40 @@ _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.""" """Convert v1 Message format to HF format for apply_chat_template."""
hf_messages = [] hf_messages = []
for message in messages: for message in messages:
tool_calls: list[dict] = [] tool_calls: list[dict] = []
reasoning_content = "" reasoning_content = ""
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 = "" text = ""
for content in message["content"]: for content in message["content"]:
if content["type"] == "text": if content["type"] == "text":
@@ -52,8 +81,12 @@ def _to_hf_messages(messages: list[Message]) -> list[dict]:
except json.JSONDecodeError as e: except json.JSONDecodeError as e:
raise ValueError(f"tool_call value is not valid JSON: {content['value']!r}") from 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: 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}") raise ValueError(
tool_calls.append({"type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}}) 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} hf_msg = {"role": message["role"], "content": text}
if tool_calls: if tool_calls:
@@ -63,3 +96,75 @@ def _to_hf_messages(messages: list[Message]) -> list[dict]:
hf_messages.append(hf_msg) hf_messages.append(hf_msg)
return hf_messages 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."
)

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with 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 - ``format`` -- v1<->HF message conversion
- ``escape`` -- special-token escaping (prompt-injection hardening) - ``escape`` -- special-token escaping (prompt-injection hardening)
Assistant supervision is located WITHOUT a per-model marker table: a training sample is rendered Note: ``position_ids`` are assigned by ``process_samples`` (1-based); multimodal (mrope) position
so that its last message is the supervised assistant turn, and that turn's token span is recovered ids are expected to be recomputed by the model/trainer.
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 ``<think>``) are handled correctly without hard-coding role markers.
""" """
import json import json
import numpy as np
import torch
from ...utils.constants import IGNORE_INDEX 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.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 .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( def _render_messages(
@@ -46,20 +51,23 @@ def _render_messages(
is_generate: bool = False, is_generate: bool = False,
**kwargs, **kwargs,
) -> ModelInput: ) -> 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. Note: ``position_ids`` are not produced here; ``process_samples`` assigns a 1-based range.
""" """
tokenizer = get_tokenizer(processor) tokenizer = get_tokenizer(processor)
if not getattr(tokenizer, "chat_template", None): is_multimodal = not is_tokenizer(processor)
tokenizer.chat_template = _FALLBACK_CHATML_JINJA
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). # 0. Neutralize special-token strings in user-controlled text (no-op for normal data).
specials = _special_token_strings(tokenizer) specials = _special_token_strings(tokenizer)
special_ids = {tid for tid, t in tokenizer.added_tokens_decoder.items() if getattr(t, "special", False)} 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) 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 tools_parsed = None
if tools: if tools:
@@ -70,36 +78,76 @@ def _render_messages(
raise ValueError(f"tools is not valid JSON: {tools!r}") from e raise ValueError(f"tools is not valid JSON: {tools!r}") from e
if not isinstance(tools_parsed, list): if not isinstance(tools_parsed, list):
tools_parsed = [tools_parsed] 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]: if not is_generate and hf_messages and hf_messages[-1]["role"] == "assistant":
text = tokenizer.apply_chat_template( kwargs["enable_thinking"] = bool(hf_messages[-1].get("reasoning_content"))
msgs, tokenize=False, add_generation_prompt=add_generation_prompt, tools=tools_parsed, **kwargs
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. # 1. Full sequence (used verbatim), plus its multimodal feature outputs.
input_ids = _encode(hf_messages, add_generation_prompt=is_generate) input_ids, outputs = _encode(hf_messages, messages, add_generation_prompt=is_generate)
n = len(input_ids) 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: if is_generate:
# Generation prompt only -- nothing is supervised. # Generation prompt only -- nothing is supervised.
return ModelInput( result = ModelInput(
input_ids=input_ids, input_ids=input_ids,
attention_mask=[1] * n, attention_mask=[1] * n,
labels=[IGNORE_INDEX] * n, labels=[IGNORE_INDEX] * n,
loss_weights=[0.0] * 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": if not messages or messages[-1]["role"] != "assistant":
raise ValueError( raise ValueError(
"training render expects the last message to be the supervised assistant turn; " "training render expects the last message to be the supervised assistant turn; "
"multi-turn conversations are split per turn in process_samples." "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: 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 # 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 # 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) labels.append(tid if supervised else IGNORE_INDEX)
loss_weights.append(weight) loss_weights.append(weight)
return ModelInput( result = ModelInput(
input_ids=input_ids, input_ids=input_ids,
attention_mask=[1] * n, attention_mask=[1] * n,
labels=labels, labels=labels,
loss_weights=loss_weights, loss_weights=loss_weights,
) )
_attach_multimodal(result)
return result
class Renderer: 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 self.processor = processor
def render_messages( def render_messages(
@@ -152,6 +205,61 @@ class Renderer:
""" """
return _render_messages(self.processor, messages, tools, is_generate, **kwargs) 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]: def process_samples(self, samples: list[Sample]) -> list[ModelInput]:
"""Process samples to model input. """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( model_input["position_ids"] = list(range(1, len(chosen_input["input_ids"]) + 1)) + list(
range(1, len(rejected_input["input_ids"]) + 1) 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) rendered.append(model_input)
else: else:
raise ValueError("No valid messages or chosen_messages/rejected_messages found in sample.") raise ValueError("No valid messages or chosen_messages/rejected_messages found in sample.")

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with 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 import StatefulDataLoader
from torchdata.stateful_dataloader.sampler import StatefulDistributedSampler from torchdata.stateful_dataloader.sampler import StatefulDistributedSampler
from ...accelerator.helper import ReduceOp
from ...accelerator.interface import Dim, DistributedInterface from ...accelerator.interface import Dim, DistributedInterface
from ...config import BatchingStrategy from ...config import BatchingStrategy
from ...utils import logging 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.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 ..rendering import Renderer
from .collation import _MULTIMODAL_PASSTHROUGH_KEYS, pad_and_truncate
logger = logging.get_logger(__name__) logger = logging.get_logger(__name__)
@@ -45,7 +48,82 @@ logger = logging.get_logger(__name__)
__all__ = ["BatchGenerator"] __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"] micro_batch_size = batch_info["micro_batch_size"]
num_micro_batch = batch_info["num_micro_batch"] num_micro_batch = batch_info["num_micro_batch"]
cutoff_len = batch_info["cutoff_len"] cutoff_len = batch_info["cutoff_len"]
@@ -54,10 +132,24 @@ def default_collate_fn(buffer: StatefulBuffer, batch_info: BatchInfo) -> list[Ba
return None return None
samples = buffer.get(batch_size) samples = buffer.get(batch_size)
batch = [] micro_batches = [samples[i * micro_batch_size : (i + 1) * micro_batch_size] for i in range(num_micro_batch)]
for i in range(num_micro_batch):
micro_batch = samples[i * micro_batch_size : (i + 1) * micro_batch_size] # Collate first; presence is judged on the *post-truncation* result, since truncation can
batch.append(default_collate(pad_and_truncate(micro_batch, cutoff_len))) # 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 return batch
@@ -227,8 +319,17 @@ class BatchGenerator(Iterator):
def _generate_batch(self) -> list[BatchInput] | None: def _generate_batch(self) -> list[BatchInput] | None:
if self.batching_strategy == BatchingStrategy.NORMAL: 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: 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 from ...plugins.trainer_plugins.batching import BatchingPlugin
return BatchingPlugin(self.batching_strategy).generate_batch(self._buffer, self._batch_info) return BatchingPlugin(self.batching_strategy).generate_batch(self._buffer, self._batch_info)

View File

@@ -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

View File

@@ -14,11 +14,13 @@
import json import json
import re
from typing import Any, Literal, NotRequired, TypedDict from typing import Any, Literal, NotRequired, TypedDict
from ...utils import logging from ...utils import logging
from ...utils.constants import AUDIO_PLACEHOLDER, IMAGE_PLACEHOLDER, VIDEO_PLACEHOLDER
from ...utils.plugin import BasePlugin 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__) logger = logging.get_logger(__name__)
@@ -29,6 +31,9 @@ class AlpacaSample(TypedDict, total=False):
instruction: str instruction: str
input: NotRequired[str] input: NotRequired[str]
output: str output: str
images: NotRequired[list[str] | str]
videos: NotRequired[list[str] | str]
audios: NotRequired[list[str] | str]
SharegptMessage = TypedDict( SharegptMessage = TypedDict(
@@ -40,6 +45,9 @@ SharegptMessage = TypedDict(
class SharegptSample(TypedDict, total=False): class SharegptSample(TypedDict, total=False):
conversations: list[SharegptMessage] conversations: list[SharegptMessage]
tools: NotRequired[str] tools: NotRequired[str]
images: NotRequired[list[str] | str]
videos: NotRequired[list[str] | str]
audios: NotRequired[list[str] | str]
class OpenaiMessage(TypedDict, total=False): class OpenaiMessage(TypedDict, total=False):
@@ -54,6 +62,65 @@ class OpenaiSample(TypedDict, total=False):
class PairSample(TypedDict, total=False): class PairSample(TypedDict, total=False):
chosen: list[OpenaiMessage] chosen: list[OpenaiMessage]
rejected: 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): class DataConverterPlugin(BasePlugin):
@@ -76,6 +143,7 @@ def alpaca_converter(raw_sample: AlpacaSample) -> SFTSample:
SFTSample: SFT sample. SFTSample: SFT sample.
""" """
messages = [] messages = []
media_iters = _build_media_iters(raw_sample)
if "system" in raw_sample: if "system" in raw_sample:
messages.append( messages.append(
{"role": "system", "content": [{"type": "text", "value": raw_sample["system"]}], "loss_weight": 0.0} {"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( messages.append(
{ {
"role": "user", "role": "user",
"content": [ "content": _to_content_blocks(
{"type": "text", "value": raw_sample.get("instruction", "") + raw_sample.get("input", "")} raw_sample.get("instruction", "") + raw_sample.get("input", ""), media_iters
], ),
"loss_weight": 0.0, "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} {"role": "assistant", "content": [{"type": "text", "value": raw_sample["output"]}], "loss_weight": 1.0}
) )
_assert_media_consumed(media_iters)
return {"messages": messages} return {"messages": messages}
@@ -121,6 +190,7 @@ def sharegpt_converter(raw_sample: SharegptSample) -> SFTSample:
} }
sample = {} sample = {}
messages = [] messages = []
media_iters = _build_media_iters(raw_sample)
for message in raw_sample.get("conversations", []): for message in raw_sample.get("conversations", []):
tag = message["from"] tag = message["from"]
if tag not in tag_mapping: if tag not in tag_mapping:
@@ -146,11 +216,12 @@ def sharegpt_converter(raw_sample: SharegptSample) -> SFTSample:
messages.append( messages.append(
{ {
"role": tag_mapping[tag], "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, "loss_weight": 1.0 if tag == "gpt" else 0.0,
} }
) )
_assert_media_consumed(media_iters)
sample["messages"] = messages sample["messages"] = messages
tools = raw_sample.get("tools") tools = raw_sample.get("tools")
@@ -178,6 +249,8 @@ def pair_converter(raw_sample: PairSample) -> DPOSample:
""" """
def process_message(raw_messages: list[OpenaiMessage]): 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 = [] messages = []
for message in raw_messages: for message in raw_messages:
if message["role"] == "tool": if message["role"] == "tool":
@@ -201,11 +274,12 @@ def pair_converter(raw_sample: PairSample) -> DPOSample:
messages.append( messages.append(
{ {
"role": message["role"], "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, "loss_weight": 1.0 if message["role"] == "assistant" else 0.0,
} }
) )
_assert_media_consumed(media_iters)
return messages return messages
sample = {} sample = {}
@@ -221,3 +295,4 @@ def pair_converter(raw_sample: PairSample) -> DPOSample:
logger.warning_rank0(f"Invalid tools format: {str(tools)}") logger.warning_rank0(f"Invalid tools format: {str(tools)}")
return sample return sample

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License. # you may not use this file except in compliance with the License.
@@ -20,8 +20,8 @@ from typing import Any
import torch import torch
from torch.utils.data import default_collate from torch.utils.data import default_collate
from ...core.utils.collation import pad_and_truncate
from ...utils.constants import IGNORE_INDEX from ...utils.constants import IGNORE_INDEX
from ...utils.helper import pad_and_truncate
from ...utils.objects import StatefulBuffer from ...utils.objects import StatefulBuffer
from ...utils.plugin import BasePlugin, ensure_methods_implemented from ...utils.plugin import BasePlugin, ensure_methods_implemented
from ...utils.types import BatchInfo, BatchInput, DataLoader, ModelInput from ...utils.types import BatchInfo, BatchInput, DataLoader, ModelInput

View File

@@ -34,6 +34,14 @@ from ...model_plugins.deepspeed_utils import infer_deepspeed_mixed_precision
logger = get_logger(__name__) 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: class DeepSpeedEngine:
"""DeepSpeed integration using accelerate's built-in capabilities. """DeepSpeed integration using accelerate's built-in capabilities.
@@ -84,6 +92,7 @@ class DeepSpeedEngine:
Internally calls deepspeed.initialize() and wraps the returned objects. Internally calls deepspeed.initialize() and wraps the returned objects.
""" """
self._fill_zero3_bucket_sizes(model)
if lr_scheduler is not None: if lr_scheduler is not None:
model, optimizer, lr_scheduler = self.accelerator.prepare(model, optimizer, lr_scheduler) model, optimizer, lr_scheduler = self.accelerator.prepare(model, optimizer, lr_scheduler)
else: else:
@@ -94,6 +103,23 @@ class DeepSpeedEngine:
logger.info_rank0("Model, optimizer, and lr_scheduler prepared via accelerate") logger.info_rank0("Model, optimizer, and lr_scheduler prepared via accelerate")
return model, optimizer, lr_scheduler 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: def backward(self, loss: torch.Tensor) -> None:
"""Backward pass using accelerate. """Backward pass using accelerate.
@@ -108,7 +134,7 @@ class DeepSpeedEngine:
"""Get the global gradient norm from the DeepSpeed engine.""" """Get the global gradient norm from the DeepSpeed engine."""
engine_wrapper = getattr(self.accelerator, "deepspeed_engine_wrapped", None) engine_wrapper = getattr(self.accelerator, "deepspeed_engine_wrapped", None)
if engine_wrapper is not 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 return 0.0

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with 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 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) no_split_modules = getattr(model, "_no_split_modules", None)
if no_split_modules: if no_split_modules:
if isinstance(no_split_modules, (list, tuple)): found: dict[str, type[nn.Module]] = {}
for name, module in model.named_modules(): for _, module in model.named_modules():
for cls_name in no_split_modules: cls_name = module.__class__.__name__
if module.__class__.__name__ == cls_name: if cls_name in no_split_modules and cls_name not in found:
return module.__class__ found[cls_name] = module.__class__
if hasattr(model, "model") and hasattr(model.model, "layers"): if len(found) == len(no_split_modules):
return type(model.model.layers[0]) break
if hasattr(model, "layers"): if found:
return type(model.layers[0]) return set(found.values())
return None return set()
def save_model(model: HFModel, output_dir: str, processor: Processor) -> None: def save_model(model: HFModel, output_dir: str, processor: Processor) -> None:
@@ -196,16 +226,15 @@ class FSDP2Engine:
return model return model
mp_policy = self.get_mp_policy() 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( logger.warning(
"Could not identify Transformer Layer class, applying FSDP to the whole model structure only." "Could not identify Transformer Layer class, applying FSDP to the whole model structure only."
) )
transformer_layer_cls_to_wrap = set()
else: else:
logger.info(f"Applying per-layer FSDP to {layer_cls.__name__}") names = ", ".join(cls.__name__ for cls in transformer_layer_cls_to_wrap)
transformer_layer_cls_to_wrap = {layer_cls} logger.info(f"Applying per-layer FSDP to: {names}")
if self.is_lora_module_wrap(model): if self.is_lora_module_wrap(model):
lora_modules = [] lora_modules = []
@@ -259,6 +288,11 @@ class FSDP2Engine:
model.get_input_embeddings().register_forward_hook(make_inputs_require_grad) 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( fully_shard(
model, model,
mesh=self.fsdp_mesh, mesh=self.fsdp_mesh,

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with 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 # See the License for the specific language governing permissions and
# limitations under the License. # limitations under the License.
import os
IGNORE_INDEX = -100 IGNORE_INDEX = -100
IMAGE_PLACEHOLDER = os.getenv("IMAGE_PLACEHOLDER", "<image>")
VIDEO_PLACEHOLDER = os.getenv("VIDEO_PLACEHOLDER", "<video>")
AUDIO_PLACEHOLDER = os.getenv("AUDIO_PLACEHOLDER", "<audio>")

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License. # you may not use this file except in compliance with the License.
@@ -23,7 +23,7 @@ from transformers import set_seed as hf_set_seed
from ..accelerator.helper import is_torch_npu_available from ..accelerator.helper import is_torch_npu_available
from ..accelerator.interface import DistributedInterface from ..accelerator.interface import DistributedInterface
from .constants import IGNORE_INDEX from .constants import IGNORE_INDEX
from .types import BatchInput, ModelInput, Processor, Tensor from .types import BatchInput, Processor
def enable_full_determinism(seed: int) -> None: def enable_full_determinism(seed: int) -> None:
@@ -79,37 +79,6 @@ def get_tokenizer(processor: Processor) -> PreTrainedTokenizer:
return processor.tokenizer if hasattr(processor, "tokenizer") else processor return processor.tokenizer if hasattr(processor, "tokenizer") else processor
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 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:
padded_sample = {}
for key, value in sample.items():
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
def compute_valid_tokens(batches: list[BatchInput]) -> int: def compute_valid_tokens(batches: list[BatchInput]) -> int:
"""Compute valid tokens in batches. """Compute valid tokens in batches.
@@ -125,3 +94,15 @@ def compute_valid_tokens(batches: list[BatchInput]) -> int:
for batch in batches for batch in batches
if "labels" in batch if "labels" in batch
) )
def model_uses_mrope(config) -> bool:
"""Whether the model uses multimodal RoPE (3D position ids built from grid_thw).
Detected from the (text) config's rope settings carrying an ``mrope_section`` (Qwen2.5-VL /
Qwen3-VL / Qwen3.5 family). Such models compute their own multimodal position ids inside
``forward`` when ``position_ids`` is not provided.
"""
text_config = getattr(config, "text_config", config)
rope = getattr(text_config, "rope_scaling", None) or getattr(text_config, "rope_parameters", None)
return isinstance(rope, dict) and "mrope_section" in rope

View File

@@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team. # Copyright 2026 the LlamaFactory team.
# #
# Licensed under the Apache License, Version 2.0 (the "License"); # Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License. # you may not use this file except in compliance with the License.
@@ -141,6 +141,20 @@ class ModelInput(TypedDict, total=False):
"""Position ids for the model (optional).""" """Position ids for the model (optional)."""
token_type_ids: NotRequired[list[int]] token_type_ids: NotRequired[list[int]]
"""Token type ids used in DPO, 1 represents the chosen messages, 2 represents the rejected messages.""" """Token type ids used in DPO, 1 represents the chosen messages, 2 represents the rejected messages."""
pixel_values: NotRequired[Any]
"""Pixel values for vision models."""
image_grid_thw: NotRequired[Any]
"""Image grid (temporal, height, width) for vision models."""
pixel_values_videos: NotRequired[Any]
"""Pixel values for video inputs."""
video_grid_thw: NotRequired[Any]
"""Video grid (temporal, height, width) for video models."""
input_features: NotRequired[Any]
"""Audio input features (e.g. mel spectrogram) for audio models."""
feature_attention_mask: NotRequired[Any]
"""Attention mask over the audio input features."""
mm_token_type_ids: NotRequired[list[int]]
"""Multimodal token type ids: 0=text, 1=image, 2=video, 3=audio."""
class BatchInput(TypedDict, total=False): class BatchInput(TypedDict, total=False):
@@ -156,6 +170,20 @@ class BatchInput(TypedDict, total=False):
"""Position ids for the model (optional).""" """Position ids for the model (optional)."""
token_type_ids: NotRequired[Tensor] token_type_ids: NotRequired[Tensor]
"""Token type ids used in DPO, 1 represents the chosen messages, 2 represents the rejected messages.""" """Token type ids used in DPO, 1 represents the chosen messages, 2 represents the rejected messages."""
pixel_values: NotRequired[Tensor]
"""Pixel values for vision models."""
image_grid_thw: NotRequired[Tensor]
"""Image grid (temporal, height, width) for vision models."""
pixel_values_videos: NotRequired[Tensor]
"""Pixel values for video inputs."""
video_grid_thw: NotRequired[Tensor]
"""Video grid (temporal, height, width) for video models."""
input_features: NotRequired[Tensor]
"""Audio input features (e.g. mel spectrogram) for audio models."""
feature_attention_mask: NotRequired[Tensor]
"""Attention mask over the audio input features."""
mm_token_type_ids: NotRequired[Tensor]
"""Multimodal token type ids: 0=text, 1=image, 2=video, 3=audio."""
class BatchInfo(TypedDict): class BatchInfo(TypedDict):

View File

@@ -370,3 +370,208 @@ def test_dynamic_padding_free_fill_buffer_restarts_until_micro_batch_is_complete
assert len(batch) == 1 assert len(batch) == 1
assert batch[0]["input_ids"].shape == (1, 18) assert batch[0]["input_ids"].shape == (1, 18)
assert len(batch_generator._buffer) == 1 assert len(batch_generator._buffer) == 1
def _image_fragment(n_pad: int = 4, merge_sq: int = 4):
"""Hand-crafted image fragment: vision_start + n_pad image_pad + vision_end."""
import torch
pad, vstart, vend = 9, 8, 7
return {
"input_ids": [vstart] + [pad] * n_pad + [vend],
"mm_token_type_ids": [0] + [1] * n_pad + [0],
"pixel_values": torch.zeros((n_pad * merge_sq, 16), dtype=torch.float32),
"image_grid_thw": torch.tensor([[1, 2, n_pad * 2]], dtype=torch.long),
}
def _text_sample(n: int, base: int = 100):
s = _make_model_input(n, start=base)
s["position_ids"] = list(range(1, n + 1))
return s
def test_inject_appends_zero_loss_dummy_into_collated_text_batch():
import torch
from llamafactory.v1.core.utils.batching import _collate_micro_batch, _inject_dummy_into_collated
collated = _collate_micro_batch([_text_sample(20), _text_sample(8)], cutoff_len=4096)
assert "pixel_values" not in collated
bsz, seqlen = collated["input_ids"].shape
frag = _image_fragment(n_pad=4)
fl = len(frag["input_ids"])
_inject_dummy_into_collated(collated, frag, marker=1)
new_len = seqlen + fl
# every sequence field grew by the fragment length, batch size unchanged
for key in ("input_ids", "attention_mask", "labels", "loss_weights", "position_ids", "mm_token_type_ids"):
assert collated[key].shape == (bsz, new_len)
# dummy lives only in row 0's tail; other rows are padding (attention 0) there
assert collated["input_ids"][0, seqlen:].tolist() == frag["input_ids"]
assert collated["attention_mask"][0, seqlen:].tolist() == [1] * fl
assert collated["attention_mask"][1, seqlen:].tolist() == [0] * fl
# zero loss contribution
assert collated["labels"][0, seqlen:].tolist() == [IGNORE_INDEX] * fl
assert torch.all(collated["loss_weights"][:, seqlen:] == 0.0)
assert collated["mm_token_type_ids"][0, seqlen:].tolist() == frag["mm_token_type_ids"]
# pixel features carried verbatim
assert torch.equal(collated["pixel_values"], frag["pixel_values"])
assert torch.equal(collated["image_grid_thw"], frag["image_grid_thw"])
def test_inject_video_concatenates_alongside_existing_image():
"""Injecting a missing modality leaves the other modality's features intact (dim-0 cat)."""
import torch
from llamafactory.v1.core.utils.batching import _collate_micro_batch, _inject_dummy_into_collated
img = _text_sample(10)
img["pixel_values"] = torch.ones((8, 16), dtype=torch.float32)
img["image_grid_thw"] = torch.tensor([[1, 2, 4]], dtype=torch.long)
img["mm_token_type_ids"] = [0] * 10
collated = _collate_micro_batch([img], cutoff_len=4096)
video_frag = {
"input_ids": [8, 6, 6, 7],
"mm_token_type_ids": [0, 2, 2, 0],
"pixel_values_videos": torch.zeros((8, 16), dtype=torch.float32),
"video_grid_thw": torch.tensor([[1, 2, 4]], dtype=torch.long),
}
_inject_dummy_into_collated(collated, video_frag, marker=2)
# image features untouched, video features added
assert torch.equal(collated["pixel_values"], torch.ones((8, 16)))
assert collated["pixel_values_videos"].shape[0] == 8
assert collated["video_grid_thw"].shape[0] == 1
assert collated["mm_token_type_ids"][0, -4:].tolist() == [0, 2, 2, 0]
def test_collate_creates_mm_token_type_ids_for_pure_text_then_inject():
"""A pure-text micro batch has no mm_token_type_ids; injection must create it."""
from llamafactory.v1.core.utils.batching import _collate_micro_batch, _inject_dummy_into_collated
collated = _collate_micro_batch([_text_sample(12)], cutoff_len=4096)
assert "mm_token_type_ids" not in collated
seqlen = collated["input_ids"].shape[1]
frag = _image_fragment(n_pad=3)
_inject_dummy_into_collated(collated, frag, marker=1)
assert "mm_token_type_ids" in collated
assert collated["mm_token_type_ids"].shape == collated["input_ids"].shape
# original region all zero (text), dummy region carries the markers
assert collated["mm_token_type_ids"][0, :seqlen].tolist() == [0] * seqlen
assert collated["mm_token_type_ids"][0, seqlen:].tolist() == frag["mm_token_type_ids"]
def _audio_fragment(n_tok: int = 2, n_frames: int = 3000):
"""Hand-crafted audio fragment: audio_bos + n_tok AUDIO + audio_eos, with feature rows."""
import torch
aud, bos, eos = 50, 51, 52
return {
"input_ids": [bos] + [aud] * n_tok + [eos],
"mm_token_type_ids": [0] + [3] * n_tok + [0],
"input_features": torch.zeros((1, 128, n_frames), dtype=torch.float32),
"feature_attention_mask": torch.ones((1, n_frames), dtype=torch.long),
}
def test_inject_audio_dummy_into_text_batch():
"""A pure-text micro batch gets an audio dummy appended so the audio tower fires on every rank."""
import torch
from llamafactory.v1.core.utils.batching import _collate_micro_batch, _inject_dummy_into_collated
collated = _collate_micro_batch([_text_sample(12)], cutoff_len=4096)
assert "input_features" not in collated
seqlen = collated["input_ids"].shape[1]
frag = _audio_fragment(n_tok=2)
fl = len(frag["input_ids"])
_inject_dummy_into_collated(collated, frag, marker=3)
# audio feature tensors carried verbatim; placeholder tokens marked 3 in the dummy tail
assert torch.equal(collated["input_features"], frag["input_features"])
assert torch.equal(collated["feature_attention_mask"], frag["feature_attention_mask"])
assert collated["mm_token_type_ids"][0, seqlen:].tolist() == frag["mm_token_type_ids"]
# zero loss contribution from the dummy
assert collated["labels"][0, seqlen:].tolist() == [IGNORE_INDEX] * fl
assert torch.all(collated["loss_weights"][:, seqlen:] == 0.0)
def test_audio_truncation_drops_orphaned_item_and_zeros_tokens():
"""Truncating mid-audio trims the orphaned feature row and zeros its in-window tokens."""
import torch
from llamafactory.v1.core.utils.collation import _align_multimodal_on_truncation
aud = 50
# text(2) + [audio#0: 4 tok] + text(1) + [audio#1: 4 tok] + text(1)
input_ids = [1, 2] + [aud] * 4 + [3] + [aud] * 4 + [4]
mm = [0, 0] + [3] * 4 + [0] + [3] * 4 + [0]
sample = {
"input_ids": input_ids,
"labels": input_ids.copy(),
"loss_weights": [1.0] * len(input_ids),
"mm_token_type_ids": mm,
"input_features": torch.zeros((2, 128, 10), dtype=torch.float32),
"feature_attention_mask": torch.ones((2, 10), dtype=torch.long),
}
# audio#1 occupies positions 7..10; cut at 9 so its last token (10) is orphaned, audio#0 intact
out = _align_multimodal_on_truncation(dict(sample), max_length=9)
assert out["input_features"].shape[0] == 1 # only the complete audio#0 survives
assert out["feature_attention_mask"].shape[0] == 1
# audio#0 tokens (positions 2..5) untouched
assert all(out["input_ids"][i] == aud and out["mm_token_type_ids"][i] == 3 for i in range(2, 6))
# audio#1's in-window tokens (positions 7,8) zeroed + delabeled (positions >= 9 cut by truncation)
for i in (7, 8):
assert out["input_ids"][i] == 0
assert out["mm_token_type_ids"][i] == 0
assert out["labels"][i] == IGNORE_INDEX
assert out["loss_weights"][i] == 0.0
def test_audio_truncation_keeps_all_when_complete():
"""No trimming when the cut falls after every audio's last token."""
import torch
from llamafactory.v1.core.utils.collation import _align_multimodal_on_truncation
aud = 50
input_ids = [1] + [aud] * 4 + [2]
sample = {
"input_ids": input_ids,
"labels": input_ids.copy(),
"loss_weights": [1.0] * len(input_ids),
"mm_token_type_ids": [0] + [3] * 4 + [0],
"input_features": torch.zeros((1, 128, 10), dtype=torch.float32),
"feature_attention_mask": torch.ones((1, 10), dtype=torch.long),
}
out = _align_multimodal_on_truncation(dict(sample), max_length=6)
assert out["input_features"].shape[0] == 1
assert out["input_ids"] == input_ids
def test_drop_unsupervised_samples():
"""Samples whose supervised tokens fall entirely beyond cutoff_len are dropped (warn once)."""
from types import SimpleNamespace
def _s(weights): # a sample's input_ids length matches its loss_weights length
return {"input_ids": list(range(len(weights))), "loss_weights": weights}
gen = SimpleNamespace(cutoff_len=4, _warned_truncation=False)
samples = [
_s([0.0, 0.0, 1.0, 1.0]), # fits cutoff (len 4), supervised -> kept
_s([0.0, 0.0, 0.0, 0.0, 1.0, 1.0]), # len 6 > 4, supervision only beyond cutoff -> dropped
_s([1.0, 1.0]), # short, fully supervised -> kept
_s([0.0, 0.0, 1.0, 1.0, 1.0, 1.0]), # len 6 > 4 but supervision within cutoff -> kept
]
kept = BatchGenerator._drop_unsupervised(gen, samples)
assert kept == [samples[0], samples[2], samples[3]]
assert gen._warned_truncation is True

View File

@@ -71,6 +71,148 @@ def test_sharegpt_converter():
assert DataConverterPlugin("sharegpt")(example) == expected_data assert DataConverterPlugin("sharegpt")(example) == expected_data
def test_sharegpt_converter_multimodal():
example = {
"conversations": [
{"from": "human", "value": "What is <image> and what happens in <video>?"},
{"from": "gpt", "value": "An image and a video."},
],
"images": ["/p/a.jpg"],
"videos": ["/p/v.mp4"],
}
expected_data = {
"messages": [
{
"role": "user",
"content": [
{"type": "text", "value": "What is "},
{"type": "image_url", "value": "/p/a.jpg"},
{"type": "text", "value": " and what happens in "},
{"type": "video_url", "value": "/p/v.mp4"},
{"type": "text", "value": "?"},
],
"loss_weight": 0.0,
},
{"role": "assistant", "content": [{"type": "text", "value": "An image and a video."}], "loss_weight": 1.0},
]
}
assert DataConverterPlugin("sharegpt")(example) == expected_data
def test_sharegpt_converter_multiple_images_in_order():
# images are a sample-level list consumed by <image> tags in document order across turns
example = {
"conversations": [
{"from": "human", "value": "<image><image>Compare these."},
{"from": "gpt", "value": "Done."},
],
"images": ["/p/a.jpg", "/p/b.jpg"],
}
user = DataConverterPlugin("sharegpt")(example)["messages"][0]
assert user["content"] == [
{"type": "image_url", "value": "/p/a.jpg"},
{"type": "image_url", "value": "/p/b.jpg"},
{"type": "text", "value": "Compare these."},
]
def test_sharegpt_converter_no_media_unchanged():
# backward compatibility: a scalar (non-list) image column and no tags is normalized; with no
# media columns at all the output is byte-identical to the text-only path.
example = {"conversations": [{"from": "human", "value": "hi"}, {"from": "gpt", "value": "yo"}]}
assert DataConverterPlugin("sharegpt")(example) == {
"messages": [
{"role": "user", "content": [{"type": "text", "value": "hi"}], "loss_weight": 0.0},
{"role": "assistant", "content": [{"type": "text", "value": "yo"}], "loss_weight": 1.0},
]
}
def test_alpaca_converter_multimodal():
example = {"instruction": "Describe <image>", "input": "", "output": "ok", "images": ["/p/a.jpg"]}
user = DataConverterPlugin("alpaca")(example)["messages"][0]
assert user["content"] == [
{"type": "text", "value": "Describe "},
{"type": "image_url", "value": "/p/a.jpg"},
]
def test_pair_converter_multimodal_shared_media():
# chosen and rejected each reference the same sample-level image
example = {
"chosen": [
{"role": "user", "content": "Look at <image>"},
{"role": "assistant", "content": "good"},
],
"rejected": [
{"role": "user", "content": "Look at <image>"},
{"role": "assistant", "content": "bad"},
],
"images": ["/p/a.jpg"],
}
out = DataConverterPlugin("pair")(example)
for side in ("chosen_messages", "rejected_messages"):
assert out[side][0]["content"] == [
{"type": "text", "value": "Look at "},
{"type": "image_url", "value": "/p/a.jpg"},
]
def test_converter_media_count_mismatch():
# more tags than media files
with pytest.raises(ValueError, match="More <image> tags"):
DataConverterPlugin("sharegpt")(
{
"conversations": [{"from": "human", "value": "<image><image>"}, {"from": "gpt", "value": "x"}],
"images": ["/p/a.jpg"],
}
)
# fewer tags than media files
with pytest.raises(ValueError, match="Fewer <image> tags"):
DataConverterPlugin("sharegpt")(
{
"conversations": [{"from": "human", "value": "<image>"}, {"from": "gpt", "value": "x"}],
"images": ["/p/a.jpg", "/p/b.jpg"],
}
)
def test_converter_audio_column_and_tag():
# an <audio> tag consumes the next path from the audios column, lifted into an audio_url block
example = {
"conversations": [
{"from": "human", "value": "hear <audio>What is this?"},
{"from": "gpt", "value": "A bell."},
],
"audios": ["/p/a.wav"],
}
user = DataConverterPlugin("sharegpt")(example)["messages"][0]
assert user["content"] == [
{"type": "text", "value": "hear "},
{"type": "audio_url", "value": "/p/a.wav"},
{"type": "text", "value": "What is this?"},
]
def test_converter_audio_count_mismatch():
# more audio tags than files
with pytest.raises(ValueError, match="More <audio> tags"):
DataConverterPlugin("sharegpt")(
{
"conversations": [{"from": "human", "value": "<audio><audio>"}, {"from": "gpt", "value": "x"}],
"audios": ["/p/a.wav"],
}
)
# fewer audio tags than files
with pytest.raises(ValueError, match="Fewer <audio> tags"):
DataConverterPlugin("sharegpt")(
{
"conversations": [{"from": "human", "value": "<audio>"}, {"from": "gpt", "value": "x"}],
"audios": ["/p/a.wav", "/p/b.wav"],
}
)
@pytest.mark.parametrize("num_samples", [16]) @pytest.mark.parametrize("num_samples", [16])
def test_pair_converter(num_samples: int): def test_pair_converter(num_samples: int):
data_args = DataArguments(train_dataset="llamafactory/v1-dataset-info/orca-dpo-pairs.yaml") data_args = DataArguments(train_dataset="llamafactory/v1-dataset-info/orca-dpo-pairs.yaml")
@@ -117,3 +259,4 @@ def test_pair_converter(num_samples: int):
], ],
} }
assert data_engine[index] == {"_dataset_name": "tiny_dataset", **expected_data} assert data_engine[index] == {"_dataset_name": "tiny_dataset", **expected_data}