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

如何在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通道这类异常值时会报错,可通过两种方式处理:

  1. 指定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"
    )
    
  2. 自定义预处理函数修复通道:
    如果指定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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 22:57:52