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

如何在JAX的Adam优化器中实现指数学习率衰减算法

解决方案

以下是修改后的代码,实现学习率从1e-2随迭代指数衰减至1e-4的需求:

import jax
from jax.example_libraries import optimizers

initial_lr = 1e-2
final_lr = 1e-4
train_iters = 100000

# 计算每一步的指数衰减率,确保最终迭代时学习率降至1e-4
decay_rate = (final_lr / initial_lr) ** (1 / train_iters)

# 创建指数衰减的学习率调度器,每一步都更新学习率
step_size = optimizers.exponential_decay(
    initial_lr,
    decay_steps=1,
    decay_rate=decay_rate
)

@jax.jit
def resnet_update(params, opt_state, step):
    value, grads = jax.value_and_grad(objective)(params)
    # 传入当前迭代步数,让优化器计算对应学习率
    opt_state = opt_update(step, grads, opt_state)
    return optimizers.get_params(opt_state), opt_state, value

# 初始化Adam优化器,传入动态学习率调度器
opt_init, opt_update, get_params = optimizers.adam(
    step_size,
    b1=0.9,
    b2=0.999,
    eps=1e-8
)
opt_state = opt_init(params)

for i in range(train_iters):
    params, opt_state, value = resnet_update(params, opt_state, i)
    if i % 1000 == 0:
        # 可选:打印当前学习率,验证衰减效果
        current_lr = step_size(i)
        print(f"Iteration {i:3d} objective {value:.6f} learning rate {current_lr:.6e}")

关键修改说明

  • 计算衰减率:通过公式decay_rate = (final_lr / initial_lr) ** (1 / train_iters)计算每一步的衰减系数,保证第100000次迭代时学习率恰好降至目标值。
  • 动态学习率调度:用optimizers.exponential_decay替代固定学习率,设置decay_steps=1让学习率每一步都更新。
  • 传递迭代步数:修改resnet_update函数,将当前迭代次数i传入opt_update,让优化器根据步数计算对应学习率。
  • 可选日志验证:添加当前学习率打印,方便直观确认衰减过程是否符合预期。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.17 08:35:10