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

BLIP-2批量生成图片标题报错,寻求原生批量处理方案

BLIP-2批量生成图片标题的问题与解决方案

问题描述

使用BLIP-2生成单张图片标题的代码可正常运行:

prompt = "this is a picture of"
inputs = processor(trainData[0]["image"], text=prompt, return_tensors="pt").to(device, torch.float16)
generated_ids = model.generate(inputs.pixel_values, input_ids=inputs.input_ids, max_new_tokens=20)

但将输入改为批量图片(如trainData[0:3])时,抛出如下错误:

RuntimeError: shape mismatch: value tensor of shape [245760] cannot be broadcast to indexing result of shape [81920]

无效尝试(Gemini提供的方案)

尝试了以下方案但运行仍报错:

prompt = "this is a picture of"

inputs = processor(trainData[0:3]["image"], text=prompt, return_tensors="pt").to(device, torch.float16)

# 获取模型编码器的图像嵌入
image_embeds = shadow_model.get_base_model().vision_model(pixel_values=inputs.pixel_values).last_hidden_state

# 替换<image> token的嵌入为图像嵌入
image_token_index = processor.tokenizer.convert_tokens_to_ids("<image>")
inputs_embeds = shadow_model.get_base_model().embeddings(inputs.input_ids)
inputs_embeds[inputs.input_ids == image_token_index] = image_embeds

# 使用修改后的inputs_embeds生成文本
generated_ids = model.generate(inputs_embeds=inputs_embeds, max_new_tokens=20)

# 解码并打印生成文本
generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)
print(generated_text)

原生解决方案

批量报错的核心原因是:处理批量输入时,需保证文本prompt与图片批量维度匹配,同时BLIP-2的generate方法支持直接传入批量输入,无需手动修改嵌入。以下是两种原生处理方式:

方法1:直接批量处理(小批量场景)

from transformers import Blip2Processor, Blip2ForConditionalGeneration
import torch

device = "cuda" if torch.cuda.is_available() else "cpu"
processor = Blip2Processor.from_pretrained("Salesforce/blip2-opt-2.7b")
model = Blip2ForConditionalGeneration.from_pretrained(
    "Salesforce/blip2-opt-2.7b",
    torch_dtype=torch.float16
).to(device)

# 提取批量图片
batch_images = [item["image"] for item in trainData[0:3]]
# 生成与图片数量匹配的prompt列表
prompt = "this is a picture of"
batch_prompts = [prompt] * len(batch_images)

# 批量预处理(启用padding对齐文本输入)
inputs = processor(
    images=batch_images,
    text=batch_prompts,
    return_tensors="pt",
    padding=True
).to(device, torch.float16)

# 批量生成标题
generated_ids = model.generate(
    pixel_values=inputs.pixel_values,
    input_ids=inputs.input_ids,
    attention_mask=inputs.attention_mask,
    max_new_tokens=20
)

# 批量解码结果
generated_texts = processor.batch_decode(generated_ids, skip_special_tokens=True)
for text in generated_texts:
    print(text)

关键注意点

  • 必须保证prompt数量与图片数量一致,即使所有prompt相同,也要转为列表形式
  • 开启padding=True,让processor自动对齐文本输入维度
  • 传入attention_mask,避免padding token干扰生成过程

方法2:DataLoader批量处理(大规模数据场景)

from torch.utils.data import Dataset, DataLoader

# 自定义数据集类
class ImageCaptionDataset(Dataset):
    def __init__(self, data, prompt):
        self.data = data
        self.prompt = prompt
    
    def __len__(self):
        return len(self.data)
    
    def __getitem__(self, idx):
        return self.data[idx]["image"], self.prompt

# 自定义数据拼接函数
def collate_fn(batch):
    images, prompts = zip(*batch)
    inputs = processor(
        images=list(images),
        text=list(prompts),
        return_tensors="pt",
        padding=True
    )
    return inputs

# 初始化数据加载器
dataset = ImageCaptionDataset(trainData, "this is a picture of")
dataloader = DataLoader(dataset, batch_size=3, collate_fn=collate_fn)

# 批量推理
model.eval()
with torch.no_grad():
    for batch_inputs in dataloader:
        batch_inputs = batch_inputs.to(device, torch.float16)
        generated_ids = model.generate(
            pixel_values=batch_inputs.pixel_values,
            input_ids=batch_inputs.input_ids,
            attention_mask=batch_inputs.attention_mask,
            max_new_tokens=20
        )
        generated_texts = processor.batch_decode(generated_ids, skip_special_tokens=True)
        print(generated_texts)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 07:05:21