TensorFlow分割模型训练报错:values与permutation尺寸不匹配求助
我使用TensorFlow 2.12.0训练分割模型,训练约6秒后报错停止,报错信息如下:
Epoch 1/100
2023-07-17 08:14:20.618828: E tensorflow/core/grappler/optimizers/meta_optimizer.cc:954] layout failed: INVALID_ARGUMENT: Size of values 0 does not match size of permutation 4 @ fanin shape inmodel_3/dropout_15/dropout/SelectV2-2-TransposeNHWCToNCHW-LayoutOptimizer
8278/8278 [==============================] - 6s 198us/step - loss: 2.1831 - accuracy: 0.8421 - val_loss: 2.2880 - val_accuracy: 0.8349
使用自定义DataGen数据生成器加载COCO2014数据集的图像和掩码,怀疑问题出在数据生成器或模型布局(尤其是Dropout层)。
相关代码片段
# Data generator class DataGen(tf.keras.utils.Sequence): def __init__(self, path_input, path_mask, class_name='person', batch_size=8, image_size=128): self.ids = os.listdir(path_mask) self.path_input = path_input self.path_mask = path_mask self.class_name = class_name self.batch_size = batch_size self.image_size = image_size self.on_epoch_end() def __load__(self, id_name): image_path = os.path.join(self.path_input, id_name) mask_path = os.path.join(self.path_mask, id_name) image = cv2.imread(image_path, 1) # 1 specifies RGB format image = cv2.resize(image, (self.image_size, self.image_size)) # resizing before inserting into the network mask = cv2.imread(mask_path, -1) mask = cv2.resize(mask, (self.image_size, self.image_size)) mask = mask.reshape((self.image_size, self.image_size, 1)) # normalize image image = image / 255.0 mask = mask / 255.0 return image, mask def __getitem__(self, index): id_name = self.ids[index] image, mask = self.__load__(id_name) if image is not None and mask is not None: images = np.expand_dims(image, axis=0) masks = np.expand_dims(mask, axis=0) else: images = np.empty((self.image_size, self.image_size, 3)) masks = np.empty((self.image_size, self.image_size, 1)) return images, masks def on_epoch_end(self): pass def __len__(self): return len(self.ids) # Configure model image_size = 128 epochs = 100 batch_size = 10 # Create data generators train_gen = DataGen(path_input="/kaggle/input/coco-2014-dataset-for-yolov3/coco2014/images/train2014", path_mask="/kaggle/working/mask_train_2014", batch_size=batch_size, image_size=image_size) val_gen = DataGen(path_input="/kaggle/input/coco-2014-dataset-for-yolov3/coco2014/images/val2014", path_mask="/kaggle/working/mask_val_2014", batch_size=batch_size, image_size=image_size) # Define model architecture inputs = Input(shape=(128, 128, 3)) # ... # Compile and train the model optimizer = tf.keras.optimizers.Adam(lr=1e-4) model.compile(optimizer=optimizer, loss='binary_crossentropy', metrics=['accuracy']) model.fit(train_gen, validation_data=val_gen, steps_per_epoch=train_steps, epochs=epochs)
解决思路与建议
修复数据生成器的批量输出逻辑
当前__getitem__仅返回单样本的batch(shape为(1,128,128,3)),但初始化时设置了batch_size=10,维度不匹配直接触发布局优化器报错。正确实现应返回对应batch大小的样本集合:def __getitem__(self, index): # 获取当前批次的ID列表 batch_ids = self.ids[index*self.batch_size : (index+1)*self.batch_size] images = [] masks = [] for id_name in batch_ids: image, mask = self.__load__(id_name) if image is not None and mask is not None: images.append(image) masks.append(mask) # 转换为符合模型输入的numpy数组 return np.array(images), np.array(masks)同时修正
__len__返回批次数量而非样本总数:def __len__(self): return len(self.ids) // self.batch_size添加数据加载有效性校验
__load__未处理cv2.imread加载失败的情况(如文件损坏、不存在),会生成空数组破坏维度一致性,建议添加校验:def __load__(self, id_name): image_path = os.path.join(self.path_input, id_name) mask_path = os.path.join(self.path_mask, id_name) image = cv2.imread(image_path, 1) if image is None: raise ValueError(f"加载图像失败: {image_path}") image = cv2.resize(image, (self.image_size, self.image_size)) mask = cv2.imread(mask_path, -1) if mask is None: raise ValueError(f"加载掩码失败: {mask_path}") mask = cv2.resize(mask, (self.image_size, self.image_size)) mask = mask.reshape((self.image_size, self.image_size, 1)) image = image / 255.0 mask = mask / 255.0 return image, mask临时禁用Layout优化器
若上述修改后仍报错,可尝试跳过NHWC与NCHW的布局转换:import tensorflow as tf tf.config.optimizer.set_experimental_options({"layout_optimizer": False})此为临时 workaround,优先修复数据生成器问题。
校验模型输入输出维度匹配
确保模型输入shape为(128,128,3),输出shape与掩码的(128,128,1)完全匹配,避免模型结构导致的维度不兼容。
内容的提问来源于stack exchange,提问作者Arman Asgharpoor

