opt_list:基于千个优化任务学习出的超参数列表,为 TensorFlow、PyTorch 与 JAX 提供即插即用的优化器
2026/9/21 14:31:37 网站建设 项目流程
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

opt_list 是 Google Research 仓库中一个"已学习优化器列表"(Learned Optimizer List)项目:它把在大量机器学习任务上表现良好的超参数整理成一张有序的、可直接按序号取用的超参数序列,让使用者跳过繁琐的搜索空间设计与超参数调优,直接按idx序号逐个尝试即可获得良好效果。本文将从列表的生成背景、存储与加载机制讲起,完整讲解它在 PyTorch、TensorFlow 1.x、TF 2.0/Keras、JAX(Flax / 经典 optimizers / Optix)六大框架中的接入方式,并深入NAdamW + 余弦学习率衰减的底层实现,帮助你快速把这套"无需调参"的优化器方案落地到自己的模型训练中。

一、什么是 Learned Optimizer List:核心思想与使用原则

1.1 它是什么

Learned Optimizer List 是一组按性能排序的超参数序列。这些超参数并非人工拍脑袋设计,而是基于在大量机器学习任务上的表现筛选出来的(其生成方法见论文《Using a thousand optimization tasks to learn hyperparameter search strategies》,对应本仓库的 task_set 目录)。项目作者为 Luke Metz,联系方式见 opt_list/README.md。

它的核心主张是:与其为每个问题纠结如何设计合理的超参数搜索空间、或运行复杂的超参数搜索方法,不如直接尝试这份已经"学习"好的超参数列表。

1.2 使用原则:两步走

官方建议的使用流程非常简洁(opt_list/README.md):

  1. 先确定训练步数:明确你的模型总共要训练多少步(training_steps),因为列表中的每个配置都内置了与之匹配的学习率调度(warmup + 恒定 + 余弦衰减),训练步数不同,调度曲线也不同。
  2. 按顺序尝试多个配置:通过idx参数从 0 开始依次取用列表中的配置。列表是排过序的,通常不到 10 次尝试(under 10 trials)就能得到不错的结果;追求最佳性能时,可以继续尝试直到idx = 100左右。

1.3 重要建议:请移除你的 weight decay

官方明确建议:为获得最佳性能,请关闭任何形式的 L2 正则化或 weight decay(opt_list/README.md)。因为这些优化器配置已经在内部替你管理了正则化——每个配置都同时包含adamw_weight_decay(AdamW 式权重衰减)与l2_weight_decay(L2 损失式权重衰减)两个通道,再叠加外部 weight decay 会造成重复正则化,反而损害效果。

二、列表的存储与加载:idx 背后发生了什么

2.1 两个数据文件 + 一个加载入口

整个列表的实现非常轻量,核心数据由两个文件承载(见 opt_list/opt_list/ 目录):

  • nadamw_opt_list.txt:1000 行优化器配置名,每行一个,顺序即idx的排序顺序(文件开头如nadamw_grid_seed251nadamw_grid_seed611…)。
  • nadamw_configs.json:1001 个优化器配置字典,每个配置名为键,值为[配置字典, "nadamw_grid"]元组,记录了该配置所属的搜索网格。

加载入口在 common.py 的get_optimizer_config(idx)

def get_optimizer_config(idx): """Get the optimizer config from the list of hparams at the given index.""" names = [x.strip() for x in _get_opt_name_content().split("\n") if x.strip()] name_to_use = names[idx] config, _ = _get_config_map()[name_to_use] logging.info("Using config:: %s", str(config)) return config

即:按idx从名字列表取第idx个配置名,再到 JSON 中查表拿到对应的超参数字典。值得注意的是配置文件位于包内(os.path.dirname(__file__)),并且 setup.py 通过package_data={'opt_list': ['*.json', '*.txt']}将这两个数据文件一并打包,因此 pip 安装后无需额外下载任何数据。

2.2 每个配置包含哪些超参数

以 JSON 中首个配置为例(nadamw_configs.json):

{ "learning_rate": 0.00013043019947775472, "beta1": 0.6184381722189567, "beta2": 0.9994057151352359, "epsilon": 5.7755507421823396e-05, "use_nesterov": false, "adamw_weight_decay": 0.010296116606084241, "l2_weight_decay": 0.000569334985961975, "warmup_fraction": 0.014360184808201766, "min_learning_rate_mult": 0.0, "constant_fraction": 0.25251538281125563 }

从 JSON 中可以看到,列表中的超参数覆盖了很广的取值分布(例如learning_rate1e-5量级到0.6量级,epsilon1e-8到数千,beta1从约0.02到接近0.999),这正是"在大量任务上学出来的多样性"。各字段含义详见后文第五节。

三、安装方式

官方提供两种安装途径(opt_list/README.md):

