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

TensorFlow中同一层不同变量设置不同学习率的标准方法

在TensorFlow中为同一层内不同变量设置不同学习率的标准方法

当然有更高效的标准方法,无需拆分层就能实现同一层内不同变量的学习率差异化,以下是两种常用方案:

方案一:自定义学习率调度器(LearningRateSchedule)

通过继承tf.keras.optimizers.schedules.LearningRateSchedule,可以根据变量名称或属性动态返回对应学习率,无需修改模型结构,直接在优化器层面实现区分。

示例代码:

import tensorflow as tf

class LayerVarLR(tf.keras.optimizers.schedules.LearningRateSchedule):
    def __init__(self, kernel_lr=0.001, bias_lr=0.005):
        self.kernel_lr = kernel_lr
        self.bias_lr = bias_lr

    def __call__(self, step, var=None):
        # 根据变量名称判断返回对应学习率
        if var is not None:
            if 'kernel' in var.name:
                return self.kernel_lr
            elif 'bias' in var.name:
                return self.bias_lr
        # 默认返回kernel的学习率
        return self.kernel_lr

# 创建带自定义学习率的优化器
optimizer = tf.keras.optimizers.Adam(learning_rate=LayerVarLR())

# 编译模型时使用该优化器
model.compile(optimizer=optimizer, loss='categorical_crossentropy')

方案二:使用tfa.optimizers.MultiOptimizer直接绑定变量

不用拆分层,直接获取目标层的kernel和bias变量,分别为它们分配不同学习率的优化器实例,通过MultiOptimizer完成绑定。这种方式精准控制单个变量的学习率,且变量更新可并行进行,不会损失训练速度。

示例代码:

import tensorflow as tf
import tensorflow_addons as tfa

# 构建模型
model = tf.keras.Sequential([
    tf.keras.layers.Dense(128, activation='relu', name='target_dense')
])

# 获取目标层的kernel和bias变量
target_layer = model.get_layer('target_dense')
kernel_var = target_layer.kernel
bias_var = target_layer.bias

# 为不同变量创建对应学习率的优化器
kernel_optimizer = tf.keras.optimizers.Adam(learning_rate=0.001)
bias_optimizer = tf.keras.optimizers.Adam(learning_rate=0.005)

# 绑定变量与优化器
multi_optimizer = tfa.optimizers.MultiOptimizer([
    (kernel_optimizer, [kernel_var]),
    (bias_optimizer, [bias_var])
])

# 编译模型
model.compile(optimizer=multi_optimizer, loss='mse')

这两种方法都避免了拆分层带来的训练速度损耗,是TensorFlow中实现同一层内不同变量差异化学习率的标准方式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 02:45:11