如何在TensorFlow Keras中识别并修复通道数异常的问题图像?
问题
我用以下代码加载本地图像数据集训练模型:
data_load = tk.utils.image_dataset_from_directory( dir, labels="inferred", batch_size=128, image_size=image_shape, shuffle=True, seed=42, validation_split=0.2, subset="training", )
其中dir是数据集本地路径。调用model.fit训练时抛出以下错误:
--------------------------------------------------------------------------- InvalidArgumentError Traceback (most recent call last) c:\Users\HP\Desktop\SBU\Courses\spring23\ese577\Labs\Lab3\lenet.ipynb Cell 8 in 2 1 epoch = 15 ----> 2 hist = model.fit(data.train, batch_size=batch_size, epochs=epoch) File c:\python10\lib\site-packages\keras\utils\traceback_utils.py:70, in filter_traceback..error_handler(*args, **kwargs) 67 filtered_tb = _process_traceback_frames(e.__traceback__) 68 # To get the full stack trace, call: 69 # `tf.debugging.disable_traceback_filtering()` ---> 70 raise e.with_traceback(filtered_tb) from None 71 finally: 72 del filtered_tb File c:\python10\lib\site-packages\tensorflow\python\eager\execute.py:52, in quick_execute(op_name, num_outputs, inputs, attrs, ctx, name) 50 try: 51 ctx.ensure_initialized() ---> 52 tensors = pywrap_tfe.TFE_Py_Execute(ctx._handle, device_name, op_name, 53 inputs, attrs, num_outputs) 54 except core._NotOkStatusException as e: 55 if name is not None: InvalidArgumentError: Graph execution error: Number of channels inherent in the image must be 1, 3 or 4, was 2 [[{{node decode_image/DecodeImage}}]] [[IteratorGetNext]] [Op:__inference_train_function_4955]
错误总是在训练到以下阶段时出现:
Epoch 1/15 9/124 [=>............................] - ETA: 6:36 - loss: 5.0484 - accuracy: 0.4852
我查过这类错误通常出现在读取BMP图像时,但我的图像都是JPG格式,还是遇到了这个问题。请问怎么修复这个错误,或者怎么识别并删除问题图像来继续训练?
解决方案
方法一:强制统一图像通道数
TensorFlow的image_dataset_from_directory默认自动推断图像通道数,但遇到2通道这类异常值时会报错,可通过两种方式处理:
- 指定color_mode参数:
根据模型输入要求,直接指定图像通道模式。如果模型需要3通道RGB图像,添加color_mode="rgb";如果是单通道灰度图,用color_mode="grayscale"。修改后的加载代码:data_load = tk.utils.image_dataset_from_directory( dir, labels="inferred", batch_size=128, image_size=image_shape, shuffle=True, seed=42, validation_split=0.2, subset="training", color_mode="rgb" # 按需选择"grayscale" ) - 自定义预处理函数修复通道:
如果指定color_mode无效,可在数据集加载后添加映射函数,手动将2通道图像转换为3通道:import tensorflow as tf def normalize_channels(image, label): # 处理2通道图像,复制最后一个通道凑成3通道 if tf.shape(image)[-1] == 2: image = tf.concat([image, image[..., -1:]], axis=-1) # 确保最终是3通道(或根据模型需求调整为1通道) return image, label data_load = data_load.map(normalize_channels)
方法二:定位并删除异常图像
如果不想修改代码,直接清理数据集,可遍历所有JPG图像,找出通道数为2的文件并删除:
import os from PIL import Image # 遍历目录找出异常图像 def scan_invalid_images(dataset_dir): invalid_paths = [] for root, _, files in os.walk(dataset_dir): for fname in files: if fname.lower().endswith(('.jpg', '.jpeg')): img_path = os.path.join(root, fname) try: with Image.open(img_path) as img: if len(img.getbands()) == 2: invalid_paths.append(img_path) except Exception as e: print(f"读取失败:{img_path},错误:{e}") invalid_paths.append(img_path) return invalid_paths # 执行扫描 bad_images = scan_invalid_images(dir) print(f"发现 {len(bad_images)} 张异常图像") # 删除异常图像(操作前请备份数据集) for path in bad_images: os.remove(path) print(f"已删除:{path}")
内容的提问来源于stack exchange,提问作者chiru player
相关产品推荐
相关产品推荐

