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

Keras U-Net图像分割中Sequential模型与数据增强报错求助

Oxford Pets图像分割示例数据增强错误修复

错误根源

  1. 核心类型不匹配:数据增强过程中,tf.switch_case的两个分支返回的segmentation_masks dtype不一致(一个为float32,一个为int64),违反了TensorFlow分支返回值必须类型、结构完全一致的要求。
  2. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.11 17:24:52