如何训练含预训练冻结网络的串联TensorFlow模型,是否需自定义训练循环
串联预训练冻结模型训练异常问题
我在训练两个串联模型时遇到问题:第一个模型的输出作为第二个模型的输入,其中第二个模型是已单独训练完成的预训练冻结网络。
我已附上训练流程和模型的架构示意图,其中模拟器(simulator)模型的预训练环节运行正常。但训练生成器(generator)时,我已正确拼接两个模型,训练结果却远不符合预期。请问我是否需要定义自定义训练循环?如果需要,应当如何定义正确的训练循环?
拼接模型实现代码
input1 = keras.Input(shape=(100,), name='noise') input2 = keras.Input(shape=(14,), name='contrast_vector') [image_output, period_output] = generator([input1, input2]) Spectrum_output = simulator([image_output, period_output]) Final_model = keras.Model(inputs=[input1, input2], outputs=[image_output, period_output, Spectrum_output], name='Final_Model') def ssim_loss(y_true, y_pred): return tf.reduce_mean(tf.image.ssim(y_true, y_pred, 1.0)) loss1 = ssim_loss losses = [loss1, 'mse', 'mse'] Final_model.compile( loss= losses, loss_weights=[0.05, 0.01, 1.0], optimizer = keras.optimizers.Adam(learning_rate=1e-3, beta_1=0.5)) history = Final_model.fit([noise_train, CT_vector_train], [y_train, period_train, full_spec_train], batch_size=256, epochs=1000, validation_split=0.2)
模型相关示意图
- 模型架构图:

- TensorFlow中Final_model的可视化图:

问题解答
你当前的训练逻辑存在两个核心问题,可先排查修正,效果仍不符合预期再切换自定义训练循环即可:
- 损失函数优化方向错误:你定义的
ssim_loss直接返回SSIM均值,SSIM取值范围为[-1,1],数值越接近1代表图像相似度越高,作为损失函数需要最小化它的相反数,你当前的写法会让优化器反向拉低SSIM值,图像生成效果自然异常。修正代码如下:
def ssim_loss(y_true, y_pred): return 1 - tf.reduce_mean(tf.image.ssim(y_true, y_pred, 1.0))
- 需要显式冻结simulator权重:构建Final_model前必须显式设置simulator所有层不可训练,否则即使是预训练完成的模型,联合训练时权重也会被更新,破坏预训练效果。拼接模型前添加以下代码:
simulator.trainable = False for layer in simulator.layers: layer.trainable = False
如果完成以上两点修正后训练效果仍不达标,可使用如下自定义训练循环,更方便监控各模块损失变化:
# 定义优化器和损失函数 optimizer = keras.optimizers.Adam(learning_rate=1e-3, beta_1=0.5) ssim_loss_fn = lambda y_true, y_pred: 1 - tf.reduce_mean(tf.image.ssim(y_true, y_pred, 1.0)) mse_loss_fn = keras.losses.MeanSquaredError() # 封装单步训练逻辑 @tf.function def train_step(inputs, labels): noise, ct_vector = inputs y_img, y_period, y_spectrum = labels with tf.GradientTape() as tape: # 前向传播,generator启用训练模式 image_output, period_output = generator([noise, ct_vector], training=True) # simulator全程冻结,关闭训练模式避免BN、Dropout层异常 spectrum_output = simulator([image_output, period_output], training=False) # 计算各分项损失 loss_ssim = ssim_loss_fn(y_img, image_output) loss_period = mse_loss_fn(y_period, period_output) loss_spectrum = mse_loss_fn(y_spectrum, spectrum_output) # 加权计算总损失 total_loss = 0.05 * loss_ssim + 0.01 * loss_period + 1.0 * loss_spectrum # 仅更新generator的权重 gradients = tape.gradient(total_loss, generator.trainable_variables) optimizer.apply_gradients(zip(gradients, generator.trainable_variables)) return {"total_loss": total_loss, "ssim_loss": loss_ssim, "period_loss": loss_period, "spectrum_loss": loss_spectrum} # 完整训练逻辑实现 epochs = 1000 batch_size = 256 val_split = 0.2 # 拆分训练集和验证集 split_idx = int(len(noise_train) * (1 - val_split)) train_noise, val_noise = noise_train[:split_idx], noise_train[split_idx:] train_ct, val_ct = CT_vector_train[:split_idx], CT_vector_train[split_idx:] train_img, val_img = y_train[:split_idx], y_train[split_idx:] train_period, val_period = period_train[:split_idx], period_train[split_idx:] train_spec, val_spec = full_spec_train[:split_idx], full_spec_train[split_idx:] # 构建数据集 train_dataset = tf.data.Dataset.from_tensor_slices(((train_noise, train_ct), (train_img, train_period, train_spec))).shuffle(1000).batch(batch_size) val_dataset = tf.data.Dataset.from_tensor_slices(((val_noise, val_ct), (val_img, val_period, val_spec))).batch(batch_size) for epoch in range(epochs): print(f"\nEpoch {epoch+1}/{epochs}") # 训练阶段 total_train_loss = 0.0 for step, (inputs, labels) in enumerate(train_dataset): train_losses = train_step(inputs, labels) total_train_loss += train_losses["total_loss"] # 每10步打印一次损失 if step % 10 == 0: print(f"Step {step}: 总损失 {train_losses['total_loss']:.4f}, SSIM损失 {train_losses['ssim_loss']:.4f}, 周期损失 {train_losses['period_loss']:.4f}, 光谱损失 {train_losses['spectrum_loss']:.4f}") avg_train_loss = total_train_loss / (step + 1) # 验证阶段 total_val_loss = 0.0 for step, (inputs, labels) in enumerate(val_dataset): noise, ct_vector = inputs y_img, y_period, y_spectrum = labels image_output, period_output = generator([noise, ct_vector], training=False) spectrum_output = simulator([image_output, period_output], training=False) loss_ssim = ssim_loss_fn(y_img, image_output) loss_period = mse_loss_fn(y_period, period_output) loss_spectrum = mse_loss_fn(y_spectrum, spectrum_output) total_val_loss += 0.05 * loss_ssim + 0.01 * loss_period + 1.0 * loss_spectrum avg_val_loss = total_val_loss / (step + 1) print(f"Epoch {epoch+1} 平均训练损失: {avg_train_loss:.4f}, 平均验证损失: {avg_val_loss:.4f}")
内容的提问来源于stack exchange,提问作者Paris
相关产品推荐
相关产品推荐

