Compare commits

...

9 Commits

Author SHA1 Message Date
WinterShiver 5f48dd7bf9
Merge 46cda3c3a9 into 713b5a3f95 2026-08-03 17:39:08 -07:00
SSSSuperC 713b5a3f95
[model] add MOSS-VL support (#10708)
docker / build (cuda) (push) Waiting to run Details
tests / tests (macos-latest, 3.11, ) (push) Waiting to run Details
tests / tests (macos-latest, 3.12, ) (push) Waiting to run Details
tests / tests (macos-latest, 3.13, ) (push) Waiting to run Details
tests / tests (ubuntu-latest, 3.11, ) (push) Waiting to run Details
tests / tests (ubuntu-latest, 3.11, 4.55.0) (push) Waiting to run Details
tests / tests (ubuntu-latest, 3.11, 4.57.1) (push) Waiting to run Details
tests / tests (ubuntu-latest, 3.12, ) (push) Waiting to run Details
tests / tests (ubuntu-latest, 3.13, ) (push) Waiting to run Details
tests / tests (windows-latest, 3.11, ) (push) Waiting to run Details
tests / tests (windows-latest, 3.12, ) (push) Waiting to run Details
tests / tests (windows-latest, 3.13, ) (push) Waiting to run Details
tests_cuda / tests (linux-x86_64-gpu-2, 3.11) (push) Waiting to run Details
tests_npu / tests (linux-aarch64-a2-4, 3.11, 2.7.1) (push) Waiting to run Details
2026-08-03 18:18:24 +08:00
浮梦 62ae362455
[v1] Support multimodal data training (#10656)
Co-authored-by: Claude Opus 4.8 <noreply@anthropic.com>
2026-07-31 18:54:13 +08:00
Kyungmin Kim 3984675dd5
fix(ci): align workflow Python version with requires-python (#10707) 2026-07-31 16:56:00 +08:00
Chaoran Wei 1b47415a2f
[train] Fix hyper parallel tail accumulation loss scaling (#10705)
Co-authored-by: wcrzlh <weichaoran@huawei.com>
2026-07-30 17:29:59 +08:00
WinterShiver 46cda3c3a9
[minor] rm redundant code 2025-11-04 10:46:31 +08:00
WinterShiver 8368f7d6d4
[fix] read model dtype from config 2025-10-31 10:21:10 +08:00
WinterShiver b83f341276
[inference] formatted hf_infer script 2025-10-29 18:01:41 +08:00
WinterShiver 9deb38347f
[inference] add hf_infer script for inference using huggingface backend 2025-10-29 17:33:56 +08:00
37 changed files with 2889 additions and 132 deletions

View File

@ -29,7 +29,7 @@ jobs:
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: '3.10'
python-version: '3.11'
- name: Install dependencies
run: |

View File

@ -0,0 +1,25 @@
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "Who are they?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They're Kane and Gretzka from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "What are they doing?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "They are celebrating on the soccer field."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "Who is he?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He's Thomas Muller from Bayern Munich."}]}, {"role": "user", "content": [{"type": "text", "value": "Why is he on the ground?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Because he's sliding on his knees to celebrate."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "Please describe this image"}]}, {"role": "assistant", "content": [{"type": "text", "value": "Chinese astronaut Gui Haichao is giving a speech."}]}, {"role": "user", "content": [{"type": "text", "value": "What has he accomplished?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "He was appointed to be a payload specialist on Shenzhou 16 mission in June 2022, thus becoming the first Chinese civilian of Group 3 in space on 30 May 2023. He is responsible for the on-orbit operation of space science experimental payloads."}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}, {"type": "text", "value": "他们是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们是拜仁慕尼黑的凯恩和格雷茨卡。"}]}, {"role": "user", "content": [{"type": "text", "value": "他们在做什么?"}, {"type": "image_url", "value": "data/mllm_demo_data/1.jpg"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他们在足球场上庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/2.jpg"}, {"type": "text", "value": "他是谁?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他是来自拜仁慕尼黑的托马斯·穆勒。"}]}, {"role": "user", "content": [{"type": "text", "value": "他为什么在地上?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "因为他正在双膝跪地滑行庆祝。"}]}]}
{"messages": [{"role": "user", "content": [{"type": "image_url", "value": "data/mllm_demo_data/3.jpg"}, {"type": "text", "value": "请描述这张图片"}]}, {"role": "assistant", "content": [{"type": "text", "value": "中国宇航员桂海潮正在讲话。"}]}, {"role": "user", "content": [{"type": "text", "value": "他取得过哪些成就?"}]}, {"role": "assistant", "content": [{"type": "text", "value": "他于2022年6月被任命为神舟十六号任务的有效载荷专家从而成为2023年5月30日进入太空的首位平民宇航员。他负责在轨操作空间科学实验有效载荷。"}]}]}

View File

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

View File

@ -0,0 +1,6 @@
### Install model-specific dependencies: `pip install -r requirements/moss-vl.txt`
model_name_or_path: OpenMOSS-Team/MOSS-VL-Instruct-0708
template: moss_vl
infer_backend: huggingface # choices: [huggingface, vllm, sglang, ktransformers]
trust_remote_code: true

View File

@ -0,0 +1,14 @@
### Install model-specific dependencies: `pip install -r requirements/moss-vl.txt`
### Note: DO NOT use quantized model or quantization_bit when merging lora adapters
### model
model_name_or_path: OpenMOSS-Team/MOSS-VL-Instruct-0708
adapter_name_or_path: saves/moss-vl-11b/lora/sft
template: moss_vl
trust_remote_code: true
### export
export_dir: saves/moss_vl_sft_merged
export_size: 5
export_device: cpu # choices: [cpu, auto]
export_legacy_format: false

View File

@ -0,0 +1,57 @@
### Install model-specific dependencies: `pip install -r requirements/moss-vl.txt`
### model
model_name_or_path: OpenMOSS-Team/MOSS-VL-Instruct-0708
image_max_pixels: 262144
video_max_pixels: 16384
video_fps: 1.0
video_maxlen: 256
use_reentrant_gc: false
trust_remote_code: true
### method
stage: sft
do_train: true
finetuning_type: full
freeze_vision_tower: true
freeze_multi_modal_projector: true
freeze_language_model: false
deepspeed: examples/deepspeed/ds_z3_config.json
### dataset
dataset: mllm_demo,identity,alpaca_en_demo # video: mllm_video_demo
template: moss_vl
cutoff_len: 4096
max_samples: 1000
preprocessing_num_workers: 16
dataloader_num_workers: 4
packing: false
### output
output_dir: saves/moss-vl-11b/full/sft
logging_steps: 10
save_steps: 500
plot_loss: true
overwrite_output_dir: true
save_only_model: false
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
### train
per_device_train_batch_size: 1
gradient_accumulation_steps: 1
gradient_checkpointing: true
gradient_checkpointing_kwargs:
use_reentrant: false
learning_rate: 1.0e-5
num_train_epochs: 3.0
lr_scheduler_type: cosine
warmup_ratio: 0.1
bf16: true
ddp_timeout: 180000000
resume_from_checkpoint: null
### eval
# val_size: 0.1
# per_device_eval_batch_size: 1
# eval_strategy: steps
# eval_steps: 500

View File

@ -0,0 +1,54 @@
### Install model-specific dependencies: `pip install -r requirements/moss-vl.txt`
### model
model_name_or_path: OpenMOSS-Team/MOSS-VL-Instruct-0708
image_max_pixels: 262144
video_max_pixels: 16384
video_fps: 1.0
video_maxlen: 256
trust_remote_code: true
### method
stage: sft
do_train: true
finetuning_type: lora
lora_rank: 8
lora_target: all
freeze_vision_tower: true
freeze_multi_modal_projector: true
freeze_language_model: false
### dataset
dataset: mllm_demo,identity,alpaca_en_demo # video: mllm_video_demo
template: moss_vl
cutoff_len: 4096
max_samples: 1000
preprocessing_num_workers: 16
dataloader_num_workers: 4
packing: false
### output
output_dir: saves/moss-vl-11b/lora/sft
logging_steps: 10
save_steps: 500
plot_loss: true
overwrite_output_dir: true
save_only_model: false
report_to: none # choices: [none, wandb, tensorboard, swanlab, mlflow]
### train
per_device_train_batch_size: 2
gradient_accumulation_steps: 1
learning_rate: 1.0e-4
num_train_epochs: 3.0
lr_scheduler_type: cosine
warmup_ratio: 0.1
bf16: true
ddp_timeout: 180000000
resume_from_checkpoint: null
### eval
# val_size: 0.1
# per_device_eval_batch_size: 1
# eval_strategy: steps
# eval_steps: 500

View File

@ -0,0 +1,27 @@
model: Qwen/Qwen3.5-0.8B
model_class: llm
kernel_config:
name: auto
quant_config: null
dist_config:
name: fsdp2
dcp_path: null
### data
train_dataset: data/v1_multimodal_demo.yaml
### training
output_dir: outputs/test_multimodal
micro_batch_size: 1
cutoff_len: 2048
learning_rate: 1.0e-4
max_steps: 5
### sample
sample_backend: hf
max_new_tokens: 128

3
requirements/moss-vl.txt Normal file
View File

@ -0,0 +1,3 @@
transformers==4.57.1
torchcodec==0.7.0
joblib

146
scripts/hf_infer.py Normal file
View File

@ -0,0 +1,146 @@
import gc
import json
from typing import Optional
import torch
import fire
from tqdm import tqdm
from transformers import (
AutoModelForCausalLM,
AutoTokenizer,
GenerationConfig,
Seq2SeqTrainingArguments,
)
from peft import PeftModel
from llamafactory.data import get_dataset, get_template_and_fix_tokenizer
from llamafactory.extras.constants import IGNORE_INDEX
from llamafactory.hparams import get_infer_args
from llamafactory.model import load_tokenizer
def hf_infer(
model_name_or_path: str,
adapter_name_or_path: str = None,
dataset: str = "alpaca_en_demo",
dataset_dir: str = "data",
template: str = "default",
cutoff_len: int = 2048,
max_samples: Optional[int] = None,
save_name: str = "generated_predictions.jsonl",
temperature: float = 0.95,
top_p: float = 0.7,
top_k: int = 50,
max_new_tokens: int = 1024,
repetition_penalty: float = 1.0,
skip_special_tokens: bool = True,
default_system: Optional[str] = None,
enable_thinking: bool = True,
seed: Optional[int] = None,
batch_size: int = 4,
device: str = None,
):
"""
Perform batch generation using Hugging Face transformers backend.
Usage:
python hf_infer.py --model_name_or_path meta-llama/Llama-2-7b-hf --template llama --dataset alpaca_en_demo --batch_size 16 --save_name test.jsonl
"""
device = device or ("cuda" if torch.cuda.is_available() else "cpu")
model_args, data_args, _, _ = get_infer_args(
dict(
model_name_or_path=model_name_or_path,
adapter_name_or_path=adapter_name_or_path,
dataset=dataset,
dataset_dir=dataset_dir,
template=template,
cutoff_len=cutoff_len,
max_samples=max_samples,
preprocessing_num_workers=8,
default_system=default_system,
enable_thinking=enable_thinking,
temperature=temperature,
top_p=top_p,
top_k=top_k,
max_new_tokens=max_new_tokens,
repetition_penalty=repetition_penalty,
)
)
training_args = Seq2SeqTrainingArguments(output_dir="dummy_dir")
tokenizer_module = load_tokenizer(model_args)
tokenizer = tokenizer_module["tokenizer"]
template_obj = get_template_and_fix_tokenizer(tokenizer, data_args)
template_obj.mm_plugin.expand_mm_tokens = False
# --- Load model ---
model = AutoModelForCausalLM.from_pretrained(
model_name_or_path,
torch_dtype=model_args.infer_dtype,
device_map="auto",
trust_remote_code=True,
)
if adapter_name_or_path is not None:
model = PeftModel.from_pretrained(model, adapter_name_or_path)
model.eval()
# --- Load dataset ---
dataset_module = get_dataset(template_obj, model_args, data_args, training_args, "ppo", **tokenizer_module)
train_dataset = dataset_module["train_dataset"]
# --- Generation configuration ---
gen_cfg = GenerationConfig(
temperature=temperature,
top_p=top_p,
top_k=top_k,
repetition_penalty=repetition_penalty,
max_new_tokens=max_new_tokens,
do_sample=True,
)
all_prompts, all_preds, all_labels = [], [], []
for i in tqdm(range(0, len(train_dataset), batch_size), desc="Processing batched inference"):
batch = train_dataset[i : min(i + batch_size, len(train_dataset))]
input_ids = [torch.tensor(x) for x in batch["input_ids"]]
input_ids = torch.nn.utils.rnn.pad_sequence(
input_ids, batch_first=True, padding_value=tokenizer.pad_token_id
).to(device)
# Generate
with torch.no_grad():
outputs = model.generate(
input_ids=input_ids,
**gen_cfg.to_dict(),
)
# Decode predictions
for j in range(len(batch["input_ids"])):
prompt = tokenizer.decode(batch["input_ids"][j], skip_special_tokens=skip_special_tokens)
label = tokenizer.decode(
list(filter(lambda x: x != IGNORE_INDEX, batch["labels"][j])),
skip_special_tokens=skip_special_tokens,
)
pred = tokenizer.decode(outputs[j][len(batch["input_ids"][j]) :], skip_special_tokens=skip_special_tokens)
all_prompts.append(prompt)
all_preds.append(pred)
all_labels.append(label)
gc.collect()
# Save all results
with open(save_name, "w", encoding="utf-8") as f:
for text, pred, label in zip(all_prompts, all_preds, all_labels):
f.write(json.dumps({"prompt": text, "predict": pred, "label": label}, ensure_ascii=False) + "\n")
print("*" * 70)
print(f"{len(all_prompts)} total generated results have been saved at {save_name}.")
print("*" * 70)
if __name__ == "__main__":
fire.Fire(hf_infer)

View File

@ -150,7 +150,9 @@ class MultiModalDataCollatorForSeq2Seq(DataCollatorForSeq2Seq):
if isinstance(self.model, PeftModel):
self.model = self.model.base_model.model
if self.model is not None and hasattr(self.model, "get_rope_index"): # for qwen2vl mrope
if getattr(getattr(self.model, "config", None), "model_type", None) == "moss_vl":
self.get_rope_func = None # MOSS-VL computes its own XRoPE positions in model.forward.
elif self.model is not None and hasattr(self.model, "get_rope_index"): # for qwen2vl mrope
self.get_rope_func = self.model.get_rope_index # transformers < 4.52.0 or qwen2.5 omni
elif self.model is not None and hasattr(self.model, "model") and hasattr(self.model.model, "get_rope_index"):
self.get_rope_func = self.model.model.get_rope_index # transformers >= 4.52.0
@ -322,6 +324,8 @@ class MultiModalDataCollatorForSeq2Seq(DataCollatorForSeq2Seq):
)
def __call__(self, features: list[dict[str, Any]]) -> dict[str, "torch.Tensor"]:
model_type = getattr(getattr(self.model, "config", None), "model_type", None)
is_moss_vl = model_type == "moss_vl"
batch_images, batch_videos, batch_audios = [], [], []
batch_imglens, batch_vidlens, batch_audlens, batch_input_ids = [], [], [], []
packing_params_list: list[dict[str, Any] | None] = []
@ -341,7 +345,10 @@ class MultiModalDataCollatorForSeq2Seq(DataCollatorForSeq2Seq):
fake_input_ids = []
has_dummy_image = False
if (
self.template.mm_plugin.image_token is not None and sum(batch_imglens) == 0 and sum(batch_vidlens) == 0
self.template.mm_plugin.image_token is not None
and sum(batch_imglens) == 0
and sum(batch_vidlens) == 0
and not is_moss_vl # MOSS-VL builds one native zero-valued dummy per text-only sample in its plugin.
): # avoid process hanging in zero3/fsdp case
fake_messages = [{"role": "user", "content": IMAGE_PLACEHOLDER}]
fake_images = [Image.new("RGB", (64, 64), (255, 255, 255))]
@ -416,7 +423,6 @@ class MultiModalDataCollatorForSeq2Seq(DataCollatorForSeq2Seq):
features: dict[str, torch.Tensor] = super().__call__(features)
bsz, seq_len = features["input_ids"].shape[:2]
model_type = getattr(self.model.config, "model_type", None) if self.model is not None else None
is_omni = model_type in [
"qwen2_5_omni_thinker",
"qwen3_omni_moe_thinker",
@ -461,12 +467,17 @@ class MultiModalDataCollatorForSeq2Seq(DataCollatorForSeq2Seq):
):
raise ValueError(f"{self.model.config.model_type} requires 3D position ids for mrope.")
if "cross_attention_mask" in mm_inputs: # for mllama inputs when pad_to_multiple_of is enabled
if (
"cross_attention_mask" in mm_inputs and mm_inputs["cross_attention_mask"].dtype != torch.bool
): # for mllama inputs when pad_to_multiple_of is enabled
cross_attention_mask = mm_inputs.pop("cross_attention_mask")
seq_len = features["input_ids"].size(1)
orig_len = cross_attention_mask.size(1)
mm_inputs["cross_attention_mask"] = F.pad(cross_attention_mask, (0, 0, 0, 0, 0, seq_len - orig_len))
if is_moss_vl:
mm_inputs = self.template.mm_plugin.post_process_mossvl_inputs(features, mm_inputs, self.processor)
features.update(mm_inputs)
if "image_bound" in features: # for minicpmv inputs

View File

@ -472,6 +472,354 @@ class BasePlugin(MMPluginMixin):
return self._get_mm_inputs(images, videos, audios, processor)
@dataclass
class MossVLPlugin(BasePlugin):
vision_bos_token: str = "<|vision_start|>"
vision_eos_token: str = "<|vision_end|>"
time_bos_token: str = "<|time_start|>"
time_eos_token: str = "<|time_end|>"
@staticmethod
def _split_pixel_values(
pixel_values: "torch.Tensor",
grid_thw: "torch.Tensor",
) -> list["torch.Tensor"]:
patch_counts = [int(grid.prod().item()) for grid in grid_thw]
return list(torch.split(pixel_values, patch_counts))
@staticmethod
def _create_cross_attention_mask(
input_ids: Union[list[list[int]], "torch.Tensor"],
grid_thw: "torch.Tensor",
media_nums_per_sample: list[int],
image_token_id: int,
attention_mask: Optional["torch.Tensor"] = None,
padding_side: Literal["left", "right"] = "right",
) -> "torch.Tensor":
r"""Create the native MOSS-VL frame-level causal cross-attention mask."""
if isinstance(input_ids, list):
max_text_len = max(len(token_ids) for token_ids in input_ids)
input_ids_tensor = torch.full((len(input_ids), max_text_len), -1, dtype=torch.long)
attention_mask_tensor = torch.zeros_like(input_ids_tensor, dtype=torch.bool)
for batch_index, token_ids in enumerate(input_ids):
seq_len = len(token_ids)
start = max_text_len - seq_len if padding_side == "left" else 0
input_ids_tensor[batch_index, start : start + seq_len] = torch.tensor(token_ids, dtype=torch.long)
attention_mask_tensor[batch_index, start : start + seq_len] = True
else:
input_ids_tensor = input_ids
attention_mask_tensor = (
torch.ones_like(input_ids_tensor, dtype=torch.bool)
if attention_mask is None
else attention_mask.bool()
)
total_frames_per_sample = []
media_index = 0
for num_media in media_nums_per_sample:
sample_grid = grid_thw[media_index : media_index + num_media]
total_frames_per_sample.append(int(sample_grid[:, 0].sum().item()))
media_index += num_media
max_num_frames = max(total_frames_per_sample)
frame_indices = torch.arange(max_num_frames, device=input_ids_tensor.device).view(1, 1, -1)
visible_mask = (input_ids_tensor == image_token_id).cumsum(dim=1).unsqueeze(-1) > frame_indices
visible_mask &= attention_mask_tensor.unsqueeze(-1)
valid_frames = frame_indices < torch.tensor(
total_frames_per_sample,
device=input_ids_tensor.device,
).view(-1, 1, 1)
visible_mask &= valid_frames
return (~visible_mask).unsqueeze(1)
def _get_video_inputs(
self,
videos: list["VideoInput"],
processor: "MMProcessor",
return_metadata: bool,
) -> dict[str, Any]:
video_kwargs = {"return_tensors": "pt", "return_metadata": return_metadata}
if getattr(processor, "video_fps", None) is not None:
video_kwargs["video_fps"] = processor.video_fps
if getattr(processor, "video_maxlen", None) is not None:
video_kwargs["max_frames"] = processor.video_maxlen
video_min_pixels = getattr(processor, "video_min_pixels", None)
video_max_pixels = getattr(processor, "video_max_pixels", None)
if video_min_pixels is not None and video_max_pixels is not None:
video_kwargs["size"] = {
"shortest_edge": video_min_pixels,
"longest_edge": video_max_pixels,
}
return dict(processor.video_processor(videos=videos, **video_kwargs))
def _get_media_order_from_ids(
self,
input_ids: list[int],
processor: "MMProcessor",
num_images: int,
num_videos: int,
expected_video_frames: Optional[list[int]] = None,
) -> list[str]:
media_order = []
video_frame_counts = []
in_video = False
current_video_frames = 0
for token_id in input_ids:
if token_id == processor.vision_start_token_id:
if in_video:
raise ValueError(
"MOSS-VL encountered nested video token blocks after tokenization. "
"Please increase `cutoff_len` if a video placeholder was truncated."
)
media_order.append("video")
in_video = True
current_video_frames = 0
elif token_id == processor.vision_end_token_id:
if not in_video:
raise ValueError(
"MOSS-VL encountered a video end token without a matching start token after tokenization. "
"Please increase `cutoff_len` if a video placeholder was truncated."
)
video_frame_counts.append(current_video_frames)
in_video = False
elif token_id == processor.image_token_id:
if in_video:
current_video_frames += 1
else:
media_order.append("image")
if in_video:
raise ValueError(
"MOSS-VL encountered an incomplete video token block after tokenization. "
"Please increase `cutoff_len` or reduce `video_maxlen`."
)
if media_order.count("image") != num_images or media_order.count("video") != num_videos:
raise ValueError(
"MOSS-VL media tokens do not match the provided media after tokenization: "
f"order={media_order}, images={num_images}, videos={num_videos}. "
"Please increase `cutoff_len` if a visual placeholder was truncated."
)
if expected_video_frames is not None and video_frame_counts != expected_video_frames:
raise ValueError(
"MOSS-VL video frame tokens do not match the processed video after tokenization: "
f"tokens={video_frame_counts}, frames={expected_video_frames}. "
"Please increase `cutoff_len` or reduce `video_maxlen`."
)
return media_order
@override
def process_messages(
self,
messages: list[dict[str, str]],
images: list["ImageInput"],
videos: list["VideoInput"],
audios: list["AudioInput"],
processor: Optional["MMProcessor"],
) -> list[dict[str, str]]:
self._validate_input(processor, images, videos, audios)
self._validate_messages(messages, images, videos, audios)
messages = deepcopy(messages)
video_inputs = self._get_video_inputs(videos, processor, return_metadata=True) if videos else {}
video_grid_thw = video_inputs.get("video_grid_thw", [])
video_metadata = video_inputs.get("video_metadata", [])
video_index = 0
for message in messages:
content = message["content"]
content = content.replace(IMAGE_PLACEHOLDER, self.image_token)
while VIDEO_PLACEHOLDER in content:
metadata = video_metadata[video_index]
if metadata.fps is None:
metadata.fps = 24
timestamps = processor._calculate_timestamps(
metadata.frames_indices,
metadata.total_num_frames,
metadata.fps,
metadata.duration,
processor.video_processor.temporal_patch_size,
actual_timestamps=getattr(metadata, "actual_timestamps", None),
)
num_frames = int(video_grid_thw[video_index][0].item())
frame_tokens = [
f"{self.time_bos_token}{timestamps[frame_idx]:.1f} seconds{self.time_eos_token}{self.image_token}"
for frame_idx in range(num_frames)
]
video_tokens = f"{self.vision_bos_token}{''.join(frame_tokens)}{self.vision_eos_token}"
content = content.replace(VIDEO_PLACEHOLDER, video_tokens, 1)
video_index += 1
message["content"] = content
return messages
@override
def get_mm_inputs(
self,
images: list["ImageInput"],
videos: list["VideoInput"],
audios: list["AudioInput"],
imglens: list[int],
vidlens: list[int],
audlens: list[int],
batch_ids: list[list[int]],
processor: Optional["MMProcessor"],
) -> dict[str, Union[list[int], "torch.Tensor"]]:
self._validate_input(processor, images, videos, audios)
if audios:
raise ValueError("MOSS-VL does not support audio inputs.")
if not (len(imglens) == len(vidlens) == len(batch_ids)):
raise ValueError("MOSS-VL batch metadata must have one entry per sample.")
final_pixel_values = []
final_grid_thw = []
media_nums_per_sample = []
image_offset = 0
video_offset = 0
for imglen, vidlen, input_ids in zip(imglens, vidlens, batch_ids):
sample_images = images[image_offset : image_offset + imglen]
sample_videos = videos[video_offset : video_offset + vidlen]
image_offset += imglen
video_offset += vidlen
image_chunks, image_grids = [], []
if sample_images:
regularized_images = self._regularize_images(
sample_images,
image_max_pixels=2**63 - 1,
image_min_pixels=1,
)["images"]
image_kwargs = {"return_tensors": "pt"}
if getattr(processor, "image_min_pixels", None) is not None:
image_kwargs["min_pixels"] = processor.image_min_pixels
if getattr(processor, "image_max_pixels", None) is not None:
image_kwargs["max_pixels"] = processor.image_max_pixels
image_inputs = processor.image_processor(images=regularized_images, **image_kwargs)
image_grids = list(image_inputs["image_grid_thw"])
image_chunks = self._split_pixel_values(image_inputs["pixel_values"], image_inputs["image_grid_thw"])
video_chunks, video_grids = [], []
if sample_videos:
video_inputs = self._get_video_inputs(sample_videos, processor, return_metadata=False)
video_grids = list(video_inputs["video_grid_thw"])
video_chunks = self._split_pixel_values(
video_inputs["pixel_values_videos"],
video_inputs["video_grid_thw"],
)
media_order = self._get_media_order_from_ids(
input_ids,
processor,
imglen,
vidlen,
expected_video_frames=[int(grid[0].item()) for grid in video_grids],
)
if not media_order:
patch_size = getattr(processor.image_processor, "patch_size", None)
if patch_size is None: # lightweight/test processors without the native MOSS-VL contract
blank_image = Image.new("RGB", (128, 128), (255, 255, 255))
blank_inputs = processor.image_processor(images=[blank_image], return_tensors="pt")
final_pixel_values.append(blank_inputs["pixel_values"])
final_grid_thw.append(blank_inputs["image_grid_thw"][0])
else:
temporal_patch_size = getattr(processor.image_processor, "temporal_patch_size", None) or 1
merge_size = getattr(processor.image_processor, "merge_size", None) or 2
factor = patch_size * merge_size
side = math.ceil(128 / factor) * factor
grid_thw = torch.tensor([1, side // patch_size, side // patch_size])
feature_dim = 3 * temporal_patch_size * patch_size * patch_size
final_pixel_values.append(torch.zeros((int(grid_thw.prod()), feature_dim), dtype=torch.float32))
final_grid_thw.append(grid_thw)
media_nums_per_sample.append(1)
continue
image_index = 0
video_index = 0
for modality in media_order:
if modality == "image":
final_pixel_values.append(image_chunks[image_index])
final_grid_thw.append(image_grids[image_index])
image_index += 1
else:
final_pixel_values.append(video_chunks[video_index])
final_grid_thw.append(video_grids[video_index])
video_index += 1
media_nums_per_sample.append(len(media_order))
if image_offset != len(images) or video_offset != len(videos):
raise ValueError("MOSS-VL media lengths do not consume all provided inputs.")
mm_inputs = {
"pixel_values": torch.cat(final_pixel_values, dim=0),
"grid_thw": torch.stack(final_grid_thw),
"media_nums_per_sample": media_nums_per_sample,
}
mm_inputs["cross_attention_mask"] = self._create_cross_attention_mask(
batch_ids,
mm_inputs["grid_thw"],
media_nums_per_sample,
processor.image_token_id,
padding_side=processor.tokenizer.padding_side,
)
return mm_inputs
def post_process_mossvl_inputs(
self,
features: dict[str, "torch.Tensor"],
mm_inputs: dict[str, Any],
processor: "MMProcessor",
) -> dict[str, Any]:
r"""Create MOSS-VL batch-only inputs after the text batch has been padded."""
input_ids = features["input_ids"]
attention_mask = features["attention_mask"].bool()
mm_inputs["cross_attention_mask"] = self._create_cross_attention_mask(
input_ids,
mm_inputs["grid_thw"],
mm_inputs["media_nums_per_sample"],
processor.image_token_id,
attention_mask,
)
dummy_image_tokens = (input_ids == processor.image_token_id) & ~attention_mask
input_ids.masked_fill_(dummy_image_tokens, processor.tokenizer.pad_token_id)
labels = features.get("labels")
if labels is not None:
control_token_ids = {
processor.image_token_id,
processor.video_token_id,
processor.vision_start_token_id,
processor.vision_end_token_id,
processor.tokenizer.convert_tokens_to_ids(self.time_bos_token),
processor.tokenizer.convert_tokens_to_ids(self.time_eos_token),
}
for batch_index, token_ids in enumerate(input_ids):
in_vision = False
for token_index, token_id in enumerate(token_ids.tolist()):
if token_id == processor.vision_start_token_id:
in_vision = True
if in_vision or token_id in control_token_ids:
labels[batch_index, token_index] = IGNORE_INDEX
if token_id == processor.vision_end_token_id:
in_vision = False
# Native MOSS-VL labels_spans supervise through <|im_end|>, but not its trailing newline.
im_end_token_id = processor.tokenizer.convert_tokens_to_ids("<|im_end|>")
labels[:, 1:].masked_fill_(input_ids[:, :-1] == im_end_token_id, IGNORE_INDEX)
features.pop("position_ids", None)
return mm_inputs
@dataclass
class ErnieVLPlugin(BasePlugin):
@override
@ -2911,6 +3259,7 @@ PLUGINS = {
"minicpm_v": MiniCPMVPlugin,
"minicpm_v_4_6": MiniCPMV4_6Plugin,
"mllama": MllamaPlugin,
"moss_vl": MossVLPlugin,
"paligemma": PaliGemmaPlugin,
"pixtral": PixtralPlugin,
"qwen2_audio": Qwen2AudioPlugin,

View File

@ -333,6 +333,50 @@ class Template:
return modelfile
@dataclass
class MossVLTemplate(Template):
@override
def _encode(
self,
tokenizer: "PreTrainedTokenizer",
messages: list[dict[str, str]],
system: Optional[str],
tools: Optional[str],
) -> list[list[int]]:
system = system or self.default_system
encoded_messages = []
for i, message in enumerate(messages):
elements = []
if i == 0:
elements += self.format_prefix.apply()
if system or tools:
tool_text = self.format_tools.apply(content=tools)[0] if tools else ""
if tools and not system:
tool_text = tool_text.lstrip("\n")
elements += self.format_system.apply(content=(system + tool_text))
if message["role"] == Role.USER:
elements += self.format_user.apply(content=message["content"], idx=str(i // 2))
elif message["role"] == Role.ASSISTANT:
elements += self.format_assistant.apply(content=message["content"])
elif message["role"] == Role.OBSERVATION:
elements += self.format_observation.apply(content=message["content"])
elif message["role"] == Role.FUNCTION:
elements += self.format_function.apply(
content=message["content"],
thought_words=self.thought_words,
tool_call_words=self.tool_call_words,
)
else:
raise NotImplementedError("Unexpected role: {}".format(message["role"]))
encoded_messages.append(self._convert_elements_to_ids(tokenizer, elements))
return encoded_messages
@dataclass
class Llama2Template(Template):
r"""A template that fuse the system message to first user message."""
@ -1526,6 +1570,32 @@ register_template(
)
# copied from qwen template
register_template(
name="moss_vl",
format_user=StringFormatter(slots=["<|im_start|>user\n{{content}}<|im_end|>\n<|im_start|>assistant\n"]),
format_assistant=StringFormatter(slots=["{{content}}<|im_end|>\n"]),
format_system=StringFormatter(slots=["<|im_start|>system\n{{content}}<|im_end|>\n"]),
format_function=FunctionFormatter(slots=["{{content}}<|im_end|>\n"], tool_format="qwen"),
format_observation=StringFormatter(
slots=["<|im_start|>user\n<tool_response>\n{{content}}\n</tool_response><|im_end|>\n<|im_start|>assistant\n"]
),
format_tools=ToolFormatter(tool_format="qwen"),
stop_words=["<|im_end|>"],
replace_eos=True,
mm_plugin=get_mm_plugin(
name="moss_vl",
image_token="<|image_pad|>",
video_token="<|video_pad|>",
vision_bos_token="<|vision_start|>",
vision_eos_token="<|vision_end|>",
time_bos_token="<|time_start|>",
time_eos_token="<|time_end|>",
),
template_class=MossVLTemplate,
)
# copied from vicuna template
register_template(
name="llava",

View File

@ -2198,6 +2198,17 @@ register_model_group(
)
register_model_group(
models={
"MOSS-VL-Instruct-0708": {
DownloadSource.DEFAULT: "OpenMOSS-Team/MOSS-VL-Instruct-0708",
},
},
template="moss_vl",
multimodal=True,
)
register_model_group(
models={
"OLMo-1B": {

View File

@ -56,7 +56,7 @@ class CompositeModel:
)
break
if project_module is not None:
if isinstance(project_module, torch.nn.Module):
mm_projectors.append(project_module)
return mm_projectors
@ -344,6 +344,15 @@ _register_composite_model(
)
_register_composite_model(
model_type="moss_vl",
projector_keys=["model.visual.merger", "model.separator_token"],
vision_model_keys=["model.visual.pos_embed", "model.visual.patch_embed", "model.visual.blocks"],
language_model_keys=["model.language_model", "lm_head"],
lora_conflict_keys=["patch_embed"],
)
_register_composite_model(
model_type="mllama",
vision_model_keys=["vision_model"],

View File

@ -23,7 +23,7 @@ from transformers.modeling_utils import is_fsdp_enabled
from transformers.utils import is_torch_cuda_available, is_torch_npu_available
from ..extras import logging
from ..extras.misc import infer_optim_dtype
from ..extras.misc import check_version, infer_optim_dtype
from ..extras.packages import is_transformers_version_greater_than
from .model_utils.attention import configure_attn_implementation, print_attn_implementation
from .model_utils.checkpointing import prepare_model_for_training
@ -418,6 +418,10 @@ def patch_config(
"pip install git+https://github.com/huggingface/transformers.git@3c2517727ce28a30f5044e01663ee204deb1cdbe"
)
if getattr(config, "model_type", None) == "moss_vl":
check_version("transformers==4.57.1", mandatory=True)
check_version("torchcodec==0.7.0", mandatory=True)
if getattr(config, "model_type", None) == "qwen3_omni_moe":
patch_qwen3_omni_moe_thinker_text_sparse_moe_block()

View File

@ -361,7 +361,12 @@ class HyperParallelTrainer(CustomSeq2SeqTrainer):
loss = loss.mean()
if not getattr(self, "model_accepts_loss_kwargs", False) and getattr(self, "compute_loss_func", None) is None:
loss = loss / self.args.gradient_accumulation_steps
accumulation_steps = getattr(
self,
"current_gradient_accumulation_steps",
self.args.gradient_accumulation_steps,
)
loss = loss / accumulation_steps
self.accelerator.backward(loss)

View File

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

View File

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

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
@ -14,13 +14,15 @@
"""Message <-> HF-template plumbing for rendering.
Pure, stateless helpers: convert v1 ``Message`` to HF chat-template format. No tokenization policy
decisions live here -- only mechanical conversion used by ``rendering.py``.
Pure, stateless helpers: convert v1 ``Message`` to HF chat-template format, extract/count media, and
guard media placeholder counts. No tokenization policy decisions live here -- only mechanical
conversion used by ``rendering.py``.
"""
import json
from ...utils.types import Message
from ...utils.helper import get_tokenizer
from ...utils.types import Message, Processor
_FALLBACK_CHATML_JINJA = (
@ -33,28 +35,59 @@ _FALLBACK_CHATML_JINJA = (
)
def _to_hf_messages(messages: list[Message]) -> list[dict]:
def _to_hf_messages(messages: list[Message], is_multimodal: bool = False) -> list[dict]:
"""Convert v1 Message format to HF format for apply_chat_template."""
hf_messages = []
for message in messages:
tool_calls: list[dict] = []
reasoning_content = ""
text = ""
for content in message["content"]:
if content["type"] == "text":
text += content["value"]
elif content["type"] == "reasoning":
reasoning_content += content["value"]
elif content["type"] == "tool_call":
try:
tc = json.loads(content["value"])
except json.JSONDecodeError as e:
raise ValueError(f"tool_call value is not valid JSON: {content['value']!r}") from e
if not isinstance(tc, dict) or "name" not in tc or "arguments" not in tc:
raise ValueError(f"tool_call must be a JSON object with 'name' and 'arguments' keys, got {tc!r}")
tool_calls.append({"type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}})
hf_msg = {"role": message["role"], "content": text}
if is_multimodal:
hf_content = []
for content in message["content"]:
if content["type"] == "text":
hf_content.append({"type": "text", "text": content["value"]})
elif content["type"] == "reasoning":
reasoning_content += content["value"]
elif content["type"] == "tool_call":
try:
tc = json.loads(content["value"])
except json.JSONDecodeError as e:
raise ValueError(f"tool_call value is not valid JSON: {content['value']!r}") from e
if not isinstance(tc, dict) or "name" not in tc or "arguments" not in tc:
raise ValueError(
f"tool_call must be a JSON object with 'name' and 'arguments' keys, got {tc!r}"
)
tool_calls.append(
{"type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}}
)
elif content["type"] == "image_url":
hf_content.append({"type": "image", "image": content["value"]})
elif content["type"] == "video_url":
hf_content.append({"type": "video", "video": content["value"]})
elif content["type"] == "audio_url":
hf_content.append({"type": "audio", "audio": content["value"]})
hf_msg = {"role": message["role"], "content": hf_content}
else:
text = ""
for content in message["content"]:
if content["type"] == "text":
text += content["value"]
elif content["type"] == "reasoning":
reasoning_content += content["value"]
elif content["type"] == "tool_call":
try:
tc = json.loads(content["value"])
except json.JSONDecodeError as e:
raise ValueError(f"tool_call value is not valid JSON: {content['value']!r}") from e
if not isinstance(tc, dict) or "name" not in tc or "arguments" not in tc:
raise ValueError(
f"tool_call must be a JSON object with 'name' and 'arguments' keys, got {tc!r}"
)
tool_calls.append(
{"type": "function", "function": {"name": tc["name"], "arguments": tc["arguments"]}}
)
hf_msg = {"role": message["role"], "content": text}
if tool_calls:
hf_msg["tool_calls"] = tool_calls
@ -63,3 +96,75 @@ def _to_hf_messages(messages: list[Message]) -> list[dict]:
hf_messages.append(hf_msg)
return hf_messages
def _extract_media_from_messages(messages: list[Message]) -> tuple[list, list, list]:
"""Extract image, video and audio paths/values from messages in order."""
images, videos, audios = [], [], []
for message in messages:
for content in message["content"]:
if content["type"] == "image_url":
images.append(content["value"])
elif content["type"] == "video_url":
videos.append(content["value"])
elif content["type"] == "audio_url":
audios.append(content["value"])
return images, videos, audios
def _count_media_in_messages(messages: list[Message]) -> tuple[int, int, int]:
"""Count total images, videos and audios in messages."""
n_images, n_videos, n_audios = 0, 0, 0
for message in messages:
for content in message["content"]:
if content["type"] == "image_url":
n_images += 1
elif content["type"] == "video_url":
n_videos += 1
elif content["type"] == "audio_url":
n_audios += 1
return n_images, n_videos, n_audios
def _load_audios(values: list, sampling_rate: int) -> list:
"""Load audio inputs into mono waveforms resampled to ``sampling_rate``."""
import numpy as np
import torchaudio
results = []
for value in values:
if isinstance(value, np.ndarray):
results.append(value)
continue
waveform, sr = torchaudio.load(value)
if waveform.shape[0] > 1: # downmix to mono
waveform = waveform.mean(dim=0, keepdim=True)
if sr != sampling_rate:
waveform = torchaudio.functional.resample(waveform, sr, sampling_rate)
results.append(waveform.squeeze(0).numpy())
return results
def _check_placeholder_counts(
processor: "Processor", full_text: str, n_images: int, n_videos: int, n_audios: int = 0
) -> None:
"""Guard: every media placeholder in the rendered text must originate from a media block."""
tokenizer = get_tokenizer(processor)
for attr, count, kind in (
("image_token_id", n_images, "image"),
("video_token_id", n_videos, "video"),
("audio_token_id", n_audios, "audio"),
):
tid = getattr(processor, attr, None)
if tid is None:
tid = getattr(tokenizer, attr, None)
if tid is None:
continue
placeholder = tokenizer.convert_ids_to_tokens(tid)
seen = full_text.count(placeholder)
if seen != count:
raise ValueError(
f"{kind} placeholder count ({seen}) != number of {kind} blocks ({count}); "
"media must be provided via image_url/video_url content blocks."
)

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# 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.")

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# 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)

View File

@ -0,0 +1,277 @@
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
"""Batch collation utils: padding/truncation and multimodal-feature alignment.
These operate on already-rendered ``ModelInput`` dicts (token lists + pixel tensors) and produce
padded ``BatchInput`` tensors. They are pure batching concerns -- independent of how a sample was
rendered -- and are consumed by the batch generators in ``core/utils/batching.py`` and
``plugins/trainer_plugins/batching.py``. Kept out of ``rendering.py`` so that file is only about
turning messages into a single tokenized sample.
"""
import torch
from ...utils.constants import IGNORE_INDEX
from ...utils.types import BatchInput, ModelInput, Tensor
# Multimodal feature keys the processor emits per sample. They are NOT padded/stacked like text
# fields: pixel/audio-feature tensors are ragged (variable patch / frame counts), so the collators
# concatenate them along dim 0 instead. Shared by rendering (which copies them verbatim from the
# processor) and the collators (which merge them across a micro batch).
_MULTIMODAL_PASSTHROUGH_KEYS = frozenset(
{
"pixel_values",
"image_grid_thw",
"pixel_values_videos",
"video_grid_thw",
"second_per_grid_ts", # Qwen2.5-VL name for the video temporal grid spacing
"video_second_per_grid", # Qwen2.5-Omni name for the same (fed to get_rope_index)
"input_features",
"feature_attention_mask",
}
)
def _pad_and_truncate(tensor: Tensor, max_seqlen: int, pad_value: int = 0) -> Tensor:
if tensor.shape[-1] >= max_seqlen:
return tensor[..., :max_seqlen]
pad_shape = list(tensor.shape)
pad_shape[-1] = max_seqlen - tensor.shape[-1]
pad_tensor = torch.full(pad_shape, pad_value, dtype=tensor.dtype, device=tensor.device)
return torch.cat([tensor, pad_tensor], dim=-1)
def _align_grid_media(
sample: ModelInput,
mm_type_ids: list[int],
max_length: int,
*,
target: int,
grid_key: str,
pixel_key: str,
) -> list[int]:
"""Trim and zero one modality's orphaned tokens for a single sample.
Layout-agnostic: a media item's placeholder tokens may be a single contiguous run or split
into per-frame sub-runs; completeness is decided per token *position*, so both are handled
identically.
Returns the (possibly updated) ``mm_token_type_ids`` so chained calls see earlier zeroing.
"""
if grid_key not in sample or pixel_key not in sample:
return mm_type_ids
grid = sample[grid_key]
n_items = len(grid)
if n_items == 0:
return mm_type_ids
positions = [i for i, t in enumerate(mm_type_ids) if t == target]
patches_per_item = [int(grid[i].prod()) for i in range(n_items)]
total_patches = sum(patches_per_item)
total_tokens = len(positions)
# merge_size**2 = pixel patches per placeholder token, derived from the data. Bail out
# untouched if the sample is inconsistent.
if total_tokens == 0 or total_patches % total_tokens != 0:
return mm_type_ids
merge_sq = total_patches // total_tokens
tokens_per_item = [p // merge_sq for p in patches_per_item]
if sum(tokens_per_item) != total_tokens:
return mm_type_ids
# Each item owns a contiguous slice of `positions`; it is complete iff its last
# placeholder token lands inside the kept window [0, max_length).
n_complete = 0
cum = 0
for n_i in tokens_per_item:
if positions[cum + n_i - 1] < max_length:
n_complete += 1
cum += n_i
else:
break
if n_complete >= n_items:
return mm_type_ids
# Trim pixel features and grid to the complete prefix.
keep_patches = sum(patches_per_item[:n_complete])
sample[pixel_key] = sample[pixel_key][:keep_patches]
sample[grid_key] = grid[:n_complete]
# Zero out orphaned placeholder tokens that fall inside the kept window; tokens
# beyond max_length are removed by truncation anyway (positions are sorted).
input_ids = list(sample["input_ids"])
mm_type_ids = list(mm_type_ids)
labels = list(sample["labels"]) if "labels" in sample else None
loss_weights = list(sample["loss_weights"]) if "loss_weights" in sample else None
for pos in positions[cum:]:
if pos >= max_length:
break
input_ids[pos] = 0
mm_type_ids[pos] = 0
if labels is not None:
labels[pos] = IGNORE_INDEX
if loss_weights is not None:
loss_weights[pos] = 0.0
sample["input_ids"] = input_ids
sample["mm_token_type_ids"] = mm_type_ids
if labels is not None:
sample["labels"] = labels
if loss_weights is not None:
sample["loss_weights"] = loss_weights
return mm_type_ids
def _align_audio(sample: ModelInput, mm_type_ids: list[int], max_length: int, *, target: int = 3) -> list[int]:
"""Trim and zero orphaned audio tokens for a single sample on truncation.
Returns the (possibly updated) ``mm_token_type_ids``.
"""
if "input_features" not in sample or "feature_attention_mask" not in sample:
return mm_type_ids
n_items = sample["input_features"].shape[0]
if n_items == 0:
return mm_type_ids
positions = [i for i, t in enumerate(mm_type_ids) if t == target]
if not positions:
return mm_type_ids
# Group the marked positions into maximal contiguous runs; each run is one audio's token span.
runs: list[tuple[int, int]] = []
run_start = prev = positions[0]
for pos in positions[1:]:
if pos != prev + 1:
runs.append((run_start, prev))
run_start = pos
prev = pos
runs.append((run_start, prev))
# Layout must match the feature rows one-to-one, else bail rather than corrupt the mapping.
if len(runs) != n_items:
return mm_type_ids
# An audio is complete iff its last placeholder token lands inside the kept window.
n_complete = 0
for _start, end in runs:
if end < max_length:
n_complete += 1
else:
break
if n_complete >= n_items:
return mm_type_ids
# Trim feature rows to the complete prefix.
sample["input_features"] = sample["input_features"][:n_complete]
sample["feature_attention_mask"] = sample["feature_attention_mask"][:n_complete]
# Zero out orphaned placeholder tokens that fall inside the kept window; tokens beyond
# max_length are removed by truncation anyway.
input_ids = list(sample["input_ids"])
mm_type_ids = list(mm_type_ids)
labels = list(sample["labels"]) if "labels" in sample else None
loss_weights = list(sample["loss_weights"]) if "loss_weights" in sample else None
for start, end in runs[n_complete:]:
for pos in range(start, end + 1):
if pos >= max_length:
break
input_ids[pos] = 0
mm_type_ids[pos] = 0
if labels is not None:
labels[pos] = IGNORE_INDEX
if loss_weights is not None:
loss_weights[pos] = 0.0
sample["input_ids"] = input_ids
sample["mm_token_type_ids"] = mm_type_ids
if labels is not None:
sample["labels"] = labels
if loss_weights is not None:
sample["loss_weights"] = loss_weights
return mm_type_ids
def _align_multimodal_on_truncation(sample: ModelInput, max_length: int) -> ModelInput:
"""Remove orphaned multimodal data when the sequence will be truncated.
When cutoff_len truncates input_ids, media whose placeholder tokens are partially cut lose
their token<->feature correspondence. Trims pixel_values/grid_thw (vision) and
input_features/feature_attention_mask (audio) to the complete items and zeros out orphaned
placeholder tokens so the model ignores them.
"""
mm_type_ids = sample.get("mm_token_type_ids")
if mm_type_ids is None:
return sample
sample = dict(sample)
mm_type_ids = _align_grid_media(
sample, mm_type_ids, max_length, target=1, grid_key="image_grid_thw", pixel_key="pixel_values"
)
mm_type_ids = _align_grid_media(
sample, mm_type_ids, max_length, target=2, grid_key="video_grid_thw", pixel_key="pixel_values_videos"
)
mm_type_ids = _align_audio(sample, mm_type_ids, max_length, target=3)
# Remove empty multimodal fields entirely
if "image_grid_thw" in sample and len(sample["image_grid_thw"]) == 0:
del sample["pixel_values"]
del sample["image_grid_thw"]
if "video_grid_thw" in sample and len(sample["video_grid_thw"]) == 0:
del sample["pixel_values_videos"]
del sample["video_grid_thw"]
if "input_features" in sample and sample["input_features"].shape[0] == 0:
del sample["input_features"]
del sample["feature_attention_mask"]
return sample
def pad_and_truncate(samples: list[ModelInput], max_seqlen: int) -> list[BatchInput]:
max_length = min(max(len(sample["input_ids"]) for sample in samples), max_seqlen)
padded_samples = []
for sample in samples:
# Align multimodal fields before truncation: remove images/videos whose
# placeholder tokens would be partially cut, preventing pixel<->token mismatch.
if len(sample["input_ids"]) > max_length and any(k in sample for k in _MULTIMODAL_PASSTHROUGH_KEYS):
sample = _align_multimodal_on_truncation(sample, max_length)
padded_sample = {}
for key, value in sample.items():
if key in _MULTIMODAL_PASSTHROUGH_KEYS:
padded_sample[key] = value
continue
if "label" in key:
pad_value = IGNORE_INDEX
else:
pad_value = 0
if not isinstance(value, str):
padded_sample[key] = _pad_and_truncate(torch.tensor(value), max_length, pad_value)
else:
padded_sample[key] = value
padded_samples.append(padded_sample)
return padded_samples

View File

@ -14,11 +14,13 @@
import json
import 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

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# 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

View File

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

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# 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,

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# 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>")

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# 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

View File

@ -1,4 +1,4 @@
# Copyright 2025 the LlamaFactory team.
# Copyright 2026 the LlamaFactory team.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# 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):

View File

@ -13,6 +13,7 @@
# limitations under the License.
import os
from types import SimpleNamespace
from typing import TYPE_CHECKING, Any
import numpy as np
@ -417,6 +418,24 @@ def test_qwen2_vl_plugin():
_check_plugin(**check_inputs)
def test_moss_vl_plugin():
messages = [
{"role": "user", "content": "First <image>, finally <image>."},
{"role": "assistant", "content": "Done."},
]
expected_messages = [
{"role": "user", "content": "First <|image_pad|>, finally <|image_pad|>."},
{"role": "assistant", "content": "Done."},
]
processor = SimpleNamespace(image_processor=object(), video_processor=object())
plugin = get_mm_plugin(name="moss_vl", image_token="<|image_pad|>", video_token="<|video_pad|>")
processed_messages = plugin.process_messages(messages, [object(), object()], [], [], processor)
assert processed_messages == expected_messages
assert messages[0]["content"] == "First <image>, finally <image>."
@pytest.mark.runs_on(["cpu", "mps"])
@pytest.mark.skipif(not is_transformers_version_greater_than("4.57.0"), reason="Requires transformers>=4.57.0")
def test_qwen3_vl_plugin():

View File

@ -0,0 +1,631 @@
# Copyright 2025 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.
from types import SimpleNamespace
import pytest
import torch
from PIL import Image
from llamafactory.data.collator import MultiModalDataCollatorForSeq2Seq
from llamafactory.data.mm_plugin import get_mm_plugin
from llamafactory.data.processor.supervised import SupervisedDatasetProcessor
from llamafactory.extras.constants import IGNORE_INDEX
IMAGE_TOKEN_ID = 101
VIDEO_TOKEN_ID = 102
VISION_START_TOKEN_ID = 103
VISION_END_TOKEN_ID = 104
TIME_START_TOKEN_ID = 105
TIME_END_TOKEN_ID = 106
IM_END_TOKEN_ID = 107
class _ImageProcessor:
def __init__(self):
self.calls = []
def __call__(self, images, return_tensors, **kwargs):
self.calls.append({"return_tensors": return_tensors, **kwargs})
values = []
for image in images:
marker = image.getpixel((0, 0))[0] + 1
values.append(torch.full((1, 3), marker, dtype=torch.float32))
return {
"pixel_values": torch.cat(values),
"image_grid_thw": torch.tensor([[1, 1, 1]] * len(images)),
}
class _VideoProcessor:
temporal_patch_size = 1
def __init__(self):
self.calls = []
def __call__(self, videos, return_tensors, return_metadata, **kwargs):
self.calls.append(
{
"return_tensors": return_tensors,
"return_metadata": return_metadata,
**kwargs,
}
)
result = {
"pixel_values_videos": torch.cat(
[torch.full((2, 3), 9 + index, dtype=torch.float32) for index in range(len(videos))]
),
"video_grid_thw": torch.tensor([[2, 1, 1]] * len(videos)),
}
if return_metadata:
result["video_metadata"] = [
SimpleNamespace(frames_indices=[0, 2], total_num_frames=2, fps=2.0, duration=2.0) for _ in videos
]
return result
class _Tokenizer:
pad_token_id = 0
padding_side = "right"
_token_ids = {
"<|time_start|>": TIME_START_TOKEN_ID,
"<|time_end|>": TIME_END_TOKEN_ID,
"<|im_end|>": IM_END_TOKEN_ID,
}
def convert_tokens_to_ids(self, token):
return self._token_ids[token]
def pad(self, features, padding, max_length, pad_to_multiple_of, return_tensors):
del padding, max_length, return_tensors
sequence_length = max(len(feature["input_ids"]) for feature in features)
if pad_to_multiple_of is not None:
sequence_length = ((sequence_length + pad_to_multiple_of - 1) // pad_to_multiple_of) * pad_to_multiple_of
padded = {"input_ids": [], "attention_mask": []}
for feature in features:
pad_length = sequence_length - len(feature["input_ids"])
if self.padding_side == "right":
padded["input_ids"].append(feature["input_ids"] + [self.pad_token_id] * pad_length)
padded["attention_mask"].append(feature["attention_mask"] + [0] * pad_length)
else:
padded["input_ids"].append([self.pad_token_id] * pad_length + feature["input_ids"])
padded["attention_mask"].append([0] * pad_length + feature["attention_mask"])
return {key: torch.tensor(value) for key, value in padded.items()}
class _Processor:
image_token_id = IMAGE_TOKEN_ID
video_token_id = VIDEO_TOKEN_ID
vision_start_token_id = VISION_START_TOKEN_ID
vision_end_token_id = VISION_END_TOKEN_ID
def __init__(self):
self.image_processor = _ImageProcessor()
self.video_processor = _VideoProcessor()
self.tokenizer = _Tokenizer()
@staticmethod
def _calculate_timestamps(*args, **kwargs):
del args, kwargs
return [0.0, 1.0]
def _get_plugin():
return get_mm_plugin(
name="moss_vl",
image_token="<|image_pad|>",
video_token="<|video_pad|>",
vision_bos_token="<|vision_start|>",
vision_eos_token="<|vision_end|>",
time_bos_token="<|time_start|>",
time_eos_token="<|time_end|>",
)
def _video_ids(seed):
return [
VISION_START_TOKEN_ID,
TIME_START_TOKEN_ID,
seed,
TIME_END_TOKEN_ID,
IMAGE_TOKEN_ID,
TIME_START_TOKEN_ID,
seed + 1,
TIME_END_TOKEN_ID,
IMAGE_TOKEN_ID,
VISION_END_TOKEN_ID,
]
def _left_pad(sequences, pad_value):
max_len = max(map(len, sequences))
return torch.tensor([[pad_value] * (max_len - len(sequence)) + sequence for sequence in sequences])
def test_moss_vl_process_messages_expands_video_frames():
plugin = _get_plugin()
processor = _Processor()
messages = [
{"role": "user", "content": "First <image>, then <video>, finally <image>."},
{"role": "assistant", "content": "Done."},
]
images = [Image.new("RGB", (2, 2)), Image.new("RGB", (2, 2), (2, 0, 0))]
processed = plugin.process_messages(messages, images, ["video.mp4"], [], processor)
video_tokens = (
"<|vision_start|>"
"<|time_start|>0.0 seconds<|time_end|><|image_pad|>"
"<|time_start|>1.0 seconds<|time_end|><|image_pad|>"
"<|vision_end|>"
)
assert processed[0]["content"] == (f"First <|image_pad|>, then {video_tokens}, finally <|image_pad|>.")
assert messages[0]["content"] == "First <image>, then <video>, finally <image>."
@pytest.mark.parametrize(
("content", "images", "videos", "error"),
[
("Missing media: <image>.", [], [], "number of images does not match"),
("Missing media: <video>.", [], [], "number of videos does not match"),
],
)
def test_moss_vl_rejects_placeholder_count_mismatch(content, images, videos, error):
plugin = _get_plugin()
with pytest.raises(ValueError, match=error):
plugin.process_messages(
[{"role": "user", "content": content}],
images,
videos,
[],
_Processor(),
)
def test_moss_vl_process_messages_expands_multiple_videos_in_order():
plugin = _get_plugin()
messages = [{"role": "user", "content": "Compare <video> with <video>."}]
processed = plugin.process_messages(messages, [], ["first.mp4", "second.mp4"], [], _Processor())
frame_tokens = (
"<|vision_start|>"
"<|time_start|>0.0 seconds<|time_end|><|image_pad|>"
"<|time_start|>1.0 seconds<|time_end|><|image_pad|>"
"<|vision_end|>"
)
assert processed[0]["content"] == f"Compare {frame_tokens} with {frame_tokens}."
def test_moss_vl_forwards_spatial_pixel_limits_to_native_processors():
plugin = _get_plugin()
processor = _Processor()
processor.image_min_pixels = 1024
processor.image_max_pixels = 262144
processor.video_min_pixels = 256
processor.video_max_pixels = 16384
processor.video_fps = 1.0
processor.video_maxlen = 8
image = Image.new("RGB", (1024, 1024))
plugin.process_messages(
[{"role": "user", "content": "Compare <image> and <video>."}],
[image],
["video.mp4"],
[],
processor,
)
plugin.get_mm_inputs(
[image],
["video.mp4"],
[],
[1],
[1],
[0],
[[IMAGE_TOKEN_ID, *_video_ids(201)]],
processor,
)
assert processor.image_processor.calls == [
{
"return_tensors": "pt",
"min_pixels": 1024,
"max_pixels": 262144,
}
]
assert processor.video_processor.calls == [
{
"return_tensors": "pt",
"return_metadata": True,
"video_fps": 1.0,
"max_frames": 8,
"size": {"shortest_edge": 256, "longest_edge": 16384},
},
{
"return_tensors": "pt",
"return_metadata": False,
"video_fps": 1.0,
"max_frames": 8,
"size": {"shortest_edge": 256, "longest_edge": 16384},
},
]
def test_moss_vl_rejects_invalid_batch_metadata():
plugin = _get_plugin()
processor = _Processor()
image = Image.new("RGB", (2, 2))
with pytest.raises(ValueError, match="batch metadata must have one entry per sample"):
plugin.get_mm_inputs([image], [], [], [1], [], [0], [[IMAGE_TOKEN_ID]], processor)
with pytest.raises(ValueError, match="media lengths do not consume all provided inputs"):
plugin.get_mm_inputs([image], [], [], [0], [0], [0], [[201]], processor)
def test_moss_vl_rejects_truncated_media_tokens():
plugin = _get_plugin()
with pytest.raises(ValueError, match="increase `cutoff_len`"):
plugin.get_mm_inputs(
[Image.new("RGB", (2, 2))],
[],
[],
[1],
[0],
[0],
[[201, 202]],
_Processor(),
)
def test_moss_vl_rejects_incomplete_video_token_block():
plugin = _get_plugin()
truncated_video_ids = _video_ids(201)[:-1]
with pytest.raises(ValueError, match="incomplete video token block"):
plugin.get_mm_inputs(
[],
["video.mp4"],
[],
[0],
[1],
[0],
[truncated_video_ids],
_Processor(),
)
def test_moss_vl_rejects_video_frame_token_count_mismatch():
plugin = _get_plugin()
incomplete_frame_ids = [VISION_START_TOKEN_ID, IMAGE_TOKEN_ID, VISION_END_TOKEN_ID]
with pytest.raises(ValueError, match="video frame tokens do not match"):
plugin.get_mm_inputs(
[],
["video.mp4"],
[],
[0],
[1],
[0],
[incomplete_frame_ids],
_Processor(),
)
def test_moss_vl_media_order_batch_mask_and_labels():
plugin = _get_plugin()
processor = _Processor()
images = [Image.new("RGB", (2, 2)), Image.new("RGB", (2, 2), (2, 0, 0))]
first_ids = [
IMAGE_TOKEN_ID,
201,
VISION_START_TOKEN_ID,
TIME_START_TOKEN_ID,
202,
TIME_END_TOKEN_ID,
IMAGE_TOKEN_ID,
TIME_START_TOKEN_ID,
203,
TIME_END_TOKEN_ID,
IMAGE_TOKEN_ID,
VISION_END_TOKEN_ID,
IMAGE_TOKEN_ID,
204,
]
second_ids = [301, 302]
assert plugin._get_media_order_from_ids(first_ids, processor, 2, 1) == ["image", "video", "image"]
mm_inputs = plugin.get_mm_inputs(
images=images,
videos=["video.mp4"],
audios=[],
imglens=[2, 0],
vidlens=[1, 0],
audlens=[0, 0],
batch_ids=[first_ids, second_ids],
processor=processor,
)
assert mm_inputs["grid_thw"].tolist() == [[1, 1, 1], [2, 1, 1], [1, 1, 1], [1, 1, 1]]
assert mm_inputs["media_nums_per_sample"] == [3, 1]
assert mm_inputs["pixel_values"][:, 0].tolist() == [1.0, 9.0, 9.0, 3.0, 256.0]
pre_padding_mask = mm_inputs["cross_attention_mask"]
assert pre_padding_mask.shape == (2, 1, len(first_ids), 4)
assert pre_padding_mask[0, 0, 0].tolist() == [False, True, True, True]
assert pre_padding_mask[0, 0, 10].tolist() == [False, False, False, True]
assert pre_padding_mask[0, 0, 12].tolist() == [False, False, False, False]
assert pre_padding_mask[1].all()
seq_len = len(first_ids)
input_ids = torch.tensor([first_ids, [0] * (seq_len - 2) + second_ids])
attention_mask = torch.tensor([[1] * seq_len, [0] * (seq_len - 2) + [1, 1]])
labels = input_ids.clone()
labels[attention_mask == 0] = IGNORE_INDEX
features = {
"input_ids": input_ids,
"attention_mask": attention_mask,
"labels": labels,
"position_ids": torch.arange(seq_len).repeat(2, 1),
}
mm_inputs = plugin.post_process_mossvl_inputs(features, mm_inputs, processor)
mask = mm_inputs["cross_attention_mask"]
assert mask.shape == (2, 1, seq_len, 4)
assert mask[0, 0, 0].tolist() == [False, True, True, True]
assert mask[0, 0, 10].tolist() == [False, False, False, True]
assert mask[0, 0, 12].tolist() == [False, False, False, False]
assert mask[1].all()
assert "position_ids" not in features
assert not torch.any((features["input_ids"] == IMAGE_TOKEN_ID) & ~features["attention_mask"].bool())
assert features["labels"][0, 1].item() == 201
assert features["labels"][0, 13].item() == 204
assert features["labels"][0, 0].item() == IGNORE_INDEX
assert features["labels"][0, 12].item() == IGNORE_INDEX
assert torch.all(features["labels"][0, 2:12] == IGNORE_INDEX)
def test_moss_vl_complex_batch_keeps_media_and_masks_sample_local():
plugin = _get_plugin()
processor = _Processor()
batch_ids = [
[IMAGE_TOKEN_ID, 211, IMAGE_TOKEN_ID, 212],
[221, *_video_ids(222), 223, *_video_ids(224), 225],
[IMAGE_TOKEN_ID, 231, *_video_ids(232), 233, IMAGE_TOKEN_ID, 234],
[241, 242, 243],
]
images = [Image.new("RGB", (2, 2), (marker, 0, 0)) for marker in range(4)]
mm_inputs = plugin.get_mm_inputs(
images=images,
videos=["first.mp4", "second.mp4", "third.mp4"],
audios=[],
imglens=[2, 0, 2, 0],
vidlens=[0, 2, 1, 0],
audlens=[0, 0, 0, 0],
batch_ids=batch_ids,
processor=processor,
)
assert mm_inputs["grid_thw"].tolist() == [
[1, 1, 1],
[1, 1, 1],
[2, 1, 1],
[2, 1, 1],
[1, 1, 1],
[2, 1, 1],
[1, 1, 1],
[1, 1, 1],
]
assert mm_inputs["media_nums_per_sample"] == [2, 2, 3, 1]
assert mm_inputs["pixel_values"][:, 0].tolist() == [
1.0,
2.0,
9.0,
9.0,
10.0,
10.0,
3.0,
9.0,
9.0,
4.0,
256.0,
]
input_ids = _left_pad(batch_ids, 0)
attention_mask = _left_pad([[1] * len(ids) for ids in batch_ids], 0)
labels = input_ids.clone()
labels[attention_mask == 0] = IGNORE_INDEX
features = {
"input_ids": input_ids,
"attention_mask": attention_mask,
"labels": labels,
"position_ids": torch.arange(input_ids.shape[1]).repeat(len(batch_ids), 1),
}
plugin.post_process_mossvl_inputs(features, mm_inputs, processor)
cross_mask = mm_inputs["cross_attention_mask"]
assert cross_mask.shape == (4, 1, input_ids.shape[1], 4)
assert (~cross_mask[0]).sum().item() > 0
assert (~cross_mask[1]).sum().item() > 0
assert (~cross_mask[2]).sum().item() > 0
assert cross_mask[3].all()
assert cross_mask[0, ..., 2:].all()
assert not cross_mask[1, ..., :4].all()
assert not cross_mask[2, ..., :4].all()
assert features["labels"][3, -3:].tolist() == [241, 242, 243]
assert torch.all(features["labels"][features["attention_mask"] == 0] == IGNORE_INDEX)
assert "position_ids" not in features
def test_moss_vl_supervised_processor_to_collator_mixed_batch(monkeypatch):
plugin = _get_plugin()
processor = _Processor()
tokenizer = processor.tokenizer
template = SimpleNamespace(mm_plugin=plugin)
dataset_processor = SupervisedDatasetProcessor(
template=template,
tokenizer=tokenizer,
processor=processor,
data_args=SimpleNamespace(),
)
first_ids = [IMAGE_TOKEN_ID, 211, *_video_ids(212), IMAGE_TOKEN_ID, 214]
second_ids = [221, 222]
def encode_example(prompt, **kwargs):
del kwargs
input_ids = first_ids if "<image>" in prompt[0]["content"] else second_ids
return input_ids, input_ids.copy()
monkeypatch.setattr(dataset_processor, "_encode_data_example", encode_example)
examples = {
"_prompt": [
[{"role": "user", "content": "Compare <image>, <video>, and <image>."}],
[{"role": "user", "content": "Text-only question."}],
],
"_response": [
[{"role": "assistant", "content": "Mixed answer."}],
[{"role": "assistant", "content": "Text answer."}],
],
"_system": ["", ""],
"_tools": ["", ""],
"_images": [
[Image.new("RGB", (2, 2)), Image.new("RGB", (2, 2), (2, 0, 0))],
None,
],
"_videos": [["video.mp4"], None],
"_audios": [None, None],
}
model_inputs = dataset_processor.preprocess_dataset(examples)
assert "media_order" not in model_inputs
collator = MultiModalDataCollatorForSeq2Seq(
tokenizer=tokenizer,
model=SimpleNamespace(config=SimpleNamespace(model_type="moss_vl")),
template=template,
processor=processor,
label_pad_token_id=IGNORE_INDEX,
)
features = [
{key: values[index] for key, values in model_inputs.items()} for index in range(len(model_inputs["input_ids"]))
]
batch = collator(features)
assert batch["grid_thw"].tolist() == [[1, 1, 1], [2, 1, 1], [1, 1, 1], [1, 1, 1]]
assert batch["media_nums_per_sample"] == [3, 1]
assert batch["pixel_values"][:, 0].tolist() == [1.0, 9.0, 9.0, 3.0, 256.0]
assert batch["cross_attention_mask"].shape == (2, 1, len(first_ids), 4)
assert batch["cross_attention_mask"][1].all()
assert torch.all(batch["labels"][1, len(second_ids) :] == IGNORE_INDEX)
assert "position_ids" not in batch
def test_moss_vl_generate_collator_keeps_left_padded_cross_attention_mask():
plugin = _get_plugin()
processor = _Processor()
processor.tokenizer.padding_side = "left"
template = SimpleNamespace(mm_plugin=plugin)
batch_ids = [
[IMAGE_TOKEN_ID, 211],
[301, IMAGE_TOKEN_ID, 302, 303],
]
features = [
{
"input_ids": input_ids,
"attention_mask": [1] * len(input_ids),
"labels": input_ids.copy(),
"images": [Image.new("RGB", (2, 2))],
}
for input_ids in batch_ids
]
collator = MultiModalDataCollatorForSeq2Seq(
tokenizer=processor.tokenizer,
model=SimpleNamespace(config=SimpleNamespace(model_type="moss_vl")),
template=template,
processor=processor,
label_pad_token_id=IGNORE_INDEX,
pad_to_multiple_of=8,
)
batch = collator(features)
assert batch["cross_attention_mask"].shape == (2, 1, 8, 1)
assert batch["cross_attention_mask"][0, 0, :, 0].tolist() == [True] * 6 + [False, False]
assert batch["cross_attention_mask"][1, 0, :, 0].tolist() == [True] * 5 + [False, False, False]
def test_moss_vl_predict_collator_uses_precomputed_cross_attention_mask_without_model():
plugin = _get_plugin()
processor = _Processor()
template = SimpleNamespace(mm_plugin=plugin)
batch_ids = [
[IMAGE_TOKEN_ID, 211],
[301, IMAGE_TOKEN_ID, 302, 303],
]
features = [
{
"input_ids": input_ids,
"attention_mask": [1] * len(input_ids),
"labels": input_ids.copy(),
"images": [Image.new("RGB", (2, 2))],
}
for input_ids in batch_ids
]
collator = MultiModalDataCollatorForSeq2Seq(
tokenizer=processor.tokenizer,
model=None,
template=template,
processor=processor,
label_pad_token_id=IGNORE_INDEX,
)
batch = collator(features)
assert batch["cross_attention_mask"].shape == (2, 1, 4, 1)
assert batch["cross_attention_mask"][0, 0, :, 0].tolist() == [False, False, True, True]
assert batch["cross_attention_mask"][1, 0, :, 0].tolist() == [True, False, False, False]
def test_moss_vl_masks_only_the_token_after_im_end():
plugin = _get_plugin()
processor = _Processor()
input_ids = torch.tensor([[301, IM_END_TOKEN_ID, 302, 303]])
features = {
"input_ids": input_ids,
"attention_mask": torch.ones_like(input_ids),
"labels": input_ids.clone(),
}
mm_inputs = plugin.get_mm_inputs([], [], [], [0], [0], [0], [input_ids[0].tolist()], processor)
plugin.post_process_mossvl_inputs(features, mm_inputs, processor)
assert features["labels"].tolist() == [[301, IM_END_TOKEN_ID, IGNORE_INDEX, 303]]
def test_moss_vl_native_text_dummy_shape_and_values():
plugin = _get_plugin()
processor = _Processor()
processor.image_processor = SimpleNamespace(patch_size=16, temporal_patch_size=1, merge_size=2)
mm_inputs = plugin.get_mm_inputs([], [], [], [0], [0], [0], [[301]], processor)
assert mm_inputs["grid_thw"].tolist() == [[1, 8, 8]]
assert mm_inputs["pixel_values"].shape == (64, 768)
assert torch.count_nonzero(mm_inputs["pixel_values"]).item() == 0
assert mm_inputs["media_nums_per_sample"] == [1]

View File

@ -0,0 +1,39 @@
# Copyright 2025 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.
from pathlib import Path
import yaml
ROOT = Path(__file__).resolve().parents[2]
def test_moss_vl_training_configs_are_unpacked_and_additive():
lora = yaml.safe_load((ROOT / "examples/train_lora/mossvl_lora_sft.yaml").read_text())
full = yaml.safe_load((ROOT / "examples/train_full/mossvl_full_sft.yaml").read_text())
assert lora["template"] == full["template"] == "moss_vl"
assert lora["packing"] is full["packing"] is False
assert lora["per_device_train_batch_size"] == 2
assert full["per_device_train_batch_size"] == 1
assert full["use_reentrant_gc"] is False
assert full["gradient_checkpointing"] is True
assert full["gradient_checkpointing_kwargs"] == {"use_reentrant": False}
for config in (lora, full):
assert config["model_name_or_path"] == "OpenMOSS-Team/MOSS-VL-Instruct-0708"
assert not any("/inspire/" in str(value) or "/tmp/" in str(value) for value in config.values())
assert config["freeze_vision_tower"] is True
assert config["freeze_multi_modal_projector"] is True
assert config["freeze_language_model"] is False

View File

@ -19,7 +19,13 @@ import pytest
from transformers import AutoTokenizer
from llamafactory.data import get_template_and_fix_tokenizer
from llamafactory.data.template import parse_template
from llamafactory.data.template import TEMPLATES, parse_template
from llamafactory.extras.constants import (
DEFAULT_TEMPLATE,
MULTIMODAL_SUPPORTED_MODELS,
SUPPORTED_MODELS,
DownloadSource,
)
from llamafactory.extras.packages import is_transformers_version_greater_than
from llamafactory.hparams import DataArguments
@ -91,6 +97,22 @@ def _check_template(
_check_tokenization(tokenizer, (prompt_ids, answer_ids), (prompt_str, answer_str))
def test_moss_vl_registration():
model_name = "MOSS-VL-Instruct-0708"
assert model_name in SUPPORTED_MODELS
assert SUPPORTED_MODELS[model_name][DownloadSource.DEFAULT] == "OpenMOSS-Team/MOSS-VL-Instruct-0708"
assert DEFAULT_TEMPLATE[model_name] == "moss_vl"
assert model_name in MULTIMODAL_SUPPORTED_MODELS
assert TEMPLATES["moss_vl"].mm_plugin.__class__.__name__ == "MossVLPlugin"
assert TEMPLATES["moss_vl"].mm_plugin.image_token == "<|image_pad|>"
assert TEMPLATES["moss_vl"].mm_plugin.video_token == "<|video_pad|>"
assert TEMPLATES["moss_vl"].mm_plugin.vision_bos_token == "<|vision_start|>"
assert TEMPLATES["moss_vl"].mm_plugin.vision_eos_token == "<|vision_end|>"
assert TEMPLATES["moss_vl"].mm_plugin.time_bos_token == "<|time_start|>"
assert TEMPLATES["moss_vl"].mm_plugin.time_eos_token == "<|time_end|>"
@pytest.mark.runs_on(["cpu", "mps"])
def test_encode_oneturn():
tokenizer = AutoTokenizer.from_pretrained(TINY_LLAMA3)

View File

@ -13,6 +13,7 @@
# limitations under the License.
import os
from types import SimpleNamespace
import pytest
import torch
@ -21,7 +22,131 @@ from transformers import AutoConfig, AutoModelForImageTextToText
from llamafactory.extras.packages import is_transformers_version_greater_than
from llamafactory.hparams import FinetuningArguments, ModelArguments
from llamafactory.model.adapter import init_adapter
from llamafactory.model.adapter import _setup_freeze_tuning, _setup_full_tuning, init_adapter
from llamafactory.model.model_utils.misc import find_all_linear_modules
from llamafactory.model.model_utils.visual import COMPOSITE_MODELS, autocast_projector_dtype, patch_target_modules
class _MossVLFixture(torch.nn.Module):
def __init__(self) -> None:
super().__init__()
self.config = SimpleNamespace(
model_type="moss_vl",
text_config=SimpleNamespace(num_hidden_layers=2),
)
self.model = torch.nn.Module()
self.model.separator_token = torch.nn.Parameter(torch.empty(4))
self.model.visual = torch.nn.Module()
self.model.visual.pos_embed = torch.nn.Embedding(4, 4)
self.model.visual.patch_embed = torch.nn.Module()
self.model.visual.patch_embed.proj = torch.nn.Linear(4, 4)
self.model.visual.blocks = torch.nn.ModuleList([self._make_block(), self._make_block()])
self.model.visual.merger = torch.nn.Module()
self.model.visual.merger.linear_fc1 = torch.nn.Linear(4, 4)
self.model.language_model = torch.nn.Module()
self.model.language_model.layers = torch.nn.ModuleList([self._make_layer(), self._make_layer()])
self.lm_head = torch.nn.Linear(4, 4)
@staticmethod
def _make_block() -> torch.nn.Module:
block = torch.nn.Module()
block.attn = torch.nn.Module()
block.attn.qkv = torch.nn.Linear(4, 4)
return block
@staticmethod
def _make_layer() -> torch.nn.Module:
layer = torch.nn.Module()
layer.self_attn = torch.nn.Module()
layer.self_attn.q_proj = torch.nn.Linear(4, 4)
return layer
@pytest.mark.parametrize("freeze_vision_tower", (False, True))
@pytest.mark.parametrize("freeze_multi_modal_projector", (False, True))
@pytest.mark.parametrize("freeze_language_model", (False, True))
def test_moss_vl_full(
freeze_vision_tower: bool,
freeze_multi_modal_projector: bool,
freeze_language_model: bool,
):
model = _MossVLFixture()
finetuning_args = FinetuningArguments(
finetuning_type="full",
freeze_vision_tower=freeze_vision_tower,
freeze_multi_modal_projector=freeze_multi_modal_projector,
freeze_language_model=freeze_language_model,
)
_setup_full_tuning(model, finetuning_args, is_trainable=True, cast_trainable_params_to_fp32=False)
for name, param in model.named_parameters():
if name.startswith("model.visual.merger") or name == "model.separator_token":
assert param.requires_grad != freeze_multi_modal_projector
elif name.startswith("model.visual"):
assert param.requires_grad != freeze_vision_tower
else:
assert param.requires_grad != freeze_language_model
@pytest.mark.parametrize("freeze_multi_modal_projector", (False, True))
def test_moss_vl_freeze(freeze_multi_modal_projector: bool):
model = _MossVLFixture()
finetuning_args = FinetuningArguments(
finetuning_type="freeze",
freeze_trainable_layers=1,
freeze_vision_tower=True,
freeze_multi_modal_projector=freeze_multi_modal_projector,
freeze_language_model=False,
)
_setup_freeze_tuning(model, finetuning_args, is_trainable=True, cast_trainable_params_to_fp32=False)
assert model.model.separator_token.requires_grad != freeze_multi_modal_projector
assert model.model.visual.merger.linear_fc1.weight.requires_grad != freeze_multi_modal_projector
assert model.model.visual.patch_embed.proj.weight.requires_grad is False
assert model.model.language_model.layers[0].self_attn.q_proj.weight.requires_grad is False
assert model.model.language_model.layers[1].self_attn.q_proj.weight.requires_grad is True
@pytest.mark.parametrize("freeze_vision_tower", (False, True))
def test_moss_vl_lora_target_all(freeze_vision_tower: bool):
model = _MossVLFixture()
finetuning_args = FinetuningArguments(
finetuning_type="lora",
lora_target="all",
freeze_vision_tower=freeze_vision_tower,
freeze_multi_modal_projector=True,
freeze_language_model=False,
)
target_modules = find_all_linear_modules(model, freeze_vision_tower)
target_modules = patch_target_modules(model, finetuning_args, target_modules)
assert any(name.startswith("model.language_model") and name.endswith("q_proj") for name in target_modules)
assert any(name.startswith("model.visual.blocks") and name.endswith("qkv") for name in target_modules) != (
freeze_vision_tower
)
assert all("patch_embed" not in name for name in target_modules)
assert all("merger" not in name for name in target_modules)
assert all("lm_head" not in name for name in target_modules)
def test_moss_vl_projector_modules():
model = _MossVLFixture()
composite_model = COMPOSITE_MODELS["moss_vl"]
assert composite_model.projector_keys == ["model.visual.merger", "model.separator_token"]
assert composite_model.get_projectors(model) == [model.model.visual.merger]
def test_moss_vl_quantized_projector_hook_skips_parameter():
model = _MossVLFixture()
model.quantization_method = "bitsandbytes"
autocast_projector_dtype(model, SimpleNamespace(compute_dtype=torch.float16))
assert len(model.model.visual.merger._forward_hooks) == 1
@pytest.mark.parametrize("freeze_vision_tower", (False, True))

View File

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

View File

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