YOLOv4自定义损失训练时维度不匹配报错该如何解决?
YOLOv4自定义损失函数维度不匹配报错解决方案
错误根因
两次报错的核心是模型输出维度、标签维度、损失函数预期维度三者不匹配,具体分析如下:
- 第一次reshape报错:损失函数预期76*76特征层的输出最后一维为
3*(NB_CLASS + 5),按报错参数计算为108,对应总元素数1247616,但你的模型实际输出该层最后一维为255(默认YOLOv4适配COCO 80类的配置:3*(80+5)=255),总元素数2945760,和自定义类别数的预期不匹配。 - 第二次维度不匹配报错:将reshape最后一维改为-1只是绕开了reshape的校验错误,本质上模型输出的255通道和标签生成的108通道仍不匹配,计算损失时触发维度对齐错误。
解决步骤
- 修改模型检测头输出维度
找到YOLOv4对应7676、3838、19*19三个特征层的检测头最后一层卷积层,将卷积核数量从默认的255修改为3 * (NB_CLASS + 5),和你自定义的类别数对齐。 - 优化损失函数动态维度适配
损失函数中不要硬编码固定的BATCH_SIZE,避免最后一个batch样本数不足触发维度错误,修改对应代码如下:
def custom_loss(y_true, y_pred): grid_h, grid_w = y_pred.shape[1:3] # 动态获取batch维度,替换硬编码BATCH_SIZE batch_size = tf.shape(y_pred)[0] if grid_h == grids[0][0]: anchor = anchors[0] elif grid_h == grids[1][0]: anchor = anchors[1] else: anchor = anchors[2] mask_shape = tf.shape(y_true)[:4] cell_x = tf.cast((tf.reshape(tf.tile(tf.range(grid_w), [grid_h]), (1, grid_h, grid_w, 1, 1))),dtype=tf.float32) cell_y = tf.transpose(cell_x, (0,2,1,3,4)) # tile的batch维度用动态值 cell_grid = tf.tile(tf.concat([cell_x,cell_y], -1), [batch_size, 1, 1, NB_BOX, 1]) ###### prediction # 动态batch维度reshape y_pred = tf.reshape(y_pred, (batch_size, grid_h, grid_w, NB_BOX, NB_CLASS+5)) # ......其余逻辑保持不变...... ###### true # 动态batch维度reshape y_true = tf.reshape(y_true, (batch_size, grid_h, grid_w, NB_BOX, NB_CLASS+5)) # ......其余逻辑保持不变......
- 确认标签生成逻辑一致性
检查训练标签y_true的生成代码,确保每个检测头输出的y_true最后一维同样为3*(NB_CLASS +5),和模型输出、损失函数预期完全对齐。
内容的提问来源于stack exchange,提问作者Cindy Zheng
相关产品推荐
相关产品推荐

