案例

Vision GRPO

分类:多模态 RL;数据集:GeoQA;Accuracy 71.0% → 87.0%

在固定 100 道 GeoQA test 题上,Qwen/Qwen3.5-4B Base Model 的 Accuracy / Format rate 为 71.0% / 75.0%,Vision GRPO step-100 checkpoint 为 87.0% / 91.0%

20-step Vision GRPO 会话的 PyTRIO token 消耗

介绍

Vision GRPO 把 GRPO 的组内相对策略更新扩展到图片输入。本案例使用 GeoQA 的中文几何选择题:模型同时读取题目、四个候选项和几何图片,经过简短推理后输出 \boxed{A}\boxed{B}\boxed{C}\boxed{D}

对同一道题,当前 LoRA policy 会采样一组回答。规则 reward 判断最终选项是否正确,再用组内均值构造 advantage:

Ai=ri1Gj=1GrjA_i = r_i - \frac{1}{G}\sum_{j=1}^{G}r_j

一次训练 step 的数据流如下:

GeoQA 题目 + 四个选项 + 图片
→ 图文 chat template
→ 同题 group rollout
→ boxed choice reward
→ group-relative advantage
→ multimodal Datum
→ importance_sampling
→ LoRA update

图片和题目 token 为 completion 提供上下文,它们的 target、old logprob 和 advantage 都填零。训练信号只落在模型实际生成的 completion token 上。

实验配置

项目本案例配置
Base ModelQwen/Qwen3.5-4B
PyTRIO0.2.7
LoRA rank32
训练数据GeoQA 原始 train,3,503 条
固定评测集GeoQA 原始 test 中按 seed=42 固定的 100 条
参考训练100 steps
每个 step8 道题
每道题8 条 completion
最大 completion1,024 tokens
Thinking modeenable_thinking=False
Reward正确的 \boxed{A-D} 得 1,否则得 0
Advantagereward - group_mean
Lossimportance_sampling
Learning rate4e-5
Checkpoint每 25 step 保存 sampler weights + training state

代码结构

Vision GRPO 由三个脚本组成:

vision-grpo/
├── download-dataset.py   # 下载 GeoQA,生成训练集与固定测试集
├── train.py              # 图文 group rollout、reward、GRPO 更新与 checkpoint
└── eval.py               # 单个 Base / checkpoint 的异步固定集评测

下面先拆解决定多模态输入、GRPO 对齐和异步调用边界的关键代码,随后给出三个完整脚本。

环境与数据

本地机器负责准备图片、组织训练循环和记录实验。多模态采样、LoRA 前向反向与参数更新由 PyTRIO 服务执行。

使用 Python 3.13 或更高版本,准备一个空目录并安装运行所需依赖:

mkdir vision-grpo
cd vision-grpo

python -m venv .venv
source .venv/bin/activate
python -m pip install \
  "pytrio==0.2.7" \
  "datasets>=5.0.0" \
  huggingface_hub numpy pillow swanlab tqdm transformers \
  torch torchvision

trio login
swanlab login

GeoQA 源数据共有 5,010 条,原始 train / test / dev 分别为 3,503 / 759 / 748 条。每条数据包含题目、四个候选项、正确选项、几何知识点、解析和图片。

GeoQA 数据集字段与图文样本

核心逻辑

1. 把 chat template 拆成图文 chunks

train.py 先用模型的 chat template 渲染包含图片占位符的消息,再沿 <|image_pad|> 拆成:

EncodedTextChunk + ImageChunk + EncodedTextChunk
def build_prompt_chunks(
    tokenizer,
    image_processor,
    image,
    subject,
    choices,
):
    messages = [
        {
            "role": "user",
            "content": [
                {"type": "text", "text": format_question(subject, choices)},
                {"type": "image", "image": "geoqa"},
            ],
        }
    ]
    prompt = tokenizer.apply_chat_template(
        messages,
        tokenize=False,
        add_generation_prompt=True,
        enable_thinking=False,
    )
    before_image, after_image = prompt.split("<|image_pad|>")
    return [
        trio.types.EncodedTextChunk(
            tokens=tokenizer.encode(before_image, add_special_tokens=False)
        ),
        encode_image(image, image_processor),
        trio.types.EncodedTextChunk(
            tokens=tokenizer.encode(after_image, add_special_tokens=False)
        ),
    ]

