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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 20:48:20