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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 14:54:24