图片会先将透明区域合成到白色背景并转成 RGB。encode_image() 使用 image processor 计算视觉 patch 数,再把视觉 token 数写入 ImageChunk.expected_tokens

patches = image_processor.get_number_of_image_patches(
    image.height,
    image.width,
    images_kwargs={},
)
expected_tokens = patches // int(image_processor.merge_size) ** 2

image_chunk = trio.ImageChunk(
    data=buffer.getvalue(),
    format=chunk_format,
    expected_tokens=expected_tokens,
)

expected_tokens 参与本地序列长度计算。脚本在 rollout 后继续校验 len(ModelInput) 与远端返回的 response.input_tokens,防止图文 token 计数错位。

2. 异步采样同题 group

一个 step 内的 8 道题共用同一版当前 policy sampler,不同题目通过 tqdm_asyncio.gather() 并发。每道题在一次请求中采样 8 条 completion:

prompt = trio.ModelInput(chunks=prompt_chunks)
prompt_length = len(prompt)

response = await sampling_client.sample_async(
    prompt=prompt,
    num_samples=group_size,
    sampling_params=sampling_params,
    return_text=True,
)

if response.input_tokens != prompt_length:
    raise ValueError(
        f"图文 prompt 长度不一致:"
        f"local={prompt_length}, remote={response.input_tokens}"
    )

reward 只读取回答中最后一个合法的 boxed 选项:

BOXED_CHOICE_PATTERN = re.compile(
    r"\\boxed\s*\{\s*([A-D])\s*\}",
    re.IGNORECASE,
)

predicted_choice = extract_choice(text)
reward = float(predicted_choice == gold_choice)
mean_reward = sum(rewards) / len(rewards)
advantage = reward - mean_reward
回答结果Reward
最后一个 \boxed{A-D} 与标签一致1.0
选项错误、格式非法或没有 boxed 选项0.0

同组 reward 全为 0 或全为 1 时,所有 advantage 都是 0,脚本会跳过整个 group。degenerate_fraction 用来观察当前 batch 中缺少组内相对信号的题目比例。

3. 对齐图片上下文与 completion loss

多模态 GRPO 仍然执行自回归右移。model_input 保存完整图文 prompt 和 completion[:-1],其余三个数组在 prompt / image 区间补零:

字段prompt / image 区间completion 区间
model_input完整图文 chunkscompletion[:-1]
target_tokens0 占位完整 completion tokens
logprobs0.0 占位rollout 时的 old logprobs
advantages0.0 占位该 completion 的 group-relative advantage

对应的 Datum 构造为:

def build_grpo_datum(group, sample):
    model_input = trio.ModelInput(
        chunks=[
            *group.prompt_chunks,
            trio.types.EncodedTextChunk(tokens=sample.tokens[:-1]),
        ]
    )
    observation_length = group.prompt_length - 1
    return trio.Datum(
        model_input=model_input,
        loss_fn_inputs={
            "target_tokens": np.asarray(
                [0] * observation_length + sample.tokens,
                dtype=np.int64,
            ),
            "logprobs": np.asarray(
                [0.0] * observation_length + sample.logprobs,
                dtype=np.float32,
            ),
            "advantages": np.asarray(
                [0.0] * observation_length
                + [sample.advantage] * len(sample.tokens),
                dtype=np.float32,
            ),
        },
    )

completion tokens、采样时的 old logprobs 和 advantages 必须严格等长。old logprobs 来自生成这条轨迹的当前 sampler,参数更新后不再重算替换。

4. 完成 rollout 后更新 LoRA

每个 step 开始时先保存当前 LoRA 权重并创建 sampler:

sampling_client = (
    await training_client.save_weights_and_get_sampling_client_async()
)

全部题目完成 group rollout、reward 和 advantage 计算后,非退化 completion 一起进入策略更新:

forward_backward = await training_client.forward_backward_async(
    datums,
    loss_fn="importance_sampling",
)
optim_step = await training_client.optim_step_async(
    trio.AdamParams(learning_rate=args.learning_rate)
)

result = await forward_backward
await optim_step

这里有两种异步返回边界:

  • sample_async() 一次 await 直接返回采样结果;
  • forward_backward_async()optim_step_async() 第一次 await 返回远端 future,还要继续 await future 才完成任务。

保存 checkpoint 时,脚本同时写出:

