TensorFlow加载图像提示未知格式错误,请求技术解决方案
问题
在Colab中运行修改自YouTube视频的TensorFlow代码时遇到错误,代码如下:
import os os.environ["TF_CPP_MIN_LOG_LEVEL"] = "2" import tensorflow as tf import pathlib from google.colab import drive drive.mount('/content/gdrive') from tensorflow import keras from tensorflow.keras import layers from tensorflow.keras.preprocessing.image import ImageDataGenerator img_height = 28 img_width = 28 batch_size = 2 proj_path = "/content/gdrive/MyDrive/G10-SIGHTHOUNDS" proj_path = pathlib.Path(proj_path) model = keras.Sequential( [ layers.Input((28, 28, 3)), layers.Conv2D(16, 3, padding="same"), layers.Conv2D(32, 3, padding="same"), layers.MaxPooling2D(), layers.Flatten(), layers.Dense(11), ] ) # METHOD 1 # ==================================================== # # Using dataset_from_directory # # ==================================================== # ds_train = tf.keras.preprocessing.image_dataset_from_directory( proj_path, labels="inferred", label_mode="int", # categorical, binary # class_names=['0', '1', '2', '3', ...] color_mode="rgb", batch_size=batch_size, image_size=(img_height, img_width), # reshape if not in this size shuffle=True, seed=42, validation_split=0.2, subset="training", ) ds_validation = tf.keras.preprocessing.image_dataset_from_directory( proj_path, labels="inferred", label_mode="int", # categorical, binary color_mode="rgb", batch_size=batch_size, image_size=(img_height, img_width), # reshape if not in this size shuffle=True, seed=42, validation_split=0.2, subset="validation", ) def augment(x, y): image = tf.image.random_brightness(x, max_delta=0.05) return image, y ds_train = ds_train.map(augment) # Custom Loops for epochs in range(10): for x, y in ds_train: # train here pass model.compile( optimizer=keras.optimizers.Adam(), loss=tf.losses.SparseCategoricalCrossentropy(from_logits=True), metrics=["accuracy"], ) model.fit(ds_train, epochs=10, verbose=2)
运行后报错:
Drive already mounted at /content/gdrive; to attempt to forcibly remount, call drive.mount("/content/gdrive", force_remount=True). Found 2749 files belonging to 11 classes. Using 2200 files for training. Found 2749 files belonging to 11 classes. Using 549 files for validation. --------------------------------------------------------------------------- InvalidArgumentError Traceback (most recent call last) <ipython-input-15-b39fbc2d162d> in <module> 69 # Custom Loops 70 for epochs in range(10): ---> 71 for x, y in ds_train: 72 # train here 73 pass 3 frames /usr/local/lib/python3.8/dist-packages/tensorflow/python/framework/ops.py in raise_from_not_ok_status(e, name) 7213 def raise_from_not_ok_status(e, name): 7214 e.message += (" name: " + name if name is not None else "") -> 7215 raise core._status_to_exception(e) from None # pylint: disable=protected-access 7216 7217 InvalidArgumentError: {{function_node __wrapped__IteratorGetNext_output_types_2_device_/job:localhost/replica:0/task:0/device:CPU:0}} Unknown image file format. One of JPEG, PNG, GIF, BMP required. [[{{node decode_image/DecodeImage}}]] [Op:IteratorGetNext]
已用Python代码检查数据集文件,确认无损坏文件,寻求解决方法。
解决方案
针对格式不兼容的错误,可按以下步骤排查修复:
清理数据集内的非图像文件:
image_dataset_from_directory会读取目录下所有文件,包括隐藏文件(如.DS_Store、Thumbs.db)、缓存文件或格式不匹配的文件(如WebP、TIFF)。手动遍历数据集目录删除这类文件,或添加过滤逻辑只保留指定后缀的文件:def filter_files(file_path): allowed_extensions = [".jpg", ".jpeg", ".png", ".gif", ".bmp"] return tf.strings.regex_full_match(tf.strings.lower(file_path), r".*(" + "|".join(allowed_extensions) + r")$") # 获取所有文件路径和标签 file_paths = list(proj_path.glob("*/*")) labels = [p.parent.name for p in file_paths] # 转换为TensorFlow数据集并过滤 ds = tf.data.Dataset.from_tensor_slices(([str(p) for p in file_paths], labels)) ds = ds.filter(lambda x, y: filter_files(x))显式指定图像解码格式:自定义加载函数替换默认解码逻辑,避免自动检测失败:
def load_image(file_path, label): img = tf.io.read_file(file_path) # 优先解码JPEG,若包含PNG可添加格式判断分支 img = tf.image.decode_jpeg(img, channels=3) img = tf.image.resize(img, (img_height, img_width)) return img, label # 改用自定义加载逻辑构建数据集 train_ds = tf.data.Dataset.from_tensor_slices((train_files, train_labels)) train_ds = train_ds.map(load_image, num_parallel_calls=tf.data.AUTOTUNE) train_ds = train_ds.batch(batch_size).shuffle(1000)验证文件扩展名与实际格式匹配:部分文件可能后缀错误(如后缀为
.jpg实际是PNG格式),用PIL批量验证并修正:from PIL import Image import os for root, dirs, files in os.walk(proj_path): for file in files: file_path = os.path.join(root, file) try: with Image.open(file_path) as img: img.verify() # 检查扩展名与实际格式是否一致 actual_format = Image.open(file_path).format.lower() file_ext = os.path.splitext(file)[1].lower() if file_ext != f".{actual_format}": print(f"格式不匹配:{file_path},实际格式{actual_format}") # 可选:重命名文件修正后缀 # os.rename(file_path, os.path.join(root, f"{os.path.splitext(file)[0]}.{actual_format}")) except Exception as e: print(f"无效文件:{file_path},错误:{e}")跳过无法解码的文件:添加错误处理逻辑,跳过无效样本避免训练中断:
def load_image_safe(file_path, label): try: img = tf.io.read_file(file_path) img = tf.image.decode_image(img, channels=3, expand_animations=False) img = tf.image.resize(img, (img_height, img_width)) return img, label except Exception as e: # 返回占位符图像和无效标签,后续过滤 return tf.zeros((img_height, img_width, 3)), -1 # 加载后过滤无效样本 ds_train = ds_train.map(load_image_safe).filter(lambda x, y: y != -1)
内容的提问来源于stack exchange,提问作者Jazaic Divina
相关产品推荐
相关产品推荐

