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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.04 06:45:22