*-sampler   # 用于后续采样和评测
*-state     # 完整训练状态,用于继续训练

完整代码

把下面三个文件保存在同一目录中。eval.py 会直接导入 train.py 的图文 prompt 与答案解析函数。

download-dataset.py:下载并固定数据集

"""下载 GeoQA 数据集到当前案例的 datasets 目录。

运行:
    python download-dataset.py
"""

from __future__ import annotations

import argparse
from pathlib import Path
from tempfile import TemporaryDirectory

from datasets import load_dataset
from huggingface_hub import snapshot_download

DATASET_ID = "hz2475/geoQA"
DEFAULT_OUTPUT_DIR = Path(__file__).resolve().parent / "datasets"
TEST_SEED = 42
TEST_SIZE = 100


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="下载 GeoQA 数据集")
    parser.add_argument("--dataset-id", default=DATASET_ID)
    parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR)
    parser.add_argument("--revision", default="main")
    parser.add_argument("--force", action="store_true", help="强制重新下载")
    return parser.parse_args()


def main() -> None:
    args = parse_args()
    output_dir = args.output_dir.expanduser().resolve()
    output_dir.mkdir(parents=True, exist_ok=True)

    with TemporaryDirectory(prefix="geoqa-") as download_dir:
        snapshot_download(
            repo_id=args.dataset_id,
            repo_type="dataset",
            revision=args.revision,
            local_dir=download_dir,
            allow_patterns=["data/*.parquet"],
            force_download=args.force,
        )
        parquet_files = sorted((Path(download_dir) / "data").glob("*.parquet"))
        dataset = load_dataset(
            "parquet",
            data_files={"train": [str(path) for path in parquet_files]},
            split="train",
        )

        train_data = dataset.filter(
            lambda split: split == "train",
            input_columns=["original_split"],
            desc="提取 GeoQA train",
        ).remove_columns("original_split")
        test_data = dataset.filter(
            lambda split: split == "test",
            input_columns=["original_split"],
            desc="提取 GeoQA test",
        ).shuffle(seed=TEST_SEED)
        test_data = test_data.select(
            range(len(test_data) - TEST_SIZE, len(test_data))
        ).remove_columns("original_split")

        train_path = output_dir / "train.parquet"
        test_path = output_dir / "test.parquet"
        train_data.to_parquet(train_path)
        test_data.to_parquet(test_path)

    print(f"dataset_dir={output_dir}")
    print(f"train_file={train_path} rows={len(train_data)}")
    print(f"test_file={test_path} rows={len(test_data)}")


if __name__ == "__main__":
    main()

train.py:运行 Vision GRPO 训练

"""使用 PyTRIO 在 GeoQA 上运行多模态 GRPO。

准备数据:
python download-dataset.py

小规模测试:
python train.py \
    --steps 20 \
    --batch-size 8 \
    --group-size 8 \
    --max-tokens 1024 \
    --save-every 10 \
    --swanlab-mode disabled
"""

from __future__ import annotations

import argparse
import asyncio
import io
import re
import time
from dataclasses import dataclass
from importlib.metadata import version
from pathlib import Path
from typing import Any

import numpy as np
import pytrio as trio
import swanlab
from datasets import Dataset, load_dataset
from PIL import Image
from tqdm.asyncio import tqdm_asyncio
from transformers import AutoImageProcessor

SCRIPT_DIR = Path(__file__).resolve().parent
DEFAULT_DATASET_DIR = SCRIPT_DIR / "datasets"
DEFAULT_MODEL = "Qwen/Qwen3.5-4B"
IMAGE_PAD_TOKEN = "<|image_pad|>"
CHOICE_LETTERS = "ABCD"
BOXED_CHOICE_PATTERN = re.compile(r"\\boxed\s*\{\s*([A-D])\s*\}", re.IGNORECASE)


@dataclass(frozen=True)
class RolloutSample:
    tokens: list[int]
    logprobs: list[float]
    text: str
    predicted_choice: str | None
    reward: float
    advantage: float


