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

cGAN图像上色训练步骤正确性验证及训练策略咨询

CIFAR10图像上色cGAN的train_step正确性验证与训练策略指导

我针对CIFAR10数据集构建了图像上色的条件生成对抗网络(cGAN),完成了数据预处理、生成器与判别器的模型定义,并编写了训练步骤函数train_step及早停机制、指标初始化代码。现需确认train_step的实现是否正确,同时寻求相关训练策略的指导。


早停、指标初始化及训练步骤函数

# Metrics initialization:
acc = ks.metrics.MeanSquaredError()
acc_ = ks.metrics.MeanSquaredError()
history:dict = {'loss':[],'acc':[], 'val_loss':[],'val_acc':[]}
# Don't bother this line of code:
max_step:int = len(train_data)//32
# Earlystopping implementation
global max_batch_patience,patience
breakpoint:int = 2
max_epochs:int = 10
patience:int = 1
max_batch_patience:list = []

def train_step(X, y, gen, disc, loss_fn, optimizer):
    X_gray, clues = X
    X_tr = y
    with tf.GradientTape() as gen_tape,tf.GradientTape() as disc_tape:
        pred = gen([X_gray,clues], training=True)
        loss = loss_fn(X_tr, pred)
        disc_pred = disc([X_gray], training=True)
        disc_loss = loss_fn(X_tr, disc_pred)
    grads = gen_tape.gradient(loss, gen.trainable_variables)
    gen.optimizer.apply_gradients(zip(grads, gen.trainable_variables))
    disc_grads = disc_tape.gradient(disc_loss, disc.trainable_variables)
    disc.optimizer.apply_gradients(zip(disc_grads, disc.trainable_variables))  

数据集预处理及cGAN模型定义

(X_tr,y_tr),(X_test,y_test) = ks.datasets.cifar10.load_data()

X_tr = X_tr.astype(np.float32)/255
X_test = X_test.astype(np.float32)/255

X_tr = X_tr.reshape(-1,32,32,3)
X_test = X_test.reshape(-1,32,32,3)

X_tr_gray = np.array([cv2.cvtColor(x,cv2.COLOR_RGB2GRAY) for x in X_tr])
X_test_gray = np.array([cv2.cvtColor(x, cv2.COLOR_RGB2GRAY) for x in X_test])

X_tr_gray = X_tr_gray.reshape(-1,32,32,1)
X_test_gray = X_test_gray.reshape(-1,32,32,1)



# Task 2, Build a conditional GAN aka cGAN to colorize the cifar10 dataset!

def cgan_layer(num_ch:int,
               num_filters:int,
               kernel_size:int|tuple,
               strides:int|tuple,
               drop:float,
               kernel_init:str='he_normal',
               bn:bool=False):

    model = ks.models.Sequential()

    model.add(Conv2D(num_filters,kernel_size,strides,padding='same',kernel_initializer=kernel_init,\
                     activation=lambda T:tf.nn.leaky_relu(T,0.3)))
    model.add(Dropout(drop))
    model.add(BatchNormalization())

    model.add(Conv2DTranspose(num_filters,kernel_size,strides,padding='same',kernel_initializer=kernel_init,\
                     activation=lambda T:tf.nn.leaky_relu(T,0.3)))
    model.add(Dropout(drop))
    model.add(BatchNormalization())

    model.add(Conv2D(num_ch,kernel_size,strides,padding='same',kernel_initializer=kernel_init,\
                     activation=lambda T:tf.nn.leaky_relu(T,0.3)))
    return model

lay = cgan_layer(3,64,5,1,0.5)



def cGAN_gen(in_shape:tuple,num_clues:tuple|int,embed_dims:tuple,concat_ax:int=-1,reshape_dim:int=256,num_channels:int=3)->ks.models.Model:
    in_layer = Input(shape=in_shape,dtype=tf.float32)
    clue_dim = Input(shape=num_clues,dtype=tf.float32)

    clue_embd = Embedding(embed_dims[0],embed_dims[1],dtype=tf.float32)(clue_dim)
    clue_embd = Flatten()(clue_embd)


    X = cgan_layer(32,32,5,1,0.5)(in_layer)
    X = cgan_layer(32,32,5,1,0.5)(X)
    X = cgan_layer(64,64,5,1,0.5)(X)
    X = cgan_layer(64,64,5,1,0.5)(X)

    X = Flatten()(X)

    Con = Concatenate(axis=concat_ax)([X,clue_embd])

    Resh = Dense(reshape_dim*in_shape[0]*in_shape[1],activation=lambda T:tf.nn.leaky_relu(T,0.3))(Con)
    Resh = Reshape((in_shape[0],in_shape[1],reshape_dim))(Resh)
    # print(Resh.shape)
    X = Conv2D(num_channels,1,1,padding='same',kernel_initializer='he_normal',\
                     activation=lambda T:tf.nn.tanh(T))(Resh)

    return ks.models.Model([in_layer,clue_dim],X)

