Keras训练U-Net时sequential层输入形状不兼容报错求助
问题成因
报错核心是传入模型的图像张量存在冗余维度:模型预期输入为4维格式(批量大小, 256, 256, 3),实际传入的是5维格式(批量大小, 1, 256, 256, 3),问题全部出在数据预处理的维度操作逻辑上,具体有四点:
- 单张图像处理完成后,你额外执行了
np.expand_dims(image_f, axis=0)给单张图增加批量维度,再将带批量维度的单图存入列表。单张测试时列表只有1个元素,转numpy数组后形状刚好是(1,256,256,3)符合要求;但加载全量数据时,每张图都自带一个长度为1的维度,所有样本堆叠后就会多出一个冗余维度,变成5维张量。 - 掩码处理存在完全相同的问题:每个单独掩码都提前加了axis=0的维度,堆叠后存入列表转数组时,掩码张量也会多出冗余维度,即使修好输入形状,后续也会触发模型输出和标签形状不匹配的错误。
- 你在循环内部每次append完样本就执行一次列表转numpy数组的操作,这步完全冗余,会随数据量上升大幅拖慢加载速度,1000张样本时加载效率会明显降低。
- 额外隐藏问题:SimpleITK读取mhd文件返回的数组默认维度顺序是
(Z轴, 高, 宽),你处理原图时取了image1[0,:,:]提取2D切片,但处理掩码时没有做相同切片操作,直接resize 3D数组也会引发维度异常。
修复方案
按以下逻辑调整预处理代码即可:
- 删除单张图像、单个掩码处理时多余的
np.expand_dims操作,单样本处理完直接保留(高, 宽, 通道数)的3维形状即可,所有样本存入列表后一次性转numpy数组时,会自动在最前面生成批量维度,得到符合模型要求的4维输入/标签张量。 - 给掩码读取步骤增加和原图一致的2D切片提取操作,避免3D数组resize带来的维度错误。
- 将numpy数组转换操作移到循环外部,所有样本加载完成后再一次性转换,提升加载效率。
修正后的完整预处理代码如下:
dataset_dir='/content/drive/MyDrive/training' image_ids = [] mascaras=[] imagens=[] for r, d, f in os.walk(dataset_dir): for file in f: if ('ED.mhd' in file) or ('ES.mhd' in file): image_ids.append(os.path.join(r, file)) image_path=os.path.join(r, file) # 处理原始图像 image1 = sitk.GetArrayFromImage(sitk.ReadImage(image_path,sitk.sitkFloat32)) image1 = image1[0,:,:] # 提取2D灰度切片 image2 = cv2.resize(image1, (256,256)) image3 = image2 / 255.0 # 灰度图复制为3通道,不额外增加批量维度 image_f = np.stack((image3, image3, image3), axis=2) imagens.append(image_f) # 处理4个分割掩码 mask0_path = os.path.join(r,file)[:-4] + "_gt0.mhd" mask0_1 = sitk.GetArrayFromImage(sitk.ReadImage(mask0_path, sitk.sitkFloat32)) mask0_2 = cv2.resize(mask0_1[0,:,:], (32,32)) # 和原图一致提取2D切片 mask1_path = os.path.join(r,file)[:-4] + "_gt1.mhd" mask1_1 = sitk.GetArrayFromImage(sitk.ReadImage(mask1_path, sitk.sitkFloat32)) mask1_2 = cv2.resize(mask1_1[0,:,:], (32,32)) mask2_path = os.path.join(r,file)[:-4] + "_gt2.mhd" mask2_1 = sitk.GetArrayFromImage(sitk.ReadImage(mask2_path, sitk.sitkFloat32)) mask2_2 = cv2.resize(mask2_1[0,:,:], (32,32)) mask3_path = os.path.join(r,file)[:-4] + "_gt3.mhd" mask3_1 = sitk.GetArrayFromImage(sitk.ReadImage(mask3_path, sitk.sitkFloat32)) mask3_2 = cv2.resize(mask3_1[0,:,:], (32,32)) # 拼接4通道掩码,不额外增加维度 mask_f = np.stack((mask0_2, mask1_2, mask2_2, mask3_2), axis=2) mascaras.append(mask_f) # 所有样本加载完成后一次性转为numpy数组 imagens = np.array(imagens, dtype=np.float32) mascaras = np.array(mascaras, dtype=np.float32) # 打印形状验证 print(f"图像数据集形状:{imagens.shape}") print(f"掩码数据集形状:{mascaras.shape}")
执行后正常输出应为:
图像数据集形状:(样本总数, 256, 256, 3) 掩码数据集形状:(样本总数, 32, 32, 4)
此时再调用model.fit()就不会再触发形状不匹配报错。
内容的提问来源于stack exchange,提问作者sprax
相关产品推荐
相关产品推荐