@dataclass(frozen=True)
class RolloutGroup:
    prompt_chunks: list[Any]
    prompt_length: int
    samples: list[RolloutSample]


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="GeoQA 多模态 GRPO")
    parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
    parser.add_argument("--base-model", default=DEFAULT_MODEL)
    parser.add_argument("--lora-rank", type=int, default=32)
    parser.add_argument("--steps", type=int, default=10)
    parser.add_argument("--batch-size", type=int, default=4)
    parser.add_argument("--group-size", type=int, default=4)
    parser.add_argument(
        "--max-samples", type=int, default=0, help="0 表示使用全部训练集"
    )
    parser.add_argument("--max-tokens", type=int, default=64)
    parser.add_argument("--temperature", type=float, default=1.0)
    parser.add_argument("--top-p", type=float, default=1.0)
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--learning-rate", type=float, default=4e-5)
    parser.add_argument(
        "--swanlab-mode",
        choices=("online", "local", "offline", "disabled"),
        default="online",
    )
    parser.add_argument("--swanlab-project", default="vision-grpo")
    parser.add_argument(
        "--experiment-name",
        default="vision-grpo-qwen35-4b-geoqa",
    )
    parser.add_argument(
        "--weights-name",
        default="vision-grpo-qwen35-4b-geoqa",
    )
    parser.add_argument(
        "--save-every",
        type=int,
        default=10,
        help="每隔多少个 step 保存一次,0 表示只保存最终 checkpoint",
    )
    parser.add_argument(
        "--save-weights",
        action=argparse.BooleanOptionalAction,
        default=True,
    )
    parser.add_argument("--show-samples", action="store_true")
    return parser.parse_args()


def load_geoqa_train(dataset_dir: Path, seed: int, max_samples: int) -> Dataset:
    """读取并打乱训练数据。"""
    dataset = load_dataset(
        "parquet",
        data_files=str(dataset_dir / "train.parquet"),
        split="train",
    ).shuffle(seed=seed)
    if max_samples > 0:
        dataset = dataset.select(range(min(max_samples, len(dataset))))
    return dataset


def pick_batch(dataset: Dataset, step: int, batch_size: int) -> Dataset:
    """按 step 顺序取 batch,走完数据后从头继续。"""
    start = step * batch_size
    indices = [(start + offset) % len(dataset) for offset in range(batch_size)]
    return dataset.select(indices)


def encode_image(image: Image.Image, image_processor: Any) -> trio.ImageChunk:
    """将图片编码成 PyTRIO chunk,并计算视觉 token 数。"""
    chunk_format = "jpeg" if image.format in {"JPG", "JPEG"} else "png"
    rgba = image.convert("RGBA")
    background = Image.new("RGBA", rgba.size, (255, 255, 255, 255))
    image = Image.alpha_composite(background, rgba).convert("RGB")

    buffer = io.BytesIO()
    image.save(buffer, format=chunk_format.upper())
    patches = image_processor.get_number_of_image_patches(
        image.height,
        image.width,
        images_kwargs={},
    )
    expected_tokens = patches // int(image_processor.merge_size) ** 2
    return trio.ImageChunk(
        data=buffer.getvalue(),
        format=chunk_format,
        expected_tokens=expected_tokens,
    )


def format_question(subject: str, choices: list[str]) -> str:
    """将题目和四个选项整理成模型指令。"""
    choice_lines = "\n".join(
        f"{letter}. {choice}"
        for letter, choice in zip(CHOICE_LETTERS, choices, strict=True)
    )
    return (
        "请根据图片解答下面的几何选择题。\n"
        f"题目:{subject.strip()}\n"
        f"选项:\n{choice_lines}\n"
        "请先进行简单逻辑推理思考,再给出最终答案。"
        "最终选项格式必须是 \\boxed{A}\\boxed{B}\\boxed{C}\\boxed{D}。"
    )


def build_prompt_chunks(
    tokenizer: Any,
    image_processor: Any,
    image: Image.Image,
    subject: str,
    choices: list[str],
) -> list[Any]:
    """先用 chat template 格式化 messages,再拆成图文 chunks。"""
    messages = [
        {
            "role": "user",
            "content": [
                {"type": "text", "text": format_question(subject, choices)},
                {"type": "image", "image": "geoqa"},
            ],
        }
    ]
    prompt = tokenizer.apply_chat_template(
        messages,
        tokenize=False,
        add_generation_prompt=True,
        enable_thinking=False,
    )
    before_image, after_image = prompt.split(IMAGE_PAD_TOKEN)
    return [
        trio.types.EncodedTextChunk(
            tokens=tokenizer.encode(before_image, add_special_tokens=False)
        ),
        encode_image(image, image_processor),
        trio.types.EncodedTextChunk(
            tokens=tokenizer.encode(after_image, add_special_tokens=False)
        ),
    ]


