如何在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
相关产品推荐
相关产品推荐

