使用tf.data.Dataset.from_generator迭代图像生成器时报错
报错核心原因
- 数据类型不匹配:
cv2.imread读取的图像默认是uint8类型(像素值范围0~255),但你在output_signature中声明的张量数据类型为tf.float32,类型校验不通过会直接抛出错误。 - OpenCV接口调用错误:你将颜色空间转换码
cv2.COLOR_BGR2RGB错误作为cv2.imread的读取参数传入。实际上cv2.imread默认读取的图像为BGR通道顺序,转RGB需要单独调用cv2.cvtColor实现,当前写法未正确修正通道顺序,还会导致图像读取逻辑异常。 - 维度处理冗余:你在生成器内部手动为单张图像新增了batch维度(
np.newaxis相关代码),后续如果调用tf.data自带的batch()方法做批量拼接,会额外新增一维,最终输入网络的张量形状不符合卷积层的输入要求。 - 缺少必要预处理:直接输入0~255范围的像素值会导致模型训练难以收敛,需要提前做归一化处理。
修复后可运行代码
修正后的生成器函数
import cv2 import numpy as np import tensorflow as tf from numpy import asarray def get_train_image(dataframe): for i in range(len(dataframe)): # 读取低分辨率图像,修正BGR转RGB逻辑 lr_img = cv2.imread(dataframe.input[i]) lr_img = cv2.cvtColor(lr_img, cv2.COLOR_BGR2RGB) # cv2.resize参数顺序为(宽, 高),输出尺寸对应高120、宽80,单张图形状为(120,80,3) lr_img = cv2.resize(lr_img, (80, 120)) # 直接转float32类型,同时将像素值归一化到0~1区间 lr_img = asarray(lr_img, dtype=np.float32) / 255.0 # 读取对应高分辨率标签图像 hr_img = cv2.imread(dataframe.output[i]) hr_img = cv2.cvtColor(hr_img, cv2.COLOR_BGR2RGB) # 输出尺寸对应高160、宽120,单张图形状为(160,120,3) hr_img = cv2.resize(hr_img, (120, 160)) hr_img = asarray(hr_img, dtype=np.float32) / 255.0 # 不再手动新增batch维度,直接返回单张图像的(H,W,C)格式数组 yield lr_img, hr_img
修正后的tf.data封装逻辑
train_data = tf.data.Dataset.from_generator( get_train_image, output_signature=( # 形状对应单张低清图:高120、宽80、3通道 tf.TensorSpec(shape=(120, 80, 3), dtype=tf.float32), # 形状对应单张高清图:高160、宽120、3通道 tf.TensorSpec(shape=(160, 120, 3), dtype=tf.float32) ), args=[train_df] ) # 配置批量大小、预加载,提升训练效率 BATCH_SIZE = 8 train_data = train_data.batch(BATCH_SIZE).prefetch(tf.data.AUTOTUNE)
额外优化建议
- 数据增强(如随机翻转、随机裁剪)需要对低清、高清图像做同步变换,可直接在生成器的循环内实现,不要在生成器内部手动拼接batch。
prefetch(tf.data.AUTOTUNE)可以让数据加载和模型训练并行执行,减少GPU等待数据的空转时间,训练速度通常能提升30%以上。- 如果你的图像路径包含中文字符,
cv2.imread可能读取失败,可替换为tf.io.read_file+tf.image.decode_image的组合做图像解码,避免路径编码问题。
内容的提问来源于stack exchange,提问作者Shubham luharuka
相关产品推荐
相关产品推荐