def extract_choice(text: str) -> str | None:
    """提取回答中最后一个 boxed 选项。"""
    matches = BOXED_CHOICE_PATTERN.findall(text)
    return matches[-1].upper() if matches else None


async def run_rollout_group(
    sampling_client: Any,
    tokenizer: Any,
    prompt_chunks: list[Any],
    gold_choice: str,
    sampling_params: trio.SamplingParams,
    group_size: int,
) -> RolloutGroup:
    """异步采样同一道题的一组回答,并计算组内 advantage。"""
    prompt = trio.ModelInput(chunks=prompt_chunks)
    prompt_length = len(prompt)
    response = await sampling_client.sample_async(
        prompt=prompt,
        num_samples=group_size,
        sampling_params=sampling_params,
        return_text=True,
    )
    if response.input_tokens != prompt_length:
        raise ValueError(
            f"图文 prompt 长度不一致:local={prompt_length}, remote={response.input_tokens}"
        )

    raw_samples: list[tuple[list[int], list[float], str, str | None, float]] = []
    rewards: list[float] = []
    for sequence in response.sequences:
        tokens = list(sequence.tokens)
        logprobs = [float(value) for value in sequence.logprobs]
        if len(tokens) != len(logprobs):
            raise ValueError("生成 token 与 logprob 长度不一致")
        text = sequence.text or tokenizer.decode(tokens, skip_special_tokens=True)
        predicted_choice = extract_choice(text)
        reward = float(predicted_choice == gold_choice)
        rewards.append(reward)
        raw_samples.append((tokens, logprobs, text, predicted_choice, reward))

    mean_reward = sum(rewards) / len(rewards)
    samples = [
        RolloutSample(
            tokens=tokens,
            logprobs=logprobs,
            text=text,
            predicted_choice=predicted_choice,
            reward=reward,
            advantage=reward - mean_reward,
        )
        for tokens, logprobs, text, predicted_choice, reward in raw_samples
    ]
    return RolloutGroup(prompt_chunks, prompt_length, samples)


def build_grpo_datum(group: RolloutGroup, sample: RolloutSample) -> trio.Datum:
    """把图文 prompt chunk 和 completion 拼成 GRPO Datum。"""
    model_input = trio.ModelInput(
        chunks=[
            *group.prompt_chunks,
            trio.types.EncodedTextChunk(tokens=sample.tokens[:-1]),
        ]
    )
    observation_length = group.prompt_length - 1
    return trio.Datum(
        model_input=model_input,
        loss_fn_inputs={
            "target_tokens": np.asarray(
                [0] * observation_length + sample.tokens,
                dtype=np.int64,
            ),
            "logprobs": np.asarray(
                [0.0] * observation_length + sample.logprobs,
                dtype=np.float32,
            ),
            "advantages": np.asarray(
                [0.0] * observation_length + [sample.advantage] * len(sample.tokens),
                dtype=np.float32,
            ),
        },
    )


def init_swanlab(args: argparse.Namespace, dataset_size: int) -> Any:
    """初始化训练日志。"""
    return swanlab.init(
        mode=args.swanlab_mode,
        project=args.swanlab_project,
        experiment_name=args.experiment_name,
        config={
            "algorithm": "vision-grpo",
            "dataset": "hz2475/geoQA",
            "dataset_size": dataset_size,
            "base_model": args.base_model,
            "pytrio_version": version("pytrio"),
            "enable_thinking": False,
            "lora_rank": args.lora_rank,
            "steps": args.steps,
            "batch_size": args.batch_size,
            "group_size": args.group_size,
            "max_tokens": args.max_tokens,
            "temperature": args.temperature,
            "learning_rate": args.learning_rate,
            "save_every": args.save_every,
        },
    )


async def save_checkpoint(
    training_client: trio.TrainingClient,
    weights_name: str,
    step: int,
) -> None:
    """同时保存推理权重和完整训练状态。"""
    prefix = f"{weights_name}-step-{step}"
    sampler_future = await training_client.save_weights_for_sampler_async(
        name=f"{prefix}-sampler"
    )
    state_future = await training_client.save_state_async(name=f"{prefix}-state")
    sampler_weights, training_state = await asyncio.gather(
        sampler_future,
        state_future,
    )
    print(f"Sampler 权重:{sampler_weights.path}")
    print(f"State 权重:{training_state.path}")


