You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将Pix2Struct单图推理代码改为批量推理以提升效率?

批量改造Pix2Struct推理代码提升效率

当然可以通过批量推理来提升效率,以下是改造后的完整代码:

from PIL import Image
from transformers import Pix2StructProcessor, Pix2StructForConditionalGeneration
import torch

device = "cuda" if torch.cuda.is_available() else "cpu"

# 加载处理器和模型(仅需初始化一次)
processor = Pix2StructProcessor.from_pretrained("google/deplot", is_vqa=True)
model = Pix2StructForConditionalGeneration.from_pretrained("google/deplot", is_vqa=True).to(device)

# 批量定义待推理图像路径
image_paths = [
    './data/test_imgs/test1.png',
    './data/test_imgs/test2.png',
    './data/test_imgs/test3.png'
    # 可按需添加更多图像路径
]

# 批量加载图像
batch_images = []
for path in image_paths:
    with open(path, "rb") as f:
        img = Image.open(f).convert("RGB")
        batch_images.append(img)

# 批量预处理数据
prompt = "Generate underlying data table of the figure below:"
inputs = processor(
    images=batch_images,
    text=[prompt] * len(batch_images),  # 为每个图像分配相同的prompt
    return_tensors="pt",
    padding=True  # 自动对齐批量内不同尺寸的图像
).to(device)

# 批量生成预测结果
predictions = model.generate(**inputs, max_new_tokens=512)

# 批量解码结果
deplot_results = [processor.decode(pred, skip_special_tokens=True) for pred in predictions]

# 输出所有推理结果
for idx, result in enumerate(deplot_results):
    print(f"图像{idx+1}的推理结果:")
    print(result)
    print("-" * 50)

关键改动说明

  • 批量加载图像:将单张图像加载逻辑改为遍历路径列表,批量收集待推理图像
  • 批量预处理:processor支持接收图像列表与文本列表,通过padding=True自动处理图像尺寸差异,生成符合模型输入要求的批量张量
  • 批量生成与解码:model.generate直接处理批量输入,返回对应数量的预测序列,最后通过列表推导式完成批量解码

额外优化建议

  • 尽量保证待推理图像尺寸统一,减少预处理阶段的padding开销,进一步提升推理速度
  • 若GPU显存不足,可调整批量大小(如每次处理4/8张图像),避免显存溢出
  • 保持model.generate的do_sample=False默认设置,在保证推理确定性的同时提升效率

内容的提问来源于stack exchange,提问作者Robin Lee

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.10 20:40:07