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

在自定义训练循环的@tf.function中实现学习率衰减解决I2I模型损失停滞

解决自定义训练循环中基于步数的学习率衰减问题

你代码的核心问题是定义了指数衰减学习率调度器,但初始化优化器时并未使用它,而是传入了固定学习率,导致学习率始终不会随训练步数变化。以下是具体修正方案:

1. 正确绑定学习率调度器到优化器

直接将你定义的ExponentialDecay调度器传给Adam优化器,替代固定学习率:

import tensorflow as tf

# 初始化学习率调度器
lr_schedule = tf.keras.optimizers.schedules.ExponentialDecay(
    initial_learning_rate=1e-3,  # 对应你的rate参数
    decay_steps=350,             # 对应steps_per_epoch
    decay_rate=0.96,
    staircase=True               # 开启后每满decay_steps才衰减一次,符合按epoch衰减的需求
)
# 绑定调度器到优化器
generator_optimizer = tf.keras.optimizers.Adam(learning_rate=lr_schedule)

2. 修正自定义训练步骤(train_step)

Keras优化器会自动根据自身的iterations属性(全局训练步数)计算当前学习率,无需手动干预。你可以添加学习率监控,确认衰减是否生效:

@tf.function
def train_step(input_image, target, step):
    with tf.GradientTape() as gen_tape:
        gen_output = generator(input_image, training=True)
        gen_l1_loss = generator_loss(gen_output, target)

    # 计算梯度并更新参数
    generator_gradients = gen_tape.gradient(gen_l1_loss, generator.trainable_variables)
    generator_optimizer.apply_gradients(zip(generator_gradients, generator.trainable_variables))

    # 记录损失和当前学习率
    with summary_writer.as_default():
        tf.summary.scalar('gen_l1_loss', gen_l1_loss, step=step)
        # 获取当前学习率并写入日志
        current_lr = generator_optimizer.learning_rate(generator_optimizer.iterations)
        tf.summary.scalar('learning_rate', current_lr, step=step)

3. 实现提前停止逻辑(利用Patience参数)

你代码中的Patience参数未实际使用,添加验证损失跟踪逻辑,当连续指定步数损失未下降时停止训练:

def fit(train_ds, val_ds, steps_per_epoch, patience, total_steps):
    best_val_loss = float('inf')
    wait = 0

    for step, (input1, input2, input3, target) in train_ds.repeat().take(total_steps).enumerate():
        input_images = tf.concat([input1, input2, input3], axis=-1)
        train_step(input_images, target, step)

        # 每个epoch结束后执行验证
        if step % steps_per_epoch == 0 and step != 0:
            epoch = step // steps_per_epoch
            val_loss = 0.0
            val_step_count = 0

            # 遍历验证集计算平均损失
            for val_in1, val_in2, val_in3, val_target in val_ds:
                val_input = tf.concat([val_in1, val_in2, val_in3], axis=-1)
                val_output = generator(val_input, training=False)
                val_loss += generator_loss(val_output, val_target).numpy()
                val_step_count += 1
            val_loss /= val_step_count

            # 记录验证损失
            with summary_writer.as_default():
                tf.summary.scalar('val_l1_loss', val_loss, step=epoch)

            # 提前停止判断
            if val_loss < best_val_loss:
                best_val_loss = val_loss
                wait = 0
                # 可选:保存最优模型权重
                generator.save_weights('best_generator_weights.h5')
            else:
                wait += 1
                if wait >= patience:
                    print(f"验证损失连续{patience}个epoch未下降,提前停止训练")
                    break

4. 调用修正后的fit函数

注意调整参数顺序,匹配函数定义:

fit(
    train_dataset, 
    val_dataset, 
    steps_per_epoch=350, 
    patience=10, 
    total_steps=NumSteps + 1
)

为什么之前的方法无效?

  • tf.config.experimental_run_functions_eagerly(True):强制eager模式会丢失图优化特性,且固定学习率会被固化为常量,无法动态更新;
  • tf.assign:手动修改优化器的学习率属性在图模式下极易出错,因为Keras优化器的学习率如果是常量张量,无法被赋值覆盖;而内置调度器是动态依赖迭代步数计算的,无需手动干预。

内容的提问来源于stack exchange,提问作者Ruben van Oosterhoudt

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 16:10:54