YOLOv5 Concat模块RuntimeError咨询:张量尺寸不匹配
YOLOv5 Concat模块RuntimeError问题排查与解决思路
问题重现
测试代码:
import torch yolo_ = torch.hub.load('ultralytics/yolov5', 'yolov5x', pretrained=True, force_reload=True) yolo_(torch.rand((2,3,1280,720)))
运行后抛出错误:
RuntimeError: Sizes of tensors must match except in dimension 1. Expected size 46 but got size 45 for tensor number 1 in the list.
错误源自~/.cache/torch/hub/ultralytics_yolov5_master/models/common.py中Concat模块forward方法的torch.cat操作,待拼接的两个张量形状如下:
x[0].shape= torch.Size([2, 640, 80, 46]) x[1].shape= torch.Size([2, 640, 80, 45])
错误原因分析
这个错误的核心是YOLOv5特征金字塔结构中,不同分支的特征图尺寸无法对齐,导致拼接操作失败,具体诱因包括:
- 输入尺寸不满足模型stride要求:YOLOv5核心卷积模块的步长(stride)为32,输入图像的宽高必须是32的整数倍。输入尺寸(1280,720)中,720不是32的整数倍(720÷32=22.5),多次下采样/上采样后,不同分支的特征图尺寸会因取整逻辑出现差异,最终导致拼接维度不匹配。
- 版本兼容性问题:PyTorch版本与YOLOv5版本不兼容时,动态尺寸张量的插值、卷积操作可能出现精度偏差,引发特征图尺寸不一致。
- 模型加载配置偏差:torch.hub加载模型时,默认配置可能未启用自动尺寸适配逻辑,导致非标准尺寸输入无法被正确处理。
解决思路
1. 调整输入尺寸为32的整数倍
将输入宽高修改为32的整数倍,比如(1280,704)(704=32×22)或(1280,736)(736=32×23),修改后测试代码:
import torch yolo_ = torch.hub.load('ultralytics/yolov5', 'yolov5x', pretrained=True, force_reload=True) yolo_(torch.rand((2,3,1280,704)))
2. 启用模型自动尺寸适配
YOLOv5提供autoshape方法,可自动处理非标准尺寸输入,内部会自动调整尺寸并添加padding,确保特征图对齐:
import torch yolo_ = torch.hub.load('ultralytics/yolov5', 'yolov5x', pretrained=True, force_reload=True).autoshape() yolo_(torch.rand((2,3,1280,720)))
3. 匹配PyTorch与YOLOv5版本
确保使用的PyTorch版本与YOLOv5版本兼容:
- YOLOv5 v6.0+建议搭配PyTorch 1.9及以上版本
- 可指定加载特定版本的YOLOv5,避免hub自动拉取的版本存在兼容问题:
yolo_ = torch.hub.load('ultralytics/yolov5', 'yolov5x', pretrained=True, force_reload=True, trust_repo=True)
4. 手动修正Concat模块的尺寸对齐逻辑
若必须使用非32倍数的输入,可修改common.py中的Concat模块,在拼接前统一特征图尺寸:
import torch.nn.functional as F from torch import nn class Concat(nn.Module): # Concatenate a list of tensors along dimension def __init__(self, dimension=1): super().__init__() self.d = dimension def forward(self, x): # 新增尺寸对齐逻辑,统一所有张量与第一个张量的尺寸一致 if len(x) > 1: target_h, target_w = x[0].shape[2], x[0].shape[3] for i in range(1, len(x)): if x[i].shape[2:] != (target_h, target_w): x[i] = F.interpolate(x[i], size=(target_h, target_w), mode='bilinear', align_corners=False) return torch.cat(x, self.d)
内容的提问来源于stack exchange,提问作者Mihai.Mehe
相关产品推荐
相关产品推荐

