使用tf.keras加载蘑菇分类数据集时遇JPEG解压错误求助
解决JPEG解压错误的方法
这个错误的核心原因是数据集中存在损坏的JPEG文件、文件格式与扩展名不匹配,或者多进程读取时出现文件访问冲突,以下是针对性的解决方法:
1. 检测并移除损坏的图片文件
先遍历数据集,验证所有JPEG文件的完整性,删除损坏的文件:
import os from PIL import Image DATADIR = "/content/drive/MyDrive/Mushrooms/dataset" # 遍历所有子目录下的图片 for root, _, files in os.walk(DATADIR): for file in files: if file.lower().endswith(('.jpg', '.jpeg')): file_path = os.path.join(root, file) try: # 验证文件完整性并尝试加载 with Image.open(file_path) as img: img.verify() img.load() except (IOError, SyntaxError): print(f"删除损坏文件:{file_path}") os.remove(file_path)
2. 让TensorFlow自动跳过损坏文件
如果不想手动检测,直接在数据集管道中添加错误忽略逻辑,修改你的数据集创建代码:
training_ds = tf.keras.utils.image_dataset_from_directory( data_dir, validation_split = 0.2, subset = "training", seed = 555, image_size = (img_height, img_width), batch_size = batch_size ).apply(tf.data.experimental.ignore_errors()) validation_ds = tf.keras.utils.image_dataset_from_directory( data_dir, validation_split = 0.2, subset = "validation", seed = 555, image_size = (img_height, img_width), batch_size = batch_size ).apply(tf.data.experimental.ignore_errors()) # 后续的cache和prefetch逻辑保持不变 AUTOTUNE = tf.data.AUTOTUNE train_ds = training_ds.cache().prefetch(buffer_size = AUTOTUNE) validation_ds = validation_ds.cache().prefetch(buffer_size = AUTOTUNE)
tf.data.experimental.ignore_errors()会自动跳过加载失败的样本,避免训练中断。
3. 处理格式与扩展名不匹配的文件
部分文件可能实际是PNG格式但后缀标为JPG,导致TensorFlow解码失败。改用自定义加载函数,让框架自动检测图片格式:
import tensorflow as tf import pathlib import os DATADIR = "/content/drive/MyDrive/Mushrooms/dataset" data_dir = pathlib.Path(DATADIR).with_suffix('') batch_size = 32 img_height = 200 img_width = 200 # 获取所有文件路径与类别名称 all_files = [str(path) for path in data_dir.glob('*/*')] class_names = sorted([item.name for item in data_dir.glob('*/')]) num_classes = len(class_names) label_map = {name: idx for idx, name in enumerate(class_names)} def load_and_preprocess_image(file_path): # 读取文件 img = tf.io.read_file(file_path) # 自动检测格式,解码为RGB并跳过动图 img = tf.image.decode_image(img, channels=3, expand_animations=False) # 调整尺寸 img = tf.image.resize(img, (img_height, img_width)) # 获取标签 label = tf.strings.split(file_path, os.path.sep)[-2] label = tf.convert_to_tensor(label_map[label.numpy().decode('utf-8')], dtype=tf.int32) return img, label # 构建并分割数据集 dataset = tf.data.Dataset.from_tensor_slices(all_files) dataset = dataset.map( lambda x: tf.py_function(load_and_preprocess_image, [x], [tf.float32, tf.int32]), num_parallel_calls=tf.data.AUTOTUNE ) dataset = dataset.shuffle(len(all_files), seed=555) train_size = int(0.8 * len(all_files)) training_ds = dataset.take(train_size).batch(batch_size).prefetch(tf.data.AUTOTUNE) validation_ds = dataset.skip(train_size).batch(batch_size).prefetch(tf.data.AUTOTUNE)
4. 临时禁用多进程训练
你当前的model.fit启用了多进程模式,可能引发文件访问冲突,先尝试禁用:
model.fit(training_ds, validation_data=validation_ds, epochs=3, use_multiprocessing=False)
如果问题解决,再重新启用多进程(需确保数据集路径为绝对路径,且所有进程可访问)。
内容的提问来源于stack exchange,提问作者Jacob Morgan
相关产品推荐
相关产品推荐

