使用TensorFlow和Keras构建CNN猫狗分类器时的图执行错误求助
问题诊断
报错核心原因是新增的图像数据中存在损坏、格式无效或不兼容的文件,PIL库无法识别这类文件,导致数据加载流程中断,进而触发Graph execution error。由于仅在图像数量提升至1500张时出现问题,说明初始500张数据正常,问题出在新增的1000张样本中。
解决方案
1. 批量检测并清理损坏图像
遍历数据集目录,用PIL尝试打开每个图像文件,删除无法识别的损坏文件:
import os from PIL import Image data_dir = "/path/to/your/dataset" # 替换为你的数据集根目录 for root, dirs, files in os.walk(data_dir): for file in files: if file.endswith(('.jpg', '.jpeg', '.png')): file_path = os.path.join(root, file) try: with Image.open(file_path) as img: img.verify() # 验证图像完整性 except (IOError, SyntaxError) as e: print(f"删除损坏文件: {file_path}") os.remove(file_path)
2. 增强数据加载的容错性
修改ImageDataGenerator配置,添加参数跳过无法加载的图像,避免训练中断:
from tensorflow.keras.preprocessing.image import ImageDataGenerator # 初始化生成器时添加skip_broken_images参数 train_datagen = ImageDataGenerator(rescale=1./255) train_batches = train_datagen.flow_from_directory( '/path/to/train', target_size=(224,224), batch_size=32, class_mode='categorical', skip_broken_images=True # 关键参数:跳过损坏图像 ) valid_datagen = ImageDataGenerator(rescale=1./255) valid_batches = valid_datagen.flow_from_directory( '/path/to/valid', target_size=(224,224), batch_size=32, class_mode='categorical', skip_broken_images=True )
3. 验证数据集完整性
清理完成后,先测试数据加载是否正常,再执行训练:
# 测试加载一个批次,确认数据无问题 x, y = next(train_batches) print(f"加载批次成功,形状: {x.shape}, {y.shape}") # 执行训练 model.fit(x=train_batches, validation_data=valid_batches, epochs=10, verbose=2)
4. 可选:改用tf.data.Dataset加载数据(更稳定的错误处理)
如果上述方法仍有问题,可使用TensorFlow原生数据集加载工具,自定义错误过滤逻辑:
import tensorflow as tf def load_and_preprocess_image(file_path): try: img = tf.io.read_file(file_path) img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, (224, 224)) img = img / 255.0 # 假设目录结构为 类别目录/图像文件 label = tf.strings.split(file_path, os.sep)[-2] label = tf.cast(label == 'cat', tf.float32) # 根据你的实际类别调整 return img, label except: return None, None # 构建训练数据集 train_files = tf.data.Dataset.list_files("/path/to/train/*/*.jpg") train_dataset = train_files.map(load_and_preprocess_image) train_dataset = train_dataset.filter(lambda x, y: x is not None) # 过滤损坏数据 train_dataset = train_dataset.batch(32).prefetch(tf.data.AUTOTUNE) # 按同样逻辑构建验证集后执行训练 model.fit(train_dataset, validation_data=valid_dataset, epochs=10, verbose=2)
额外提示
- 针对训练精度不理想的问题,可在
ImageDataGenerator中添加数据增强(旋转、翻转、缩放等),提升模型泛化能力; - 也可尝试使用VGG16、ResNet等预训练模型进行迁移学习,快速提升分类精度。
内容的提问来源于stack exchange,提问作者Resai
相关产品推荐
相关产品推荐

