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

如何训练含预训练冻结网络的串联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的可视化图:
    TensorFlow中Final_model可视化图

问题解答

你当前的训练逻辑存在两个核心问题,可先排查修正,效果仍不符合预期再切换自定义训练循环即可:

  1. 损失函数优化方向错误:你定义的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))
  1. 需要显式冻结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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 19:27:00