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

