如何将Numpy数组加载到TensorFlow Dataset实现图像着色分批训练
解决tf.data.Dataset加载NPY文件并拆分LAB分量的问题
核心方案:用tf.py_function包装numpy读取逻辑
因为tf.data.Dataset直接读取文件会得到二进制数据,需要自定义解析函数,借助tf.py_function兼容numpy的np.load操作,完成NPY文件读取、LAB分量拆分,并转换成TensorFlow可用的张量。
完整代码示例
1. 加载预生成的LAB格式NPY文件
import tensorflow as tf import numpy as np # 获取所有NPY文件路径,shuffle=True打乱顺序 file_paths = tf.data.Dataset.list_files("/your/npy/directory/*.npy", shuffle=True) def parse_npy_file(file_path): # 定义numpy侧的读取逻辑 def load_npy(path): # 将TensorFlow路径转字符串,读取NPY文件 lab_image = np.load(path.decode("utf-8")) # 拆分L分量(单通道)和AB分量(双通道) l_channel = lab_image[..., 0:1] # 形状(256,256,1) ab_channels = lab_image[..., 1:] # 形状(256,256,2) # 转float32适配TensorFlow模型 return l_channel.astype(np.float32), ab_channels.astype(np.float32) # 用tf.py_function包装numpy操作,指定输出张量类型 l, ab = tf.py_function(load_npy, [file_path], [tf.float32, tf.float32]) # 手动设置张量形状,避免模型训练时因形状不确定报错 l.set_shape((256, 256, 1)) ab.set_shape((256, 256, 2)) return l, ab # 构建数据集,设置并行处理、批次、预取 batch_size = 32 dataset = file_paths.map(parse_npy_file, num_parallel_calls=tf.data.AUTOTUNE) dataset = dataset.shuffle(buffer_size=1000).batch(batch_size).prefetch(tf.data.AUTOTUNE)
2. 可选:直接读取原始图像并实时转LAB(无需预存NPY)
如果还没将原始图像转成NPY格式,也可以在解析阶段直接完成RGB到LAB的转换,省去预生成NPY的步骤:
def parse_image_file(file_path): def load_and_convert(path): # 读取原始图像并转成RGB数组 img = tf.keras.preprocessing.image.load_img(path.decode("utf-8"), target_size=(256,256)) img_rgb = tf.keras.preprocessing.image.img_to_array(img) # 转LAB颜色空间,同时归一化数值范围 img_lab = tf.image.rgb_to_lab(img_rgb / 255.0) # 拆分分量并归一化到0-1区间(适配模型训练) l = img_lab[..., 0:1] / 100.0 # L分量原始范围0-100 ab = (img_lab[..., 1:] + 128.0) / 255.0 # AB分量原始范围-128到127 return l.astype(np.float32), ab.astype(np.float32) l, ab = tf.py_function(load_and_convert, [file_path], [tf.float32, tf.float32]) l.set_shape((256,256,1)) ab.set_shape((256,256,2)) return l, ab # 用原始图像路径构建数据集 file_paths = tf.data.Dataset.list_files("/your/image/directory/*.jpg", shuffle=True) dataset = file_paths.map(parse_image_file, num_parallel_calls=tf.data.AUTOTUNE).shuffle(1000).batch(32).prefetch(tf.data.AUTOTUNE)
关键细节说明
tf.py_function是核心:它允许在TensorFlow图模式中执行numpy代码,解决了np.load无法直接在tf.data流水线中运行的问题。- 手动设置张量形状:必须显式指定
set_shape,否则数据集输出的张量形状会是未知的,导致模型训练报错。 - 数值归一化:LAB分量的原始范围不适合模型训练,转成0-1区间能提升模型收敛效率,根据你的需求可以调整归一化方式。
内容的提问来源于stack exchange,提问作者Evermus
相关产品推荐
相关产品推荐

