多模态
此为预览文档,多模态功能将在未来发布
PyTRIO 支持图像输入,实现多模态推理与训练。
输入处理
图像输入与文本输入在 PyTRIO 中均使用 ModelInput 作为封装。
区别主要在于图像输入需要使用 Chunk(块)来构建,而文本输入往往只需用from_ints()方法。
在多模态输入中,ModelInput 由一个或多个 chunk 组成。每个 chunk 表示 prompt 中的一段内容,模型会按照 chunks 列表中的顺序读取它们。
比如在本文的图片描述任务中,输入由三部分组成:
- 图片前的文本 token,其中包含用户消息和视觉输入的开始标记
- 图片数据
- 图片后的文本 token,其中包含视觉输入的结束标记、用户问题和 assistant 的生成起点
对应的结构如下:
prompt = trio.ModelInput(
chunks=[
trio.types.EncodedTextChunk(tokens=...),
trio.ImageChunk(data=Path(), format="png"),
trio.types.EncodedTextChunk(tokens=...),
]
)chunk 的排列顺序就是模型实际看到的内容顺序,因此图片前后的特殊 token 不能随意调换。
文本块(EncodedTextChunk)
EncodedTextChunk 用于存放已经经过 tokenizer 编码的文本 token。在一些场景中,它也可以理解为 TextChunk:文本需要先转换成 token ids,再放入 tokens 字段。
tokenizer = sampler.get_tokenizer()
tokens = tokenizer.encode("请描述这张图片。", add_special_tokens=False)
text_chunk = trio.types.EncodedTextChunk(tokens=tokens)纯文本推理中常用的 ModelInput.from_ints(...),本质上是构造只有一个 EncodedTextChunk 的 ModelInput。下面两种写法表达的是同一种纯文本输入:
input_ids = tokenizer.encode("你好")
# 简写方式
prompt = trio.ModelInput.from_ints(input_ids)
# 显式使用文本 chunk
prompt = trio.ModelInput(
chunks=[trio.types.EncodedTextChunk(tokens=input_ids)]
)当 prompt 中还包含图片时,就需要使用第二种写法,显式组织不同类型的 chunk。
图片块(ImageChunk)
ImageChunk 用于传递图片的二进制内容,包含两个主要字段:
data:图片的二进制数据,可以使用Path.read_bytes()读取format:图片格式,当前示例使用png或jpeg
from pathlib import Path
image_path = Path("example.png")
image_chunk = trio.ImageChunk(
data=image_path.read_bytes(),
format="png",
)这里传入的是原始图片字节,不需要提前将图片转换成 token,也不需要手动进行 Base64 编码。图片的格式应与实际文件内容一致,例如 .jpg 和 .jpeg 文件都使用 jpeg。
多模态推理
首先,创建一个使用视觉语言模型的推理客户端,并获取对应的 tokenizer:
import pytrio as trio
sampler = trio.ServiceClient().create_sampling_client(
base_model="Qwen/Qwen3.5-4B"
)
tokenizer = sampler.get_tokenizer()
encode = lambda text: tokenizer.encode(text, add_special_tokens=False)然后,将文本和图片按模型要求的顺序组装为 ModelInput:
prompt = trio.ModelInput(
chunks=[
trio.types.EncodedTextChunk(
tokens=encode("<|im_start|>user\n<|vision_start|>")
),
trio.ImageChunk(
data=image_path.read_bytes(),
format=image_format,
),
trio.types.EncodedTextChunk(
tokens=encode(
"<|vision_end|>请描述这张图片。<|im_end|>\n"
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
)
),
]
)<|vision_start|> 和 <|vision_end|> 用于标记图片在对话中的位置;<|im_start|> 和 <|im_end|> 用于标记对话角色消息的边界。这些特殊 token 属于模型的 prompt 格式,不同模型可能使用不同的模板,应以所用模型的要求为准。
最后,和文本推理一样调用 sample,并通过 .result() 获取远程推理结果:
response = sampler.sample(
prompt=prompt,
num_samples=1,
sampling_params=trio.SamplingParams(
max_tokens=512,
temperature=0.5,
stop="<|im_end|>",
),
).result()
print(response.sequences[0].text)完整示例
将下面的代码保存为 vlm_sample.py,并传入一张本地图片:
python vlm_sample.py /path/to/image.png完整代码如下:
"""Usage:python vlm_sample.py /path/to/image.png"""
import sys
from pathlib import Path
import pytrio as trio
if len(sys.argv) != 2:
raise SystemExit(f"Usage:python {Path(__file__).name} /path/to/image.png")
image_path = Path(sys.argv[1]).expanduser()
image_format = {".png": "png", ".jpg": "jpeg", ".jpeg": "jpeg"}.get(
image_path.suffix.lower()
)
if image_format is None:
raise SystemExit("仅支持 PNG/JPEG 图片")
sampler = trio.ServiceClient().create_sampling_client(
base_model="Qwen/Qwen3.5-4B"
)
tokenizer = sampler.get_tokenizer()
encode = lambda text: tokenizer.encode(text, add_special_tokens=False)
prompt = trio.ModelInput(
chunks=[
trio.types.EncodedTextChunk(
tokens=encode("<|im_start|>user\n<|vision_start|>")
),
trio.ImageChunk(
data=image_path.read_bytes(),
format=image_format,
),
trio.types.EncodedTextChunk(
tokens=encode(
"<|vision_end|>请描述这张图片。<|im_end|>\n"
"<|im_start|>assistant\n<think>\n\n</think>\n\n"
)
),
]
)
response = sampler.sample(
prompt=prompt,
num_samples=1,
sampling_params=trio.SamplingParams(
max_tokens=512,
temperature=0.5,
stop="<|im_end|>",
),
).result()
print(response.sequences[0].text)注意事项
base_model必须是支持图片输入的多模态模型;普通文本模型无法处理ImageChunk- 图片格式必须与
ImageChunk的format参数一致 - 图片前后的特殊 token 取决于模型模板,切换模型时需要同步调整
EncodedTextChunk接收的是 token ids,而不是未经编码的字符串- 多模态推理的采样参数和返回结构与文本推理一致,输出文本仍可通过
response.sequences[0].text获取
多模态训练
为了方便解释,本节主要以SFT为例
Datum 构建
多模态训练和纯文本训练在代码上只有一个地方不同,那就是 Datum 的构建。
二者使用相同的 cross_entropy 损失函数,也都通过 weights 让模型只学习 assistant 的回复。区别在于,纯文本 SFT 可以把所有 token 放进一个 EncodedTextChunk 或直接使用from_inits()方法,而多模态 SFT 需要在 ModelInput 中同时保留文本 chunk 和图片 chunk。
一个多模态 SFT Datum 仍然包含以下三部分:
model_input:图片、prompt 文本以及右移后的 completiontarget_tokens:模型在每个输入位置需要预测的 tokenweights:每个位置的损失权重,prompt 和图片部分为0,completion 部分为1
构建 ImageChunk
推理时只需传入图片数据和格式;构建训练数据时,还需要通过 expected_tokens 告诉 PyTRIO 视觉编码器将为这张图片生成多少个 token。这样 PyTRIO 才能正确计算 ModelInput 的长度,并让 target_tokens、weights 与输入一一对齐。
以 Qwen3.5-4B 为例,可以使用对应的 image processor 计算图片 patch 数量:
import io
from PIL import Image
from transformers import AutoImageProcessor
processor_source = getattr(tokenizer, "name_or_path", "Qwen/Qwen3.5-4B")
image_processor = AutoImageProcessor.from_pretrained(
processor_source,
use_fast=False,
)
def encode_image(image: Image.Image, processor) -> trio.ImageChunk:
image = image.convert("RGB")
buffer = io.BytesIO()
image.save(buffer, format="PNG")
patches = processor.get_number_of_image_patches(
image.height,
image.width,
images_kwargs={},
)
expected_tokens = patches // processor.merge_size**2
return trio.ImageChunk(
data=buffer.getvalue(),
format="png",
expected_tokens=expected_tokens,
)expected_tokens 与具体模型的视觉预处理方式有关。切换模型时,应使用与远程模型一致的 image processor 进行计算,不要直接写死。
构建 Prompt Chunks
建议先使用 tokenizer 的 chat template 生成模型要求的完整 prompt,再把其中的图片占位符替换为真实的 ImageChunk。
比如 Qwen3.5 的图片占位符是 <|image_pad|>。对包含一张图片的样本,模板生成的 prompt 可以被拆成“图片前文本”和“图片后文本”两部分:
IMAGE_PAD = "<|image_pad|>"
PROMPT = (
"Transcribe the mathematical formula in this image into LaTeX. "
"Output only the LaTeX."
)
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": PROMPT},
{"type": "image", "image": "formula"},
],
}
]
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
parts = prompt.split(IMAGE_PAD)
if len(parts) != 2:
raise ValueError(f"期望 1 个图片占位符,实际得到 {len(parts) - 1} 个")
before_image, after_image = parts
prompt_chunks = [
trio.types.EncodedTextChunk(
tokens=tokenizer.encode(before_image, add_special_tokens=False)
),
image_chunk,
trio.types.EncodedTextChunk(
tokens=tokenizer.encode(after_image, add_special_tokens=False)
),
]chat template 已经加入了对话和图片所需的特殊 token,因此再次编码文本时需要设置 add_special_tokens=False,避免重复添加特殊 token。
自回归右移与 Loss Mask
假设模型需要学习的答案 token 为 completion,model_input 只拼接 completion[:-1]。第一个答案 token 由 prompt 的最后一个位置预测,后续答案 token 则由前一个答案 token 预测。
target_tokens 在不参与训练的位置使用 0 占位,而不是 HuggingFace 中常见的 -100。真正决定一个位置是否参与损失计算的是 weights。
完整的 process_example 如下,其中数据集的 image 字段是 PIL 图片,text 字段是目标 LaTeX:
import numpy as np
def process_example(example, tokenizer, processor) -> trio.Datum:
image_chunk = encode_image(example["image"], processor)
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": PROMPT},
{"type": "image", "image": "formula"},
],
}
]
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
parts = prompt.split(IMAGE_PAD)
if len(parts) != 2:
raise ValueError(f"期望 1 个图片占位符,实际得到 {len(parts) - 1} 个")
before_image, after_image = parts
prompt_chunks = [
trio.types.EncodedTextChunk(
tokens=tokenizer.encode(before_image, add_special_tokens=False)
),
image_chunk,
trio.types.EncodedTextChunk(
tokens=tokenizer.encode(after_image, add_special_tokens=False)
),
]
prompt_length = len(trio.ModelInput(chunks=prompt_chunks))
completion = tokenizer.encode(
str(example["text"]).strip() + "<|im_end|>",
add_special_tokens=False,
)
model_input = trio.ModelInput(
chunks=[
*prompt_chunks,
trio.types.EncodedTextChunk(tokens=completion[:-1]),
]
)
target_tokens = np.zeros(len(model_input), dtype=np.int64)
weights = np.zeros(len(model_input), dtype=np.float32)
start = prompt_length - 1
target_tokens[start : start + len(completion)] = completion
weights[start : start + len(completion)] = 1.0
return trio.Datum(
model_input=model_input,
loss_fn_inputs={
"target_tokens": target_tokens,
"weights": weights,
},
)这里把 <|im_end|> 也加入了 completion,并将其权重设为 1,让模型同时学习在答案结束时停止生成。
构建好 Datum 后,训练过程和纯文本 SFT 完全相同:
processed_examples = [
process_example(example, tokenizer, image_processor)
for example in dataset
]
fwdbwd = training_client.forward_backward(
processed_examples,
loss_fn="cross_entropy",
)
optim = training_client.optim_step(
trio.AdamParams(learning_rate=1e-4)
)
result = fwdbwd.result()
optim.result()完整示例
这是一个使用Qwen3.5-4B在 LaTeX_OCR 数据集上进行多模态微调的案例:
"""用 PyTRIO 在 LaTeX_OCR/small 上对 Qwen3.5-4B 做多模态 LoRA SFT。
环境:
pip install pytrio numpy datasets torch torchvision pillow
运行:
python train.py --epochs 3 --batch-size 2
"""
from __future__ import annotations
import argparse
import io
from pathlib import Path
import numpy as np
import pytrio as trio
from datasets import load_dataset
from huggingface_hub import snapshot_download
from PIL import Image
from transformers import AutoImageProcessor
DATASET_ID = "linxy/LaTeX_OCR"
IMAGE_PAD = "<|image_pad|>"
PROMPT = (
"Transcribe the mathematical formula in this image into LaTeX. "
"Output only the LaTeX."
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument("--model", default="Qwen/Qwen3.5-4B")
parser.add_argument("--epochs", type=int, default=1)
parser.add_argument("--batch-size", type=int, default=1)
parser.add_argument("--learning-rate", type=float, default=1e-4)
parser.add_argument("--rank", type=int, default=32)
parser.add_argument("--max-samples", type=int, default=0, help="0 表示全部")
parser.add_argument("--max-length", type=int, default=8192)
parser.add_argument("--dataset-dir", default="data/LaTeX_OCR")
parser.add_argument("--seed", type=int, default=42)
parser.add_argument("--checkpoint-name", default="latex-ocr-small-sft")
return parser.parse_args()
def load_train_dataset(dataset_dir: str):
"""首次运行时下载 small 子集,之后只读取本地 parquet。"""
local_dir = Path(dataset_dir).expanduser().resolve()
files = sorted((local_dir / "small").glob("train-*.parquet"))
if not files:
snapshot_download(
repo_id=DATASET_ID,
repo_type="dataset",
local_dir=local_dir,
allow_patterns=["README.md", "small/*.parquet"],
)
files = sorted((local_dir / "small").glob("train-*.parquet"))
if not files:
raise FileNotFoundError(f"本地没有找到 train parquet:{local_dir}")
return load_dataset(
"parquet",
data_files={"train": [str(path) for path in files]},
split="train",
)
def encode_image(image: Image.Image, processor) -> trio.ImageChunk:
"""把 PIL 图片转换为 PyTRIO 的多模态 ImageChunk。"""
image = image.convert("RGB")
buffer = io.BytesIO()
image.save(buffer, format="PNG")
# TRIO 需要提前知道视觉编码器将产生多少个 token。
patches = processor.get_number_of_image_patches(
image.height,
image.width,
images_kwargs={},
)
return trio.ImageChunk(
data=buffer.getvalue(),
format="png",
expected_tokens=patches // processor.merge_size**2,
)
def process_example(example, tokenizer, processor) -> trio.Datum:
"""将一条 image/text 样本转换为只训练答案部分的 SFT Datum。"""
image_chunk = encode_image(example["image"], processor)
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": PROMPT},
{"type": "image", "image": "formula"},
],
}
]
prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True,
enable_thinking=False,
)
# apply_chat_template 先生成图片占位符,再替换为真实 ImageChunk。
parts = prompt.split(IMAGE_PAD)
if len(parts) != 2:
raise ValueError(f"期望 1 个图片占位符,实际得到 {len(parts) - 1} 个")
before_image, after_image = parts
prompt_chunks = [
trio.types.EncodedTextChunk(
tokens=tokenizer.encode(before_image, add_special_tokens=False)
),
image_chunk,
trio.types.EncodedTextChunk(
tokens=tokenizer.encode(after_image, add_special_tokens=False)
),
]
prompt_length = len(trio.ModelInput(chunks=prompt_chunks))
completion = tokenizer.encode(
str(example["text"]).strip() + "<|im_end|>",
add_special_tokens=False,
)
# 自回归右移:prompt 最后一个位置开始预测 completion。
model_input = trio.ModelInput(
chunks=[
*prompt_chunks,
trio.types.EncodedTextChunk(tokens=completion[:-1]),
]
)
target_tokens = np.zeros(len(model_input), dtype=np.int64)
weights = np.zeros(len(model_input), dtype=np.float32)
start = prompt_length - 1
target_tokens[start : start + len(completion)] = completion
weights[start : start + len(completion)] = 1.0
return trio.Datum(
model_input=model_input,
loss_fn_inputs={"target_tokens": target_tokens, "weights": weights},
)
def loss_per_token(result, batch: list[trio.Datum]) -> float:
"""使用服务返回的 logprobs 计算有监督 token 的平均 NLL。"""
logprobs = np.concatenate(
[output["logprobs"].tolist() for output in result.loss_fn_outputs]
)
weights = np.concatenate(
[datum.loss_fn_inputs["weights"].tolist() for datum in batch]
)
return float(-np.dot(logprobs, weights) / weights.sum())
def main() -> None:
args = parse_args()
if args.batch_size < 1:
raise ValueError("batch-size 必须大于 0")
# 1. 下载并从本地加载 small/train 数据集。
dataset = load_train_dataset(args.dataset_dir)
if args.max_samples > 0:
dataset = dataset.select(range(min(args.max_samples, len(dataset))))
# 2. 与 TRIO 建立连接并创建 LoRA 训练客户端。
service_client = trio.ServiceClient()
training_client = service_client.create_lora_training_client(
base_model=args.model,
rank=args.rank,
seed=args.seed,
)
# 3. 获取与远程模型一致的 tokenizer 和图片处理器。
tokenizer = training_client.get_tokenizer()
processor_source = getattr(tokenizer, "name_or_path", args.model)
image_processor = AutoImageProcessor.from_pretrained(
processor_source,
use_fast=False,
)
# 4. 数据集很小,一次性转换为 PyTRIO Datum,避免每个 epoch 重复编码。
processed_examples = [
process_example(example, tokenizer, image_processor) for example in dataset
]
processed_examples = [
datum
for datum in processed_examples
if len(datum.model_input) <= args.max_length
]
if not processed_examples:
raise RuntimeError("没有可训练样本,请检查 max-samples/max-length")
# 5. 每个 epoch 打乱 Datum,然后按 batch_size 切片训练。
step = 0
for epoch in range(args.epochs):
indices = np.random.default_rng(args.seed + epoch).permutation(
len(processed_examples)
)
for start in range(0, len(indices), args.batch_size):
batch_indices = indices[start : start + args.batch_size]
batch = [processed_examples[index] for index in batch_indices]
fwdbwd = training_client.forward_backward(batch, "cross_entropy")
optim = training_client.optim_step(
trio.AdamParams(learning_rate=args.learning_rate)
)
result = fwdbwd.result()
optim.result()
step += 1
loss = loss_per_token(result, batch)
print(
f"epoch={epoch + 1} step={step} loss={loss:.4f}",
flush=True,
)
# 7. 保存可直接传给 SamplingClient 的推理权重。
saved = training_client.save_weights_for_sampler(
name=args.checkpoint_name
).result()
print(f"saved_weights={saved.path}")
if __name__ == "__main__":
main()