PyTorch实现YOLOv4 CAM时3D张量转5D张量reshape报错解决
YOLOv4 CAM特征图reshape适配方案
报错根因
PyTorch的reshape操作强制要求输入输出张量的总元素个数完全相等,你当前的报错来自两个维度设置错误:
- 硬编码网格尺寸为13*13,和当前输入特征图对应的步长不匹配
- 目标形状最后一维保10,和输入特征图最后一维的25(单锚框预测属性数)不匹配
数值校验:
输入总元素数:1*1083*25 = 27075
你写的目标形状总元素数:1*3*13*13*10 = 5070,二者数值差5倍以上,必然触发形状不匹配错误。
动态适配计算逻辑
YOLOv4单检测头的输出维度有3个固定值,不需要硬编码网格尺寸,可以通过输入张量自动计算:
- 固定维度1:batch size,和输入第一维保持一致
- 固定维度2:每个网格对应的锚框数量,YOLOv4所有检测头该值固定为3
- 固定维度3:单个锚框的预测属性数量,和输入最后一维保持一致(你当前输入是25,对应4个边框坐标+1个目标置信度+20个类别得分)
- 动态维度:网格边长
grid_size,计算公式为:grid_size = sqrt(输入总元素数 / (batch_size * 单网格锚框数 * 单锚框属性数))
可直接运行的实现代码
import torch # 输入特征图 input_map = torch.randn(1, 1083, 25) # 固定参数 B = input_map.shape[0] ANCHOR_PER_GRID = 3 ATTR_NUM = input_map.shape[-1] # 动态计算网格尺寸 total_elements = input_map.numel() grid_total = total_elements // (B * ANCHOR_PER_GRID * ATTR_NUM) grid_size = int(grid_total ** 0.5) # 校验输入合法性,避免非正方形特征图报错 assert grid_size * grid_size == grid_total, "输入特征图维度不符合YOLOv4检测头格式,请检查取特征层的位置" # 执行reshape output_map = input_map.reshape(B, ANCHOR_PER_GRID, grid_size, grid_size, ATTR_NUM) print(output_map.shape) # 针对你给出的测试输入,输出为 torch.Size([1, 3, 19, 19, 25]),无报错
补充说明
如果你确实需要得到1313尺寸的网格输出,说明你当前取特征图的位置不对:1313网格对应检测头的输入第二维长度应该是3*13*13=507,即输入形状为torch.Size([1, 507, 25])时,动态计算得到的grid_size才是13,转换后形状为torch.Size([1,3,13,13,25])。不要随意修改最后一维的长度,该值必须和输入特征图最后一维完全相等。
内容的提问来源于stack exchange,提问作者arunmenon
相关产品推荐
相关产品推荐

