Keras U-Net图像分割中Sequential模型与数据增强报错求助
Oxford Pets图像分割示例数据增强错误修复
错误根源
- 核心类型不匹配:数据增强过程中,
tf.switch_case的两个分支返回的segmentation_masksdtype不一致(一个为float32,一个为int64),违反了TensorFlow分支返回值必须类型、结构完全一致的要求。 - Sequential模型不兼容多输入:原代码用Sequential模型处理字典格式的多输入(图像+掩码),触发官方警告,建议改用Functional API。
修复方案及代码示例
步骤1:用Functional API重构数据增强Pipeline
替代Sequential模型,完美支持多输入格式,同时统一掩码数据类型:
import tensorflow as tf from tensorflow import keras from tensorflow.keras import layers import keras_cv # 定义输入张量,明确指定dtype input_images = layers.Input(shape=(160, 160, 3), dtype=tf.float32) input_masks = layers.Input(shape=(160, 160, 1), dtype=tf.int32) # 初始化RandAugment层,指定掩码格式 rand_augment = keras_cv.layers.RandAugment( value_range=(0, 255), augmentations_per_image=3, magnitude=0.5, segmentation_mask_format="channels_last" ) # 执行增强,返回增强后的图像和掩码 augmented_output = rand_augment( {'images': input_images, 'segmentation_masks': input_masks}, training=True ) augmented_images = augmented_output['images'] # 强制掩码保持int32类型,消除类型漂移 augmented_masks = tf.cast(augmented_output['segmentation_masks'], tf.int32) # 构建Functional API模型 augment_model = keras.Model( inputs={'images': input_images, 'segmentation_masks': input_masks}, outputs={'images': augmented_images, 'segmentation_masks': augmented_masks} )
步骤2:定义增强函数并重构数据集
# 用构建好的增强模型定义augment_fn def augment_fn(sample): return augment_model(sample, training=True) # 重新构建增强训练集 BATCH_SIZE = 32 AUTOTUNE = tf.data.AUTOTUNE augmented_train_ds = ( train_ds.shuffle(BATCH_SIZE * 2) .map(augment_fn, num_parallel_calls=AUTOTUNE) .batch(BATCH_SIZE) .map(unpackage_inputs) .prefetch(buffer_size=AUTOTUNE) )
可选:数据加载阶段统一掩码类型
如果希望从源头避免类型问题,可在数据加载时就将掩码转为统一类型:
def load_image_mask_pair(image_path, mask_path): # 加载图像(原有逻辑) image = tf.io.read_file(image_path) image = tf.image.decode_jpeg(image, channels=3) image = tf.image.resize(image, (160, 160)) # 加载掩码并强制转为int32 mask = tf.io.read_file(mask_path) mask = tf.image.decode_png(mask, channels=1) mask = tf.image.resize(mask, (160, 160), method=tf.image.ResizeMethod.NEAREST_NEIGHBOR) mask = tf.cast(mask, tf.int32) return {'images': image, 'segmentation_masks': mask}
修复说明
- 类型统一:通过
tf.cast强制掩码保持int32类型,确保tf.switch_case所有分支返回值类型一致。 - 多输入支持:Functional API天然支持字典格式的多输入,解决了Sequential模型的兼容性警告。
- 同步增强:RandAugment层通过
segmentation_mask_format参数确保图像和掩码同步应用增强操作,避免数据错位。
内容的提问来源于stack exchange,提问作者Ashley
相关产品推荐
相关产品推荐

