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

TensorFlow数据扩增报错及图像数据集扩容方案咨询

问题解决:TensorFlow数据扩增错误与数据集扩容方案

一、解决ValueError错误

错误根源是重复执行批处理操作:tf.keras.preprocessing.image_dataset_from_directory创建数据集时,已经通过batch_size=32完成了批处理,返回的是4D张量(形状为(32,180,180,3))。但你的prepare函数里再次调用ds.batch(batch_size),导致数据被二次打包,变成5D张量((None, None, 180,180,3)),不符合扩增层的输入要求。

修改后的prepare函数需移除重复的批处理步骤:

def prepare(ds, shuffle=False, augment=False):
  if shuffle:
    ds = ds.shuffle(1000)

  # 移除重复的batch调用——dataset_from_directory已完成批处理
  if augment:
    ds = ds.map(lambda x, y: (data_augmentation(x, training=True), y), 
                num_parallel_calls=AUTOTUNE)

  return ds.prefetch(buffer_size=AUTOTUNE)

修改后扩增层可正确接收4D批量张量,错误会被修复。

二、数据集扩容方案

针对你需要物理扩容数据集(而非训练时动态随机增强)的需求,提供两种可行方案:

方案1:动态扩增+数据集拼接,模拟扩容效果

无需修改磁盘文件,通过拼接原数据集与扩增后的数据集实现扩容。比如要实现2倍扩容:

# 创建未做扩增的原训练集副本
train_ds_original = prepare(train_ds, shuffle=True, augment=False)
# 创建做扩增的训练集
train_ds_augmented = prepare(train_ds, shuffle=True, augment=True)
# 拼接得到2倍规模的训练集
train_ds_2x = train_ds_original.concatenate(train_ds_augmented)

后续用train_ds_2x训练模型即可,更高倍数可重复拼接多次。

方案2:保存扩增图片到磁盘,实现物理扩容

如果需要实际增加数据集文件数量,可遍历原数据集,将扩增后的图片保存到新目录:

import pathlib

# 创建扩增数据集根目录
augmented_dataset_path = "MY_AUGMENTED_DATASET_PATH"
pathlib.Path(augmented_dataset_path).mkdir(parents=True, exist_ok=True)

# 遍历每个类别
for class_name in CLASS_NAMES:
    class_dir = os.path.join(DATASET_PATH, class_name)
    augmented_class_dir = os.path.join(augmented_dataset_path, class_name)
    pathlib.Path(augmented_class_dir).mkdir(parents=True, exist_ok=True)
    
    # 处理类别下的每张图片
    for img_name in os.listdir(class_dir):
        img_path = os.path.join(class_dir, img_name)
        img = tf.keras.preprocessing.image.load_img(img_path, target_size=(img_height, img_width))
        img_array = tf.keras.preprocessing.image.img_to_array(img)
        img_array = tf.expand_dims(img_array, 0)  # 增加batch维度
        
        # 生成2张扩增图片(可调整数量)
        for i in range(2):
            augmented_img_array = data_augmentation(img_array, training=True)[0].numpy()
            augmented_img = tf.keras.preprocessing.image.array_to_img(augmented_img_array)
            augmented_img.save(os.path.join(augmented_class_dir, f"{os.path.splitext(img_name)[0]}_aug_{i}.jpg"))

运行后会得到包含原图片+扩增图片的新数据集,直接用它训练即可。

三、额外优化建议

  • 训练时保持验证集无扩增(当前代码已实现),避免干扰精度评估。
  • 若模型精度仍不理想,可尝试解冻ResNet50顶部几层预训练层进行微调,通常能进一步提升性能。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 19:41:05