CNN灰度图像上色:12万张图像训练内存错误解决方案问询
问题描述
我正基于R. Zhang的研究实现CNN灰度图像上色模型,当前采用5000张64×64像素图像(对应维度为(5000×64×64)的numpy数组)时可正常运行。现计划将训练图像扩展至128000张,但对颜色标签执行soft encoding操作时出现内存错误,特此询问是否可通过修改代码(如使用tf占位符等方式)解决该内存问题。
解决方案
你遇到的情况非常典型——当数据集规模暴涨时,一次性把所有编码结果都存在内存里,必然会撑爆机器内存。咱们一步步调整代码来解决这个问题:
问题根源分析
你当前的soft_encode方法是把所有图像的编码结果Z存在列表里,最后转成numpy数组。拿128000张图来算,每个Z的维度是(16,16,313)(64/4=16,pts_in_hull.npy里是313个ab量化点),总内存需求超过40GB,这显然超出了普通机器的内存上限。
具体修改方案
1. 改用TensorFlow数据集做批量处理
不要一次性加载128000张图到内存,用tf.data.Dataset分批加载、处理,每次只保留一个batch的数据,内存占用会大幅降低。
2. 将Soft Encoding迁移到TensorFlow图中
把原本用Scikit-learn和NumPy实现的KNN、高斯核计算改成TensorFlow原生操作,既可以利用TF的内存优化,还能配合数据集实现流水线式处理。
3. 重写后的核心代码示例
from __future__ import print_function import pickle import numpy as np import matplotlib.pyplot as plt from keras.models import Sequential from keras.layers import Conv2D, UpSampling2D, BatchNormalization, Dropout from skimage import color import cv2 as cv from keras.regularizers import l2 import tensorflow as tf from keras import backend as K class network: def __init__(self, points_n, epochs, batch_size): self.interpolation_factor = 4 self.batch_size = batch_size # 把量化点转为TF张量 self.Q = tf.convert_to_tensor(np.load('pts_in_hull.npy').astype(np.float32)) self.sigma = 5.0 # 加载数据集为tf.data.Dataset self.dataset = self.load_dataset() def unpickle(self, file): with open(file, 'rb') as fo: dict = pickle.load(fo) return dict def load_dataset(self): d = self.unpickle('train_data_batch_1') x = d['data'] area_img = int((x.shape[1]) / 3) h = w = int(np.sqrt(area_img)) x = np.dstack((x[:, :area_img], x[:, area_img:2*area_img], x[:, 2*area_img:])) # 归一化到0-1,转为float32 x = x.reshape(x.shape[0], h, w, 3).astype(np.float32) / 255.0 # 转为TF数据集,分批并预取优化 dataset = tf.data.Dataset.from_tensor_slices(x) dataset = dataset.batch(self.batch_size).prefetch(tf.data.AUTOTUNE) return dataset # TF图内高斯核计算 @tf.function def gaussian_kernel(self, distances): return tf.exp(-0.5 * (tf.square(distances) / tf.square(self.sigma))) # 批量处理Soft Encoding @tf.function def soft_encode_batch(self, batch_imgs): # 批量转换RGB到Lab lab_imgs = tf.image.rgb_to_lab(batch_imgs) # 提取L通道并归一化 L = (lab_imgs[..., 0:1] - 50.0) / 50.0 # 缩小ab通道到目标尺寸 new_h = tf.shape(L)[1] // self.interpolation_factor new_w = tf.shape(L)[2] // self.interpolation_factor resized_lab = tf.image.resize(lab_imgs, (new_h, new_w), method=tf.image.ResizeMethod.AREA) ab_channels = resized_lab[..., 1:3] # TF原生计算KNN距离 ab_flat = tf.reshape(ab_channels, (-1, 2)) q_expanded = tf.expand_dims(self.Q, 0) ab_expanded = tf.expand_dims(ab_flat, 1) distances = tf.norm(ab_expanded - q_expanded, axis=-1) # 取最近的5个邻居(用负距离取topK等价于取最小距离) _, indices = tf.math.top_k(-distances, k=5) distances = tf.gather(distances, indices, batch_dims=1) # 高斯核加权并归一化 knn_weights = self.gaussian_kernel(distances) knn_weights = knn_weights / tf.reduce_sum(knn_weights, axis=1, keepdims=True) # 构建加权后的soft编码 z_flat = tf.scatter_nd( indices=tf.stack([tf.range(tf.shape(ab_flat)[0])[:, tf.newaxis], indices], axis=-1), updates=tf.reshape(knn_weights, (-1)), shape=(tf.shape(ab_flat)[0], tf.shape(self.Q)[0]) ) # 恢复batch维度 Z = tf.reshape(z_flat, (tf.shape(batch_imgs)[0], new_h, new_w, tf.shape(self.Q)[0])) return Z, L def process_data_for_training(self): # 遍历数据集批量处理,直接喂给模型训练,无需存储所有结果 for batch_idx, batch_imgs in enumerate(self.dataset): Z_batch, L_batch = self.soft_encode_batch(batch_imgs) # 这里插入你的模型训练逻辑,比如model.train_on_batch([L_batch], Z_batch) print(f"Processed batch {batch_idx}: Z shape {Z_batch.shape}, L shape {L_batch.shape}") # 初始化并运行 network1 = network(500, 1, 50) network1.process_data_for_training()
关键优化点说明
- 批量处理:用
tf.data.Dataset分批加载图像,内存占用从几十GB降到几百MB,完全适配普通机器。 - TF图内计算:所有操作都在TensorFlow图中完成,避免了NumPy与TF张量的来回转换,
@tf.function还会自动优化计算流程。 - 无全局存储:不再把所有
Z结果存在内存里,处理一个batch就用一个batch,彻底解决内存溢出问题。
如果确实需要保存所有编码结果,可以把每个batch的结果写入磁盘(比如用TFRecord格式),训练时再从磁盘读取,同样不会占用大量内存。
内容的提问来源于stack exchange,提问作者Konrad S
相关产品推荐
相关产品推荐

