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

