指南

多模态

此为预览文档,多模态功能将在未来发布

PyTRIO 支持图像输入,实现多模态推理与训练。

输入处理

图像输入与文本输入在 PyTRIO 中均使用 ModelInput 作为封装。

区别主要在于图像输入需要使用 Chunk(块)来构建,而文本输入往往只需用from_ints()方法。

在多模态输入中,ModelInput 由一个或多个 chunk 组成。每个 chunk 表示 prompt 中的一段内容,模型会按照 chunks 列表中的顺序读取它们。

比如在本文的图片描述任务中,输入由三部分组成:

  1. 图片前的文本 token,其中包含用户消息和视觉输入的开始标记
  2. 图片数据
  3. 图片后的文本 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(...),本质上是构造只有一个 EncodedTextChunkModelInput。下面两种写法表达的是同一种纯文本输入:

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:图片格式,当前示例使用 pngjpeg
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
  • 图片格式必须与 ImageChunkformat 参数一致
  • 图片前后的特殊 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 文本以及右移后的 completion
  • target_tokens:模型在每个输入位置需要预测的 token
  • weights:每个位置的损失权重,prompt 和图片部分为 0,completion 部分为 1

构建 ImageChunk

推理时只需传入图片数据和格式;构建训练数据时,还需要通过 expected_tokens 告诉 PyTRIO 视觉编码器将为这张图片生成多少个 token。这样 PyTRIO 才能正确计算 ModelInput 的长度,并让 target_tokensweights 与输入一一对齐。

以 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 为 completionmodel_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-4BLaTeX_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()
这篇文档对你有帮助吗?

本页目录