PyTorch训练边界框检测模型交叉熵损失类型不匹配如何解决
问题根因
PyTorch的F.cross_entropy交叉熵损失函数要求传入的第二个参数(类别标签target)必须为Long长整型张量,你当前代码中从数据集加载的类别标签y_class是Int整型张量,类型不匹配触发本次运行时错误。
修复方案
提供两种可选修改方式,任选其一即可:
- 方式一:从数据集源头统一标签类型
修改RoadDataset类的__getitem__方法,将读取到的类别标签转换为numpy.int64格式,后续自动转换为PyTorch张量时就会对应为Long类型:
def __getitem__(self, idx): path = self.paths[idx] print(path) y_class = self.y[idx].astype(np.int64) # 新增类型转换 x, y_bb = transformsXY(path, self.bb[idx], self.transforms) x = normalize(x) x = np.rollaxis(x, 2) return x, y_class, y_bb
- 方式二:在训练/验证阶段直接转换类型
分别找到train_epocs和val_metrics函数中标签移到GPU的代码,直接添加long()转换:
# train_epocs中修改 y_class = y_class.cuda().long() # val_metrics中修改 y_class = y_class.cuda().long()
额外注意事项
你当前定义的边界框预测头输出维度有误:边界框通常为4个坐标值,你现在写的是nn.Linear(512, 26),和类别数一致,会导致后续边界框损失计算维度不匹配,需要修改为nn.Linear(512, 4)。
内容的提问来源于stack exchange,提问作者Nike
相关产品推荐
相关产品推荐

