使用SRGAN实现图像超分辨率时出现Graph execution error的排查求助
图像超分辨率训练Graph execution error问题分析
问题背景
尝试基于SRGAN实现图像超分辨率,采用DIV2K数据集作为数据源,此前有CNN图像分类任务经验,但本次数据集无类别标签。已将数据集解压至D:\Unzipped\DIV2K_train_HR,使用以下代码读取数据:
img_dataset = tensorflow.keras.utils.image_dataset_from_directory("D:\\unzipped")
构建了如下不符合SRGAN架构的模型:
model = Sequential() model.add(Conv2D(filters=64,kernel_size=(3,3),activation="relu",input_shape=(256,256,3))) model.add(AveragePooling2D(pool_size=(2,2))) model.add(Conv2D(filters=64,kernel_size=(3,3),activation="relu")) model.add(MaxPooling2D(pool_size=(2,2))) model.compile(optimizer='sgd', loss='mse') model.fit(img_dataset,batch_size=32, epochs=10)
训练时触发Graph execution error,需确认错误根源是数据集无类别标签还是图像尺寸不统一。
错误根源分析
两个因素都可能触发报错,且当前代码完全不符合超分辨率任务的核心逻辑:
- 数据集读取逻辑错误:
image_dataset_from_directory是为分类任务设计的API,它会将目标目录下的子文件夹识别为类别,返回(图像张量, 类别标签)格式的数据集。你将根目录设为D:\\unzipped,API会把DIV2K_train_HR当作单个类别,输出的标签维度与模型输出的下采样特征图维度完全不匹配,直接导致MSE损失计算时的维度冲突,触发Graph execution error。 - 图像尺寸不统一:
DIV2K数据集的原始图像尺寸并不一致(多为2048×1024或其他尺寸),而image_dataset_from_directory默认不会自动统一图像尺寸,你的模型输入固定为(256,256,3),尺寸不匹配的图像会直接触发张量形状错误,这也是报错的关键原因。 - 模型架构完全偏离任务需求:
超分辨率是输入低分辨率图像,输出对应高分辨率图像的回归任务,而你当前的模型通过多次池化不断下采样图像,与任务目标完全相反,即便解决数据问题,也无法完成超分辨率任务。
修正方案
1. 修正数据读取逻辑
放弃分类任务的image_dataset_from_directory,手动构建超分辨率所需的(低分辨率图, 高分辨率图)训练对,同时统一图像尺寸:
import os import cv2 import numpy as np # 配置路径和尺寸 hr_dir = "D:\\Unzipped\\DIV2K_train_HR" # 设定高分辨率图尺寸,低分辨率图为其1/2(对应2倍超分) hr_size = (256, 256) lr_size = (128, 128) x_train = [] # 低分辨率图像 y_train = [] # 对应高分辨率图像 for img_name in os.listdir(hr_dir): img_path = os.path.join(hr_dir, img_name) # 读取并统一高分辨率图像尺寸 hr_img = cv2.cvtColor(cv2.imread(img_path), cv2.COLOR_BGR2RGB) hr_img = cv2.resize(hr_img, hr_size, interpolation=cv2.INTER_LANCZOS4) # 通过双三次插值生成低分辨率图像 lr_img = cv2.resize(hr_img, lr_size, interpolation=cv2.INTER_CUBIC) # 归一化到[0,1]区间 hr_img = hr_img / 255.0 lr_img = lr_img / 255.0 # 加入训练集 x_train.append(lr_img) y_train.append(hr_img) # 转换为numpy数组 x_train = np.array(x_train) y_train = np.array(y_train)
2. 重构符合SRGAN的模型
SRGAN核心是生成器+判别器的对抗架构:
- 生成器:以低分辨率图像为输入,通过残差块提取特征,再通过PixelShuffle或转置卷积完成上采样,输出高分辨率图像。
- 判别器:区分生成的超分辨率图像与真实高分辨率图像,辅助生成器优化。
示例生成器简化结构:
from tensorflow.keras import Sequential, layers from tensorflow.keras.layers import Conv2D, PReLU, BatchNormalization, Add, PixelShuffle def build_generator(): model = Sequential() # 输入层:低分辨率图像(128,128,3) model.add(Conv2D(64, kernel_size=9, padding='same', input_shape=(128,128,3))) model.add(PReLU()) # 残差块(示例设为4个,SRGAN原论文用16个) for _ in range(4): residual = Sequential() residual.add(Conv2D(64, kernel_size=3, padding='same')) residual.add(BatchNormalization()) residual.add(PReLU()) residual.add(Conv2D(64, kernel_size=3, padding='same')) residual.add(BatchNormalization()) model.add(Add()([model.output, residual.output])) # 上采样层(2倍超分) model.add(Conv2D(256, kernel_size=3, padding='same')) model.add(PixelShuffle(2)) model.add(PReLU()) # 输出层:高分辨率图像(256,256,3) model.add(Conv2D(3, kernel_size=9, padding='same', activation='tanh')) return model
内容的提问来源于stack exchange,提问作者Mr Learner
相关产品推荐
相关产品推荐

