mirror of
https://github.com/hiyouga/LLaMA-Factory.git
synced 2026-08-17 05:25:44 +08:00
[v1] Support multimodal data training (#10656)
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
25
data/v1_multimodal_demo.jsonl
Normal file
25
data/v1_multimodal_demo.jsonl
Normal 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日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]}
|
||||
|
||||
4
data/v1_multimodal_demo.yaml
Normal file
4
data/v1_multimodal_demo.yaml
Normal file
@@ -0,0 +1,4 @@
|
||||
multimodal_demo:
|
||||
path: data/v1_multimodal_demo.jsonl
|
||||
source: local
|
||||
|
||||
27
examples/v1/train_full/train_multimodal.yaml
Normal file
27
examples/v1/train_full/train_multimodal.yaml
Normal 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
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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,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."""
|
||||
hf_messages = []
|
||||
for message in messages:
|
||||
tool_calls: list[dict] = []
|
||||
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 = ""
|
||||
for content in message["content"]:
|
||||
if content["type"] == "text":
|
||||
@@ -52,8 +81,12 @@ def _to_hf_messages(messages: list[Message]) -> list[dict]:
|
||||
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"]}})
|
||||
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:
|
||||
@@ -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."
|
||||
)
|
||||
|
||||
@@ -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 ``<think>``) 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.")
|
||||
|
||||
@@ -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)
|
||||
|
||||
277
src/llamafactory/v1/core/utils/collation.py
Normal file
277
src/llamafactory/v1/core/utils/collation.py
Normal 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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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", "<image>")
|
||||
VIDEO_PLACEHOLDER = os.getenv("VIDEO_PLACEHOLDER", "<video>")
|
||||
AUDIO_PLACEHOLDER = os.getenv("AUDIO_PLACEHOLDER", "<audio>")
|
||||
|
||||
@@ -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.
|
||||
@@ -23,7 +23,7 @@ from transformers import set_seed as hf_set_seed
|
||||
from ..accelerator.helper import is_torch_npu_available
|
||||
from ..accelerator.interface import DistributedInterface
|
||||
from .constants import IGNORE_INDEX
|
||||
from .types import BatchInput, ModelInput, Processor, Tensor
|
||||
from .types import BatchInput, Processor
|
||||
|
||||
|
||||
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
|
||||
|
||||
|
||||
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:
|
||||
"""Compute valid tokens in batches.
|
||||
|
||||
@@ -125,3 +94,15 @@ def compute_valid_tokens(batches: list[BatchInput]) -> int:
|
||||
for batch in batches
|
||||
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
|
||||
|
||||
@@ -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.
|
||||
@@ -141,6 +141,20 @@ class ModelInput(TypedDict, total=False):
|
||||
"""Position ids for the model (optional)."""
|
||||
token_type_ids: NotRequired[list[int]]
|
||||
"""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):
|
||||
@@ -156,6 +170,20 @@ class BatchInput(TypedDict, total=False):
|
||||
"""Position ids for the model (optional)."""
|
||||
token_type_ids: NotRequired[Tensor]
|
||||
"""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):
|
||||
|
||||
@@ -370,3 +370,208 @@ def test_dynamic_padding_free_fill_buffer_restarts_until_micro_batch_is_complete
|
||||
assert len(batch) == 1
|
||||
assert batch[0]["input_ids"].shape == (1, 18)
|
||||
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
|
||||
|
||||
|
||||
@@ -71,6 +71,148 @@ def test_sharegpt_converter():
|
||||
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])
|
||||
def test_pair_converter(num_samples: int):
|
||||
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}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user