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

PyTorch训练CPM姿态模型遇通道不匹配RuntimeError求助

解决CPM模型训练中通道不匹配的RuntimeError

问题本质

你的输入张量([64,57,28,28])通道数(57)与Mconv1_stage2卷积层定义的输入通道数(55)不匹配,导致报错。核心要从数据拼接逻辑和模型层参数设置两个方向排查。

排查与修复步骤

1. 确认输入通道数的组成

在模型的forward方法中,找到进入Mconv1_stage2前的张量,打印其形状及拼接的各部分张量形状,明确57通道的来源:

# 在Mconv1_stage2前添加打印语句
print("Pre-Mconv1_stage2 input shape:", x.shape)
# 打印拼接的各分量形状
print("Image feature shape:", image_feat.shape)
print("Stage1 heatmap shape:", stage1_heatmap.shape)
print("Center map shape:", center_map.shape)

核对输出结果:

  • 若热图通道数不是21(freihand数据集对应21个关键点),说明数据预处理时热图生成错误,修正热图通道数;
  • 若存在重复拼接的张量(比如中心图被加入两次),去掉重复项;
  • 若各分量通道数相加确实为57,说明模型层参数设置错误。

2. 修正模型层的输入通道数

找到Mconv1_stage2的定义代码,将in_channels参数从55修改为57:

# 原错误代码
self.Mconv1_stage2 = nn.Conv2d(in_channels=55, out_channels=128, kernel_size=3, padding=1)

# 修改后代码
self.Mconv1_stage2 = nn.Conv2d(in_channels=57, out_channels=128, kernel_size=3, padding=1)

3. 对齐CPM阶段输入规范

CPM的stage2及后续阶段输入通常是初始图像特征 + 上一阶段热图 + 中心图,确保拼接逻辑符合这个规范:

# 示例正确拼接逻辑(根据实际特征通道数调整)
stage2_input = torch.cat([initial_image_feat, stage1_output_heatmap, center_map], dim=1)

计算拼接后的通道数,确保和对应stage的第一个卷积层in_channels参数完全一致。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 10:25:18