async def main(args: argparse.Namespace) -> None:
    train_data = load_geoqa_train(
        args.dataset_dir.expanduser().resolve(),
        args.seed,
        args.max_samples,
    )
    print(f"加载 GeoQA train 数据:{len(train_data)} 条")
    print(f"PyTRIO:{version('pytrio')}")

    service_client = trio.ServiceClient()
    training_client = await service_client.create_lora_training_client_async(
        base_model=args.base_model,
        rank=args.lora_rank,
        seed=args.seed,
    )
    tokenizer = training_client.get_tokenizer()
    image_processor = AutoImageProcessor.from_pretrained(
        args.base_model,
        use_fast=False,
    )
    sampling_params = trio.SamplingParams(
        max_tokens=args.max_tokens,
        temperature=args.temperature,
        top_p=args.top_p,
        stop="<|im_end|>",
    )
    adam_params = trio.AdamParams(learning_rate=args.learning_rate)
    swanlab_run = init_swanlab(args, len(train_data))
    last_saved_step = 0

    try:
        for step in range(args.steps):
            batch_rows = list(pick_batch(train_data, step, args.batch_size))
            sampling_client = (
                await training_client.save_weights_and_get_sampling_client_async()
            )
            datums: list[trio.Datum] = []
            all_samples: list[RolloutSample] = []
            prompt_rewards: list[float] = []
            degenerate_groups = 0

            # 一个 step 内的不同题目使用同一版 sampler 并发 rollout。
            rollout_groups = await tqdm_asyncio.gather(
                *(
                    run_rollout_group(
                        sampling_client,
                        tokenizer,
                        build_prompt_chunks(
                            tokenizer,
                            image_processor,
                            row["image"],
                            str(row["subject"]),
                            [str(choice) for choice in row["choices"]],
                        ),
                        CHOICE_LETTERS[int(row["label"])],
                        sampling_params,
                        args.group_size,
                    )
                    for row in batch_rows
                ),
                desc=f"Step {step + 1}/{args.steps} rollout",
                unit="题",
            )

            for row, group in zip(batch_rows, rollout_groups, strict=True):
                gold_choice = CHOICE_LETTERS[int(row["label"])]
                all_samples.extend(group.samples)
                rewards = [sample.reward for sample in group.samples]
                prompt_rewards.append(sum(rewards) / len(rewards))

                if args.show_samples:
                    print(f"\nGeoQA id={row['id']} gold={gold_choice}")
                    for index, sample in enumerate(group.samples):
                        print(
                            f"  sample={index} predicted={sample.predicted_choice} "
                            f"reward={sample.reward:.0f} text={sample.text!r}"
                        )

                # 整组 reward 相同时没有相对优势,不参与更新。
                if len(set(rewards)) == 1:
                    degenerate_groups += 1
                    continue
                datums.extend(
                    build_grpo_datum(group, sample)
                    for sample in group.samples
                    if sample.tokens
                )

            mean_output_tokens = sum(
                len(sample.tokens) for sample in all_samples
            ) / len(all_samples)
            tqdm_asyncio.write(
                f"本 batch 平均输出长度:{mean_output_tokens:.1f} tokens"
            )

            trainer_metrics: dict[str, float] = {}
            if datums:
                forward_backward = await training_client.forward_backward_async(
                    datums,
                    loss_fn="importance_sampling",
                )
                optim_step = await training_client.optim_step_async(adam_params)
                result = await forward_backward
                await optim_step
                trainer_metrics = {
                    key: float(value) for key, value in result.metrics.items()
                }

            mean_reward = sum(prompt_rewards) / len(prompt_rewards)
            format_rate = sum(
                sample.predicted_choice is not None for sample in all_samples
            ) / len(all_samples)
            degenerate_fraction = degenerate_groups / len(prompt_rewards)
            metrics = {
                "reward": mean_reward,
                "format_rate": format_rate,
                "degenerate_fraction": degenerate_fraction,
                "train_datums": len(datums),
                "rollout/completion_tokens_mean": mean_output_tokens,
                **{f"trainer/{key}": value for key, value in trainer_metrics.items()},
            }
            swanlab.log(metrics, step=step)

            loss_mean = trainer_metrics.get("loss_mean")
            loss_text = "n/a" if loss_mean is None else f"{loss_mean:.4f}"
            print(
                f"Step {step + 1}/{args.steps} | reward={mean_reward:.3f} | "
                f"format={format_rate:.1%} | degenerate={degenerate_fraction:.1%} | "
                f"datums={len(datums)} | loss_mean={loss_text}",
                flush=True,
            )

            current_step = step + 1
            if (
                args.save_weights
                and args.save_every > 0
                and current_step % args.save_every == 0
            ):
                await save_checkpoint(training_client, args.weights_name, current_step)
                last_saved_step = current_step

        if args.save_weights and last_saved_step != args.steps:
            await save_checkpoint(training_client, args.weights_name, args.steps)
    finally:
        swanlab_run.finish()


