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

如何为Keras DataGenerator指定自定义图像加载器?解决类型不匹配问题

解决Keras DataGenerator与OpenCV预处理的类型不匹配问题

错误原因分析

你遇到的错误核心是**cv2.findContours仅支持8位单通道(CV_8UC1)或32位整数单通道(CV_32SC1)图像**,而你的代码中传入的thresh图像格式不符合要求。大概率是因为输入图像被转换成了浮点型(比如float32),导致后续生成的灰度图、二值图也变成了浮点格式,触发了OpenCV的格式校验错误。

另外还有两个小细节问题需要修正:

  • cv2.flip(crop,40)中的flipCode参数无效:OpenCV的flipCode只能是0(垂直翻转)、1(水平翻转)或-1(双向翻转)
  • plt.imshow(val_gen[0])调用错误:flow_from_dataframe返回的迭代器每个元素是(批量图像, 批量标签)的元组,直接传入plt.imshow会报错

解决方案

1. 修复预处理函数的格式问题

在read_and_process_image函数中,先强制将输入图像转换为uint8类型,确保后续OpenCV操作的图像格式完全兼容:

import numpy as np
import cv2
from PIL import Image
import keras

def read_and_process_image(im, im_size):
    # 强制转换为uint8类型,消除格式不兼容风险
    im = im.astype(np.uint8)
    
    gray = cv2.cvtColor(im, cv2.COLOR_RGB2GRAY) # Keras默认加载RGB格式,转换逻辑正确
    _, thresh = cv2.threshold(gray, 10, 255, cv2.THRESH_BINARY)
    # 确保二值图是uint8格式,符合findContours要求
    thresh = thresh.astype(np.uint8)
    
    contours, hierarchy = cv2.findContours(thresh, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
    
    if len(contours) != 0:
        cnt = max(contours, key=cv2.contourArea)
        x, y, w, h = cv2.boundingRect(cnt)
        r = int(w * 0.12)
        crop = im[y+r:y+h-r, x+r:x+w-r]
        # 修正flipCode参数,示例用水平翻转(1),可根据需求调整为0或-1
        crop = cv2.flip(crop, 1)
        crop1 = normalize_histograms(crop)
        return crop1
    else:
        resized_im = cv2.resize(im, (im_size, im_size))
        return normalize_histograms(resized_im)

# 保持你的preprocessing_image函数不变
IM_SIZE = 256
def preprocessing_image(image):
    global IM_SIZE
    image = read_and_process_image(image, IM_SIZE)
    return image

2. 修复DataGenerator与可视化代码

# 初始化DataGenerator
train_datagen = keras.preprocessing.image.ImageDataGenerator(
    rescale=1./255, 
    featurewise_center=True, 
    featurewise_std_normalization=True, 
    preprocessing_function=preprocessing_image
)

val_gen = train_datagen.flow_from_dataframe(
    dataframe=val_data, 
    directory="D:/PROJECTS/MLPC2019/dataset/train/train", 
    x_col="filename", 
    y_col="label", 
    class_mode="categorical", 
    shuffle=False, 
    target_size=(IM_SIZE, IM_SIZE),  # 统一变量名,避免不一致
    batch_size=BATCH_SIZE
)

# 修复可视化代码:取批量中的第一张灰度图
plt.imshow(val_gen[0][0][0], cmap='gray')
plt.show()

3. 自定义图像加载函数(回答你的核心问题)

当然可以指定自定义的图像加载函数!Keras的flow_from_dataframe和flow_from_directory都支持loader参数,你可以直接用OpenCV加载图像,避免格式转换的潜在问题:

# 自定义OpenCV图像加载函数
def opencv_image_loader(path):
    # OpenCV默认加载BGR格式,转成Keras常用的RGB格式
    img = cv2.imread(path)
    img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)
    return img

# 在flow_from_dataframe中指定自定义加载器
val_gen = train_datagen.flow_from_dataframe(
    dataframe=val_data, 
    directory="D:/PROJECTS/MLPC2019/dataset/train/train", 
    x_col="filename", 
    y_col="label", 
    class_mode="categorical", 
    shuffle=False, 
    target_size=(IM_SIZE, IM_SIZE), 
    batch_size=BATCH_SIZE,
    loader=opencv_image_loader  # 指定自定义加载器
)

这样加载的图像直接是OpenCV处理的RGB格式,和你的预处理逻辑完全匹配,能进一步减少格式转换带来的问题。

内容的提问来源于stack exchange,提问作者n_prime

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.06 13:47:37