- 人工智能
- 深度学习
- NLP
- 计算机视觉
- 强化学习
【免费下载链接】google-research
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):
- 先确定训练步数:明确你的模型总共要训练多少步(
training_steps),因为列表中的每个配置都内置了与之匹配的学习率调度(warmup + 恒定 + 余弦衰减),训练步数不同,调度曲线也不同。 - 按顺序尝试多个配置:通过
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_seed251、nadamw_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_rate从1e-5量级到0.6量级,epsilon从1e-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 的共同点:传入idx与training_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 为例,每步更新可分解为:
- 学习率调度:
lr = get_cosine_learning_rate_fn(training_steps, ...)(step),即按当前步数动态取学习率(见 5.2)。 - L2 权重衰减注入梯度:
grad = grad - param * l2_weight_decay,模拟在 loss 中加入 L2 项的效果。 - 一阶/二阶矩更新:
grad_ema = beta1 * grad_ema + (1-beta1) * grad、grad_sq_ema = beta2 * grad_sq_ema + (1-beta2) * grad^2。 - 偏差校正(可选,默认开启):
lr_t = lr * sqrt(1 - beta2^t) / (1 - beta1^t),对应 Kingma & Ba 论文中 Section 2.1 之前的 "epsilon hat" 校正。 - Nesterov 加速(可选):
use_nesterov为真时用(beta1 * grad_ema + (1-beta1) * grad)作为分子,否则用grad_ema。 - AdamW 权重衰减:
step = step + lr_t * adamw_weight_decay * param,即衰减与梯度解耦、直接作用于参数更新量。
可见配置表中的use_nesterov、adamw_weight_decay、l2_weight_decay、use_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_decay | AdamW 式权重衰减(与梯度解耦) | 0~0.09 |
l2_weight_decay | L2 损失式权重衰减(注入梯度) | 0~0.03 |
warmup_fraction | 线性 warmup 占训练总步数的比例 | 0~0.04 |
constant_fraction | 恒定学习率段占训练总步数的比例(含 warmup) | 0.05~0.94 |
min_learning_rate_mult | 余弦衰减终点的学习率乘数 | 0~1 |
七、注意事项与适用边界
- 训练步数要如实传入:
training_steps决定学习率调度曲线的形状,请传入实际训练总步数;torch_opt_list.optimizer_for_idx与jax_flax_opt_list.optimizer_for_idx会在内部把training_steps注入配置。 - Optix 接口的例外:如 4.6 节所述,Optix 版本因不支持 AdamW 权重衰减而并非完全即插即用,使用
update_with_params而非opt.update。 - 不要叠加外部正则化:配置已内含两种 weight decay 通道,外部再开启 L2 / weight decay 会影响效果(见 1.3 节)。
- 按序尝试、控制次数:从
idx=0开始按顺序尝试,10 次以内通常足够;追求最优可尝试到 100,但收益递减。 - 适用前提:列表中配置基于 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
相关推荐
思源笔记 v3.1.3 版本解析:细节缺陷修复、索引与数据一致性改进全景解读
思源笔记 v3.1.3 版本解析:细节缺陷修复、索引与数据一致性改进全景解读 本文基于思源笔记(SiYuan)开源仓库中 v3.1.3 版本变更记录 https
人工智能深度学习NLP计算机视觉强化学习超参数调优:PyTorch深度学习模型优化策略
超参数调优:PyTorch深度学习模型优化策略 深度学习模型训练过程中,超参数调优(Hyperparameter Tuning)是决定模型性能的关键环节。本文将
示例工程教程3步告别挂号焦虑:91160-cli医疗预约自动化实践指南
3步告别挂号焦虑:91160 cli医疗预约自动化实践指南 当城市清晨的第一缕阳光还未洒进窗户,你已经在手机屏幕前紧张地等待。手指悬停在"预约"按钮上方,心跳随
CLI网页爬虫
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考