if __name__ == "__main__":
    start_time = time.perf_counter()
    asyncio.run(main(parse_args()))
    print(f"训练耗时:{time.perf_counter() - start_time:.2f}s")

eval.py:异步评测单个模型

"""在固定的 100 条 GeoQA test 样本上评测单个模型。

评测 Base:
python eval.py

评测训练后模型:
python eval.py \
    --model-path trio://run_xxx/sampler_weights/xxx-step-100-sampler
"""

from __future__ import annotations

import argparse
import asyncio
import json
import time
from pathlib import Path
from typing import Any

import pytrio as trio
from datasets import Dataset, load_dataset
from tqdm.asyncio import tqdm_asyncio
from transformers import AutoImageProcessor

from train import (
    CHOICE_LETTERS,
    DEFAULT_DATASET_DIR,
    DEFAULT_MODEL,
    build_prompt_chunks,
    extract_choice,
)

SCRIPT_DIR = Path(__file__).resolve().parent
EVAL_SEED = 42
EVAL_SIZE = 100
DEFAULT_OUTPUT = SCRIPT_DIR / "eval-results.json"


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="评测 GeoQA 多模态 GRPO")
    parser.add_argument("--dataset-dir", type=Path, default=DEFAULT_DATASET_DIR)
    parser.add_argument("--base-model", default=DEFAULT_MODEL)
    parser.add_argument(
        "--model-path",
        help="训练脚本输出的 Sampler 权重路径,不传则评测 Base 模型",
    )
    parser.add_argument("--max-tokens", type=int, default=1024)
    parser.add_argument(
        "--limit",
        type=int,
        default=EVAL_SIZE,
        help="默认评测固定的 100 条;调试时可缩小",
    )
    parser.add_argument("--output", type=Path, default=DEFAULT_OUTPUT)
    return parser.parse_args()


def load_eval_data(dataset_dir: Path, limit: int) -> Dataset:
    """读取下载阶段固定的 100 条测试数据。"""
    dataset = load_dataset(
        "parquet",
        data_files=str(dataset_dir / "test.parquet"),
        split="train",
    )
    return dataset.select(range(min(limit, len(dataset))))


def parse_response(response: Any, tokenizer: Any) -> tuple[str, str | None]:
    """读取单条采样结果中的文本和选项。"""
    sequence = response.sequences[0]
    text = sequence.text or tokenizer.decode(sequence.tokens, skip_special_tokens=True)
    return text, extract_choice(text)


