使用GeneralizedRCNN执行torch.jit.trace报too many indices for tensor of dimension 3
报错根因
Detectron2的GeneralizedRCNN系列模型原生forward方法要求输入为字典组成的列表,每个字典必须包含image键对应单张图像张量,可选传入height/width指定图像原始尺寸。你传入的是形状为[1, 3, 519, 1038]的批量张量,模型遍历输入时尝试从3维张量中取["image"]键,触发索引错误。
修复方案
方案1:使用符合原生接口的输入格式trace
首先在构建模型前开启TorchScript支持,关闭动态逻辑避免trace异常:
# 新增配置项 cfg.TORCHSCRIPT = True model = build_model(cfg) model.eval()
调整输入格式为模型要求的字典列表,再执行trace:
from PIL import Image from torchvision import transforms input_image = Image.open("model/xxx.jpg") to_tensor = transforms.ToTensor() input_tensor = to_tensor(input_image) # 构造符合要求的输入 batched_inputs = [{"image": input_tensor}] # trace时输入需打包为元组 trace = torch.jit.trace(model, (batched_inputs,))
方案2:封装模型适配批量张量输入
如果需要直接传入[B, 3, H, W]格式的批量张量推理,可以封装一层转发逻辑再trace:
import torch import torch.nn as nn class RCNNWrapper(nn.Module): def __init__(self, raw_model): super().__init__() self.raw_model = raw_model self.raw_model.eval() def forward(self, batch_tensor): # 将批量张量转为模型要求的字典列表 batched_inputs = [{"image": img} for img in batch_tensor] return self.raw_model(batched_inputs) # 封装后trace wrapped_model = RCNNWrapper(model) input_batch = input_tensor.unsqueeze(0) trace = torch.jit.trace(wrapped_model, input_batch)
注意事项
- 若需要模型返回原图尺寸的预测坐标,需要在输入字典中新增
height、width字段,传入图像的原始高宽 - 如果推理时会传入不同尺寸的输入,建议使用
torch.jit.script替代trace,避免trace固定输入尺寸导致的推理错误
内容的提问来源于stack exchange,提问作者Honey Kumar
相关产品推荐
相关产品推荐