def cGAN_disc(in_shape:tuple,num_filters:tuple,kernel_size:int|tuple,process_dims:tuple,\
                  kernel_init:str='he_normal',dr:float=0.5):

    in_layer = Input(shape=in_shape,dtype=tf.float32)

    X = Conv2D(num_filters[0],kernel_size,1,padding='same',kernel_initializer=kernel_init,\
                     activation=lambda T:tf.nn.leaky_relu(T,0.3))(in_layer)
    X = Dropout(dr)(X)
    X = BatchNormalization()(X)
    X = Conv2D(num_filters[1],kernel_size,1,padding='same',kernel_initializer=kernel_init,\
                     activation=lambda T:tf.nn.leaky_relu(T,0.3))(X)
    X = Dropout(dr)(X)
    X = BatchNormalization()(X)
    X = Conv2D(num_filters[2],kernel_size,1,padding='same',kernel_initializer=kernel_init,\
                     activation=lambda T:tf.nn.leaky_relu(T,0.3))(X)
    X = Dropout(dr)(X)
    X = BatchNormalization()(X)

    F = Flatten()(X)

    X = Dense(process_dims[0],activation=lambda T:tf.nn.leaky_relu(T,0.3))(F)
    X = Dropout(dr)(X)
    X = Dense(process_dims[1],activation=lambda T:tf.nn.leaky_relu(T,0.3))(X)

    L = Dense(1,activation=lambda T:tf.nn.sigmoid(T))(X)

    return ks.models.Model(in_layer,L)

gen = cGAN_gen(in_shape=(32,32,1),num_clues=10,embed_dims=(10,10),concat_ax=1,reshape_dim=2,num_channels=3)
disc = cGAN_disc((32,32,3),(64,128,256),5,(256,128))

for x in disc.layers:
    print(x.output_shape)

for x in gen.layers:
    print(x.output_shape)

gen.summary()

disc.summary()

# checking if the models react as i expected 
disc(tf.random.normal((32,32,32,3),dtype=tf.float32))

gen([tf.random.normal((32,32,32,1)),tf.random.uniform(shape=(32,10),minval=0,maxval=10)])
#                            b  h  w  c                            b embd_dim



y_tr = np.array([tf.one_hot(x,10) for x in y_tr])
y_test = np.array([tf.one_hot(x,10) for x in y_test])

# checking the shape of the data
X_tr_gray.shape,y_tr.shape,X_tr.shape

train_data = tf.data.Dataset.from_tensor_slices(((X_tr_gray,y_tr),X_tr)).shuffle(3000).batch(32).repeat(6)

test_data = tf.data.Dataset.from_tensor_slices(((X_test_gray,y_test),X_test)).batch(32)

train_data = train_data.prefetch(tf.data.AUTOTUNE)

test_data = test_data.prefetch(tf.data.AUTOTUNE)

train_step实现问题分析

当前train_step存在多处核心逻辑错误,完全不符合cGAN的训练范式:

  • 判别器输入逻辑错误:判别器的任务是区分真实彩色图和生成的彩色图,且作为cGAN,必须结合条件信息(灰度图+类别线索)。但当前代码中disc([X_gray])仅输入灰度图,既没有传入真实/生成的彩色图,也没有融入类别线索,完全无法完成真假判断的任务。
  • 损失函数使用错误:
    • 生成器仅用MSE和真实图计算损失,缺少对抗损失(即判别器对生成图的判断结果),无法让生成器学会生成能欺骗判别器的真实图像,最终生成图会非常模糊。
    • 判别器用MSE计算真实图和自身输出的损失,但判别器是二分类任务,应该用二元交叉熵损失,分别计算对真实图(标签为1)和生成图(标签为0)的判别损失。
  • 优化器逻辑混乱:函数参数传入的optimizer未被使用,反而直接调用gen.optimizer和disc.optimizer,导致外部配置的优化器无法生效,训练逻辑不可控。
  • 指标更新缺失:初始化的acc、acc_和history字典在train_step中完全没有更新,无法记录训练过程中的损失和指标变化,无法监控训练效果。
  • 训练顺序不合理:GAN训练通常需要先训练判别器多次(2-3次),再训练生成器1次,防止判别器过早收敛导致生成器无法学习。当前代码同步各训一次,极易出现模式崩溃或训练不稳定的问题。