方式一:pip 直接从 GitHub 安装(命令形如pip install git+<仓库地址>#subdirectory=opt_list,即从仓库的 opt_list 子目录安装)。

方式二:克隆仓库后本地安装

git clone https://gitcode.com/gh_mirrors/go/google-research cd google-research/opt_list/ pip install -e .

可编辑模式(-e)安装后即可在任意项目中使用from opt_list import ...。setup.py 要求 Python >= 3.6,依赖项见 requirements.txt。安装后,六个框架的完整示例可以通过以下命令直接运行(opt_list/opt_list/examples/ 目录,需先进入google-research/opt_list/目录):

python3 -m opt_list.examples.torch python3 -m opt_list.examples.tf_v1 python3 -m opt_list.examples.tf_keras python3 -m opt_list.examples.jax_flax python3 -m opt_list.examples.jax_optimizers python3 -m opt_list.examples.jax_optix

四、六大框架接入示例(完整代码)

所有示例 API 的共同点:传入idxtraining_steps两个参数即可获得一个现成优化器,无需任何额外配置。以下代码均取自 opt_list/README.md。

4.1 PyTorch

from opt_list import torch_opt_list opt = torch_opt_list.optimizer_for_idx(model.parameters(), idx=0, training_steps=training_steps) for i in range(training_steps): loss = forward() opt.zero_grad() loss.backward() opt.step()

optimizer_for_idx(params, idx, training_steps)返回的NadamWCosineDecay继承自torch.optim.optimizer.Optimizer(见 torch_opt_list.py),因此与标准 torch 优化器 API 完全兼容,可无缝替换torch.optim.Adam等。完整的可运行脚本见 examples/torch.py,其中用一个 2→256→256→2 的全连接网络演示了从创建优化器到zero_grad / backward / step的完整训练循环。

4.2 TensorFlow 1.x

from opt_list import tf_opt_list global_step = tf.train.get_or_create_global_step() opt = tf_opt_list.optimizer_for_idx(0, training_steps, iteration=global_step) train_op = opt.minimize(loss) with tf.Session() as sess: for i in range(training_iters): sess.run([train_op])

optimizer_for_idx(idx, training_steps, iteration)返回基于tf1.train.AdamOptimizer封装的NAdamWOptimizer(tf_opt_list.py)。源码特别说明:iteration(全局步数)用于驱动学习率调度,不传时默认取tf.train.get_or_create_global_step()并打印警告,因此建议显式传入。

4.3 TensorFlow 2.0 / Keras

from opt_list import tf_opt_list opt = tf_opt_list.keras_optimizer_for_idx(0, training_steps) model.compile(loss='mse', optimizer=opt, metrics=[]) for i in range(training_steps): model.train_on_batch(inp, target)

keras_optimizer_for_idx(tf_opt_list.py)内部会用配置中的learning_rate / min_learning_rate_mult / constant_fraction / warmup_fraction构造一个CustomCosineDecay学习率调度器,再包装成 Keras 优化器,因此能直接用于model.compile(optimizer=...)

4.4 JAX:Flax

from opt_list import jax_flax_opt_list optimizer_def = jax_flax_opt_list.optimizer_for_idx(idx=0, training_steps) optimizer = optimizer_def.create(model) for i in range(training_steps): optimizer, loss = optimizer.optimize(loss_fn)

optimizer_for_idx返回一个NAdamWCosineDecayOptimizerDef(jax_flax_opt_list.py),使用方式与标准 Flax Optimizer 一致。

4.5 JAX:经典 Optimizers(jax.example_libraries.optimizers)

from opt_list import jax_optimizers_opt_list opt_init, opt_update, get_params = jax_optimizers_opt_list.optimizer_for_idx( 0, training_iters) opt_state = opt_init(params) for i in range(training_steps): params = get_params(opt_state) opt_state = opt_update(i, jax.grad(loss_fn)(params, batch), opt_state)

返回经典 JAX 优化器风格的(init_fun, update_fun, get_params)三元组(jax_optimizers_opt_list.py),其中update_fun的第一个参数是当前步数i,用于驱动学习率调度。

4.6 JAX:Optix(jax.experimental.optix)

from opt_list import jax_optix_opt_list opt = jax_optix_opt_list.optimizer_for_idx(idx=0, training_steps) opt_state = opt.init(params) for i in range(training_steps): grads = jax.grad(loss_fn)(params, batch) # Not opt.update! We need parameter values too! updates, opt_state = opt.update_with_params(grads, params, opt_state) params = optix.apply_updates(params, updates)

注意:官方明确指出,Optix 目前不支持 AdamW 式权重衰减,因此该接入方式不是完全即插即用(drop-in)的替代品,而是遵循类似 API 的实现(opt_list/README.md)。其update_with_params相比标准 Optix 多接收一个params参数,因为 AdamW 权重衰减需要用到参数值本身(jax_optix_opt_list.py)。

五、底层原理:NAdamW 优化器与余弦学习率调度

5.1 统一的 NAdamW 更新规则

列表中的所有配置最终都作用于同一个优化器族:NAdamW(涵盖 Nadam / Adam / AdamW / NadamW 四种形态,由use_nesterov与两种 weight decay 开关组合而成)。以 PyTorch 实现 torch_opt_list.py 与 JAX 实现 jax_common.py 为例,每步更新可分解为:

  1. 学习率调度lr = get_cosine_learning_rate_fn(training_steps, ...)(step),即按当前步数动态取学习率(见 5.2)。
  2. L2 权重衰减注入梯度grad = grad - param * l2_weight_decay,模拟在 loss 中加入 L2 项的效果。
  3. 一阶/二阶矩更新grad_ema = beta1 * grad_ema + (1-beta1) * gradgrad_sq_ema = beta2 * grad_sq_ema + (1-beta2) * grad^2
  4. 偏差校正(可选,默认开启):lr_t = lr * sqrt(1 - beta2^t) / (1 - beta1^t),对应 Kingma & Ba 论文中 Section 2.1 之前的 "epsilon hat" 校正。
  5. Nesterov 加速(可选):use_nesterov为真时用(beta1 * grad_ema + (1-beta1) * grad)作为分子,否则用grad_ema
  6. AdamW 权重衰减step = step + lr_t * adamw_weight_decay * param,即衰减与梯度解耦、直接作用于参数更新量。

可见配置表中的use_nesterovadamw_weight_decayl2_weight_decayuse_bias_correction等字段直接对应上述算法分支——同一个优化器类,靠配置字典切换出不同形态的优化器

5.2 三段式学习率调度:warmup → 恒定 → 余弦衰减

get_cosine_learning_rate_fn(torch_opt_list.py、jax_common.py)返回一个"输入步数、输出学习率"的函数,其调度曲线分三个阶段:

  • Warmup 阶段:前warmup_fraction * training_steps步内,学习率从 0 线性上升到learning_rate(若warmup_fraction为 0 则跳过)。
  • 恒定阶段:从 warmup 结束直到constant_fraction * training_steps,学习率保持为learning_rate
  • 余弦衰减阶段:从constant_steps开始,沿余弦函数的单调递减段(0 到 π/2)将学习率从learning_rate衰减到min_learning_rate_mult * learning_rate

源码注释指出,这种余弦衰减已被证实在大规模语言建模(GPT、Megatron-LM)与图像分类任务中行之有效(对应经典论文《SGDR: Stochastic Gradient Descent with Warm Restarts》)。这正是"必须预先告知训练步数"的原因:调度曲线完全由training_steps归一化,步数不同则同一idx对应的学习率轨迹也不同。

六、配置字段速查表

综合 nadamw_configs.json 与源码中NAdamWHyperParams的定义(jax_common.py),列表中的 10 个超参数字段含义如下:

字段含义典型取值(来自列表)
learning_rate基础学习率,warmup 后的目标值、衰减起点1e-5~0.6
beta1一阶矩(梯度均值)指数衰减率0.02~0.999
beta2二阶矩(梯度平方均值)指数衰减率0.04~0.99998
epsilon数值稳定常数(论文中的 "epsilon hat")1e-8~ 数千
use_nesterov是否启用 NAdam 式 Nesterov 加速true/false
adamw_weight_decayAdamW 式权重衰减(与梯度解耦)0~0.09
l2_weight_decayL2 损失式权重衰减(注入梯度)0~0.03
warmup_fraction线性 warmup 占训练总步数的比例0~0.04
constant_fraction恒定学习率段占训练总步数的比例(含 warmup)0.05~0.94
min_learning_rate_mult余弦衰减终点的学习率乘数0~1

七、注意事项与适用边界

  1. 训练步数要如实传入training_steps决定学习率调度曲线的形状,请传入实际训练总步数;torch_opt_list.optimizer_for_idxjax_flax_opt_list.optimizer_for_idx会在内部把training_steps注入配置。
  2. Optix 接口的例外:如 4.6 节所述,Optix 版本因不支持 AdamW 权重衰减而并非完全即插即用,使用update_with_params而非opt.update
  3. 不要叠加外部正则化:配置已内含两种 weight decay 通道,外部再开启 L2 / weight decay 会影响效果(见 1.3 节)。
  4. 按序尝试、控制次数:从idx=0开始按顺序尝试,10 次以内通常足够;追求最优可尝试到 100,但收益递减。
  5. 适用前提:列表中配置基于 Nadam/Adam 优化器族与余弦衰减调度设计,适用于"给定固定训练步数"的标准监督训练场景;其适用范围与局限可在自己的任务上直接验证(README 也邀请使用者反馈"哪里好用、哪里不好用")。

八、延伸阅读

  • 列表生成方法对应的论文与任务集:本仓库的 task_set 目录(README 中明确指出列表基于 task_set 中的千级优化任务学习得到)。
  • 数据文件:优化器配置名排序表 nadamw_opt_list.txt(1000 行)、配置字典 nadamw_configs.json(1001 个配置)。
  • 核心实现:common.py(配置加载)、torch_opt_list.py(PyTorch 实现)、tf_opt_list.py(TF1 / Keras 实现)、jax_common.py(JAX 共享更新逻辑)。
  • 六个框架的完整可运行示例:opt_list/opt_list/examples/。
  • 人工智能
  • 深度学习
  • NLP
  • 计算机视觉
  • 强化学习

【免费下载链接】google-research

Google Research

项目地址:https://gitcode.com/gh_mirrors/go/google-research
点击查看免费下载

相关推荐

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询