如何将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
相关产品推荐
相关产品推荐

