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

机器学习新手使用Retina-Unet时model.fit遇Graph execution error求助

问题描述

我是机器学习新手,正在尝试使用Retina-Unet模型,执行model.fit时遇到错误。

运行代码

model.fit(patches_imgs_train, patches_masks_train, epochs=20, batch_size=32, verbose=1, shuffle=True, validation_split=0.1, callbacks=[checkpointer])

数据维度

  • patches_imgs_train: (190000,1,48,48)
  • patches_masks_train: (190000,2304,2)

核心错误提示

Conv2DCustomBackpropInputOp only supports NHWC.
解决步骤

1. 转换输入图像的通道顺序为NHWC

TensorFlow默认采用NHWC(批量数-高度-宽度-通道数)格式,你的输入是NCHW(批量数-通道数-高度-宽度),需调整维度:

import numpy as np
# 把通道维度从第二位移到最后一位
patches_imgs_train = np.transpose(patches_imgs_train, (0, 2, 3, 1))
# 调整后尺寸应为:(190000, 48, 48, 1)

2. 重构掩码数据维度匹配模型输出

当前掩码是扁平化的48×48=2304像素,需要还原为空间维度,和模型输出对齐:

# 将掩码从(样本数, 扁平化像素数, 类别数)转为(样本数, 高度, 宽度, 类别数)
patches_masks_train = patches_masks_train.reshape((190000, 48, 48, 2))

3. 确认模型输入输出格式

检查Retina-Unet的输入层是否接受(48,48,1)的张量,输出层是否对应(48,48,2)的分割结果。如果模型原本是基于NCHW编写的,可以在所有Conv2D层添加data_format='channels_first'参数,但更推荐统一使用TensorFlow默认的NHWC格式,避免后续兼容问题。

4. 验证调整后的数据维度

修改后打印尺寸确认:

print(patches_imgs_train.shape)  # 预期输出:(190000, 48, 48, 1)
print(patches_masks_train.shape) # 预期输出:(190000, 48, 48, 2)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 16:20:27