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
相关产品推荐
相关产品推荐