async def main(args: argparse.Namespace) -> None:
    eval_data = load_eval_data(args.dataset_dir.expanduser().resolve(), args.limit)
    service_client = trio.ServiceClient()
    sampling_client = await service_client.create_sampling_client_async(
        base_model=args.base_model,
        model_path=args.model_path,
    )
    tokenizer = sampling_client.get_tokenizer()
    image_processor = AutoImageProcessor.from_pretrained(
        args.base_model,
        use_fast=False,
    )
    sampling_params = trio.SamplingParams(
        max_tokens=args.max_tokens,
        seed=EVAL_SEED,
        temperature=0.0,
        stop="<|im_end|>",
    )

    async def evaluate_row(row: dict[str, Any]) -> dict[str, Any]:
        choices = [str(choice) for choice in row["choices"]]
        gold_choice = CHOICE_LETTERS[int(row["label"])]
        prompt = trio.ModelInput(
            chunks=build_prompt_chunks(
                tokenizer,
                image_processor,
                row["image"],
                str(row["subject"]),
                choices,
            )
        )
        response = await sampling_client.sample_async(
            prompt=prompt,
            num_samples=1,
            sampling_params=sampling_params,
            return_text=True,
        )
        text, predicted_choice = parse_response(response, tokenizer)
        return {
            "id": int(row["id"]),
            "gold": gold_choice,
            "prediction": predicted_choice,
            "text": text,
        }

    # 固定测试集共享同一个 sampler 并发评测。
    results = await tqdm_asyncio.gather(
        *(evaluate_row(row) for row in eval_data),
        desc="评测 GeoQA",
        unit="sample",
    )

    total = len(results)
    correct = sum(result["prediction"] == result["gold"] for result in results)
    formatted = sum(result["prediction"] is not None for result in results)
    metrics = {
        "accuracy": correct / total,
        "format_rate": formatted / total,
    }
    output = args.output.expanduser().resolve()
    output.write_text(
        json.dumps(
            {
                "base_model": args.base_model,
                "model_path": args.model_path,
                "eval_seed": EVAL_SEED,
                "eval_size": total,
                "metrics": metrics,
                "samples": results,
            },
            ensure_ascii=False,
            indent=2,
        ),
        encoding="utf-8",
    )

    print(f"模型:{args.model_path or args.base_model}")
    print(f"Accuracy:{metrics['accuracy']:.1%}")
    print(f"Format rate:{metrics['format_rate']:.1%}")
    print(f"评测结果:{output}")


if __name__ == "__main__":
    start_time = time.perf_counter()
    asyncio.run(main(parse_args()))
    print(f"评测耗时:{time.perf_counter() - start_time:.2f}s")

运行训练

建议先用 20 step 验证数据、图文采样、reward、训练和 checkpoint 全链路:

python train.py \
  --steps 20 \
  --batch-size 8 \
  --group-size 8 \
  --max-tokens 1024 \
  --save-every 10 \
  --swanlab-mode online

复现本文 step-100 配置:

python train.py \
  --steps 100 \
  --batch-size 8 \
  --group-size 8 \
  --max-tokens 1024 \
  --save-every 25 \
  --swanlab-mode online

20-step 小规模会话中的 rewardformat_rate 如下。这两个指标来自每个在线训练 batch,会随题目难度和采样结果波动;固定测试集负责衡量同一批题上的前后变化。

20-step Vision GRPO 的 reward 与 format rate

训练时建议重点观察:

指标作用
reward当前 batch 各题的平均规则奖励
format_rate可解析出 \boxed{A-D} 的 completion 比例
degenerate_fraction没有组内相对信号的题目比例
train_datums实际进入 forward_backward 的 completion 数
rollout/completion_tokens_mean当前 batch 的平均生成长度
trainer/*PyTRIO 服务端返回的训练指标

固定集评测

eval.py 每次只创建一个 sampler。省略 --model-path 时评测 Base Model;传入一个 trio://... 路径时评测对应 LoRA checkpoint。

评测 Base:

python eval.py \
  --output eval-results-base.json

评测 step-100 sampler weights:

python eval.py \
  --model-path 'trio://YOUR_STEP_100_SAMPLER_WEIGHTS' \
  --output eval-results-step-100.json

同一次运行中的 100 道题共享该 sampler,并通过 sample_async() 并发评测:

sampling_client = await service_client.create_sampling_client_async(
    base_model=args.base_model,
    model_path=args.model_path,
)

results = await tqdm_asyncio.gather(
    *(evaluate_row(row) for row in eval_data),
    desc="评测 GeoQA",
    unit="sample",
)

评测参数固定为 temperature=0.0seed=42max_tokens=1024。结果如下:

GeoQA Base Model 与 Vision GRPO step-100 评测对比

模型AccuracyFormat rate采样速度总耗时
Base71.0%75.0%2.20 sample/s48.65s
Vision GRPO step 10087.0%91.0%2.35 sample/s49.45s
提升+16.0 pp+16.0 pp+0.15 sample/s+0.80s

进一步拆分:

  • Base:71 条正确,4 条格式正确但答案错误,25 条没有生成可解析的 boxed 选项;
  • step 100:87 条正确,4 条格式正确但答案错误,9 条没有生成可解析的 boxed 选项;
  • 已格式化回答的准确率从 71 / 75 = 94.7% 上升到 87 / 91 = 95.6%
这篇文档对你有帮助吗?

本页目录