使用ImageDataGenerator加载肺炎X光数据集遇UnidentifiedImageError求助
解决ImageDataGenerator加载肺炎X光数据集时的UnidentifiedImageError问题
问题场景
使用ImageDataGenerator加载肺炎X光数据集时,调用next()触发UnidentifiedImageError,但手动遍历所有.jpeg文件并用PIL打开均正常。
可能的原因与解决方案
1. 路径不一致导致加载无效文件
代码中使用的相对路径(chest_xray/train)与手动遍历的绝对路径(/content/drive/MyDrive/...)不匹配,可能导致加载到目录外的无效文件或隐藏文件。
解决方法:将代码中的路径替换为绝对路径,确保与数据集实际存储路径一致:
train_data_path = "/content/drive/MyDrive/FlatironProjects/Phase-4/chest_xray/train" test_data_path = "/content/drive/MyDrive/FlatironProjects/Phase-4/chest_xray/test" val_data_path = "/content/drive/MyDrive/FlatironProjects/Phase-4/chest_xray/val"
2. 数据集存在隐藏文件或损坏文件
即使PIL能打开文件,仍可能存在:
- 系统隐藏文件(如Mac的
.DS_Store、Windows的Thumbs.db)被当成图片加载 - 0字节的空文件或完整性受损的文件
解决方法:
- 手动删除数据集目录下的系统隐藏文件
- 运行以下代码检查并清理损坏/空文件:
import PIL from pathlib import Path from PIL import UnidentifiedImageError path = Path("/content/drive/MyDrive/FlatironProjects/Phase-4/chest_xray").rglob("*.jpeg") for img_p in path: # 检查空文件 if img_p.stat().st_size == 0: print(f"删除空文件: {img_p}") img_p.unlink() continue # 验证图片完整性 try: img = PIL.Image.open(img_p) img.verify() # 深度验证图片编码 except (PIL.UnidentifiedImageError, PIL.Image.DecompressionBombError) as e: print(f"删除损坏文件: {img_p}, 错误信息: {e}") img_p.unlink()
3. 过大的batch_size引发内存不足
设置的batch_size等于数据集总样本数(如训练集5215),可能导致内存溢出,引发文件读取异常。
解决方法:减小batch_size,并添加灰度模式适配X光图像的单通道特性:
train_generator = ImageDataGenerator().flow_from_directory( train_data_path, target_size=(64, 64), batch_size=32, # 调整为合理的批量大小 color_mode='grayscale', # 匹配X光图像的单通道格式 class_mode='categorical' )
4. 替换为tf.data加载(更稳定的替代方案)
如果上述方法无效,建议使用TensorFlow原生的tf.data API加载数据集,避免ImageDataGenerator的潜在兼容性问题:
import tensorflow as tf def preprocess_image(file_path): # 提取标签 label = tf.strings.split(file_path, '/')[-2] label = tf.cast(label == 'PNEUMONIA', tf.int32) # 读取并预处理图像 img = tf.io.read_file(file_path) img = tf.image.decode_jpeg(img, channels=1) img = tf.image.resize(img, (64, 64)) img = img / 255.0 # 归一化 return img, label # 构建训练数据集 train_dataset = tf.data.Dataset.list_files(train_data_path + '/*/*.jpeg') train_dataset = train_dataset.map(preprocess_image, num_parallel_calls=tf.data.AUTOTUNE) train_dataset = train_dataset.batch(32).prefetch(tf.data.AUTOTUNE)
内容的提问来源于stack exchange,提问作者Kacey Clougher
相关产品推荐
相关产品推荐