修正后的train_step示例

以下是符合cGAN图像上色任务的train_step实现,包含对抗损失+内容损失的组合、正确的判别器训练逻辑、指标更新:

# 定义二元交叉熵损失(用于对抗损失)
bce_loss = tf.keras.losses.BinaryCrossentropy(from_logits=False)
# 定义MSE损失(用于内容损失)
mse_loss = tf.keras.losses.MeanSquaredError()

def train_step(X, y, gen, disc, gen_optimizer, disc_optimizer, lambda_content=100):
    X_gray, clues = X
    real_color = y
    batch_size = tf.shape(X_gray)[0]
    # 真实标签和伪造标签
    real_labels = tf.ones((batch_size, 1))
    fake_labels = tf.zeros((batch_size, 1))

    # ---------------------
    # 训练判别器
    # ---------------------
    with tf.GradientTape() as disc_tape:
        # 生成伪造彩色图
        fake_color = gen([X_gray, clues], training=True)
        # 判别真实图:cGAN的判别器需要接收条件+图像,需先修改disc模型输入为[灰度图, 类别线索, 彩色图]或拼接后的张量
        disc_real = disc([X_gray, clues, real_color], training=True)
        disc_fake = disc([X_gray, clues, fake_color], training=True)
        # 判别器损失:真实图损失+伪造图损失
        disc_loss_real = bce_loss(real_labels, disc_real)
        disc_loss_fake = bce_loss(fake_labels, disc_fake)
        disc_loss = (disc_loss_real + disc_loss_fake) / 2

    # 更新判别器梯度
    disc_grads = disc_tape.gradient(disc_loss, disc.trainable_variables)
    disc_optimizer.apply_gradients(zip(disc_grads, disc.trainable_variables))

    # ---------------------
    # 训练生成器
    # ---------------------
    with tf.GradientTape() as gen_tape:
        fake_color = gen([X_gray, clues], training=True)
        disc_fake = disc([X_gray, clues, fake_color], training=True)
        # 生成器损失:对抗损失 + lambda*内容损失
        gen_adv_loss = bce_loss(real_labels, disc_fake)
        gen_content_loss = mse_loss(real_color, fake_color)
        gen_loss = gen_adv_loss + lambda_content * gen_content_loss

    # 更新生成器梯度
    gen_grads = gen_tape.gradient(gen_loss, gen.trainable_variables)
    gen_optimizer.apply_gradients(zip(gen_grads, gen.trainable_variables))

    # 更新指标
    acc.update_state(real_color, fake_color)
    history['loss'].append(gen_loss.numpy())
    history['disc_loss'].append(disc_loss.numpy())
    
    return gen_loss, disc_loss

注意:上述代码需要先修改判别器cGAN_disc的输入,使其接受条件信息(灰度图+类别线索)和彩色图作为输入,比如将输入拼接成一个张量,或者设计多输入模型。


训练策略指导

  1. 损失函数组合:采用**对抗损失(二元交叉熵)+ 内容损失(MSE/MAE/感知损失)**的组合,其中内容损失的权重(如lambda_content=100)可以根据生成效果调整,感知损失(用预训练CNN提取特征计算损失)能生成更逼真的图像。
  2. 判别器训练频次:每次训练迭代中,先训练判别器2-3次,再训练生成器1次,避免判别器过于强大导致生成器无法学习。
  3. 优化器与学习率:生成器和判别器使用独立的优化器(如Adam),判别器学习率可设置为1e-4,生成器设置为5e-5;加入学习率衰减(如余弦退火),防止训练后期震荡。
  4. 早停机制优化:监控验证集的PSNR(峰值信噪比)或SSIM(结构相似性)指标,当连续patience轮次指标没有提升时停止训练,同时保存最优模型权重。
  5. 数据增强:对灰度图和彩色图同步进行随机水平翻转、随机平移等增强操作,提升模型泛化能力,避免过拟合。
  6. 梯度裁剪:对生成器和判别器的梯度进行裁剪(如tf.clip_by_norm(grads, 1.0)),防止梯度爆炸导致训练不稳定。
  7. 模型验证:每轮训练后,在验证集上生成图像,计算PSNR、SSIM等指标,同时可视化生成结果,直观判断训练效果。

内容的提问来源于stack exchange,提问作者Araz Alizadeh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 16:27:02