当前位置: 首页 > news >正文

TorchOpt API参考速查:从优化器到Transform的完整函数地图

TorchOpt API参考速查:从优化器到Transform的完整函数地图

【免费下载链接】torchoptTorchOpt is an efficient library for differentiable optimization built upon PyTorch.项目地址: https://gitcode.com/gh_mirrors/to/torchopt

TorchOpt是一个构建在 PyTorch 之上的高效可微优化(differentiable optimization)库,提供 Optax 风格的功能式 API:用"梯度变换"(Transform)自由组合出优化器,并内置隐式梯度、零阶梯度等可微回传能力。本文整理 TorchOpt 从优化器、Transform 到工具函数、元优化的完整函数地图,帮你 10 分钟定位任意 API。

🗺️ 一分钟看懂:TorchOpt 的双模式 API

TorchOpt 的所有 API 围绕一个核心抽象GradientTransformation(定义在 torchopt/base.py)展开——它是一对纯函数(init_fn, update_fn)init根据参数生成优化器状态,update把梯度变换成最终更新量。

风格入口特点源码位置
功能式(推荐)opt.sgd(lr)opt.adam(lr)返回GradientTransformation,可任意组合torchopt/alias/__init__.py
面向对象opt.optim.Adam(model.parameters(), lr=...)兼容torch.optim使用习惯torchopt/optim/__init__.py

两种风格共享同一套 Transform 底层,功能式风格还能无缝对接chainclip_grad_norm等组合工具。

🚀 优化器API速查表:8大预设,两种调用方式

TorchOpt 内置 8 类主流优化器,功能式别名(torchopt/alias/)与面向对象类(torchopt/optim/)一一对应,另有Meta*元优化变体支持二阶可微:

优化器功能式 APIOO 类元优化变体
SGDopt.sgd(lr, momentum, nesterov)SGDMetaSGD
AdaGradopt.adagrad(lr, eps)AdaGradMetaAdaGrad
AdaDeltaopt.adadelta(eps)AdaDeltaMetaAdaDelta
Adamopt.adam(lr, betas, eps)AdamMetaAdam
AdamWopt.adamw(lr, betas, eps, wd)AdamWMetaAdamW
AdaMaxopt.adamax(lr, betas)AdamaxMetaAdaMax
RAdamopt.radam(lr, betas, eps)RAdamMetaRAdam
RMSPropopt.rmsprop(lr, decay, eps)RMSPropMetaRMSProp

💡 元优化变体统一位于 torchopt/optim/meta/,例如MetaAdam对超参数也保持可微,是 MAML 等元学习算法的基础。

🧩 Transform梯度变换组件:14个预设函数

torchopt.transform模块(torchopt/transform/init.py)提供 14 个预设变换,是搭建自定义优化器的"乐高积木":

函数作用典型用途
scale(step_size)按固定系数缩放更新量学习率
scale_by_schedule(scheduler)按调度函数缩放学习率衰减
scale_by_adam(betas, eps)Adam 自适应缩放组合 Adam
scale_by_adamax(betas, eps)AdaMax 自适应缩放组合 AdaMax
scale_by_radam(betas, eps)RAdam 自适应缩放组合 RAdam
scale_by_rms / scale_by_rss / scale_by_adadelta基于均方根等统计量缩放组合 RMS 类优化器
scale_by_stddev按标准差缩放自适应学习率
add_decayed_weights(wd)L2 权重衰减AdamW 组合
masked(mask)按掩码屏蔽更新稀疏/冻结参数
nan_to_num(nan, posinf, neginf)替换nan/inf梯度训练稳定性
trace(order)一阶迹估计(随机迹)隐式梯度计算
chain(*transforms)串联多个变换组装优化器

🔗 用 chain 一行组合优化器

chain(torchopt/combine.py)把多个 Transform 串成流水线,再叠加clip_grad_norm(torchopt/clip.py)做梯度裁剪,三行代码即可等价于一个带裁剪的 Adam:

import torchopt as opt adam_with_clip = opt.chain( opt.transform.scale(0.001), opt.transform.scale_by_adam(), opt.clip_grad_norm(1.0) )

配套的opt.apply_updates(params, updates)(torchopt/update.py)负责把变换后的更新量写回参数,支持inplace原地更新。

⏱️ 调度器与实用工具函数

  • 学习率调度(torchopt/schedule/init.py):linear_schedulepolynomial_scheduleexponential_decay,可直接传给scale_by_schedule
  • 停梯度stop_gradient(torchopt/utils.py)阻断对张量的反向传播,是构建可微优化器时的关键技巧;
  • 状态管理extract_state_dict/recover_state_dict用于取出和恢复模块状态,module_clone/module_detach_提供模块克隆与原地 detach;
  • 梯度 Hookregister_hooknan_to_num_hookzero_nan_hook(torchopt/hook.py),在update前对梯度做自定义拦截;
  • PyTree 工具tree_maptree_flatten等(torchopt/pytree.py),基于 optree 处理参数树结构。

🧠 元优化与三种可微回传模式

TorchOpt 的最大亮点是可微优化:把优化步骤本身也放进计算图。它提供三种回传模式(源码位于 torchopt/diff/):

模式API原理适用场景
显式梯度默认torch.autograd直接对整段优化代码求导步数少、结构简单
隐式梯度opt.diff.implicit.custom_rootImplicitMetaGradientModule隐函数梯度定理,跳过展开步数多、内存友好(iMAML)
零阶梯度opt.diff.zero_order.zero_order有限差分离散扰动估计黑箱、不可导场景

配套模块 torchopt/nn/module.py 提供MetaGradientModuleImplicitMetaGradientModuleZeroOrderGradientModule以及reparameterizeswap_state等工具,让你把任意nn.Module包装成支持元梯度回传的模块。

以隐式 MAML(iMAML)为例,结合MetaSGDcustom_root训练出的 few-shot 模型在 Omniglot 上的精度曲线如下(完整实现见 examples/iMAML/imaml_omniglot_functional.py):

📊 可视化与进阶模块

  • 计算图可视化make_dotresize_graph(torchopt/visual.py),用 Graphviz 渲染 TorchOpt 的计算图,比 torchviz 更能保留元梯度结构:

  • 加速算子torchopt.accelerated_op提供 CUDA 加速的 Adam 算子,与torch.optim兼容,可用opt.accelerated_op_available()检测;
  • 线性求解:torchopt/linalg/ 提供共轭梯度(cg)与非线性共轭梯度(ns),torchopt/linear_solve/ 封装solve_cgsolve_invsolve_normal_cg,用于隐式梯度的线性系统求解;
  • 分布式训练:torchopt/distributed/ 提供parallelizeparallelize_sync等 RPC 并行原语,支持多进程并行训练元优化器(参考 examples/distributed/few-shot/maml_omniglot.py)。

⚡ 速查小结

我想……用什么
快速搭建 Adamopt.chain(opt.transform.scale(lr), opt.transform.scale_by_adam())
兼容 torch.optim 写法opt.optim.Adam(params, lr=...)
训练可微调的优化器opt.MetaAdam/MetaSGD
减少长序列优化的显存opt.diff.implicit.custom_root
处理黑箱不可导目标opt.diff.zero_order.zero_order
裁剪 / 清洗梯度opt.clip_grad_norm(1.0)opt.nan_to_num
调试计算图opt.visual.make_dot

所有公开 API 均从 torchopt/init.py 统一导出,配套测试用例在 tests/ 目录下,按模块一一对应(如test_optim.pytest_transform.pytest_implicit.py),可作为每个函数的用法示例快速查阅。

【免费下载链接】torchoptTorchOpt is an efficient library for differentiable optimization built upon PyTorch.项目地址: https://gitcode.com/gh_mirrors/to/torchopt

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

http://www.cnnetsun.cn/news/4225345.html

相关文章:

  • 单靠死工资不够?这本《AI创富手册》教你用AI开启第二收入曲线
  • JobRadar:基于本地大模型的职位匹配评分与智能筛选指南
  • 推理增强工程实践:DeepSeek-Reasonix与esengine组合落地指南
  • 贪心算法的思路和典型例题
  • AI找矿实战:如何用机器学习圈定高纯石英靶区
  • Vue3数字输入框(InputNumber)
  • 基于MCP Server构建AI可查询的错误知识库:从协议到实践
  • MCP Server实践:AI错误诊断工具,让报错不再难懂
  • 来看界面控件DevExtreme如何实现数据表单的高效动态更新
  • AI网站构建器背后:结构化页面表示与生成链路解析
  • 基于高德API与异步Python构建智能选址AI Skill实战
  • Matlab数据导入实战:从格式兼容到内存优化
  • 数维杯B题建模思路1.0:数据沼泽中的最小可行闭环
  • Codex限流与配置故障排查:从429到config.toml修复指南
  • DeepSeek V4 Flash测评框架:性能、延迟与成本控制实战
  • Gemini反代API工程指南:密钥、协议转换与排查
  • 用户价值分析最小闭环:从埋点到RFM分群与流失预警
  • 20天高效备战大厂面试:策略与实战指南
  • VMware Workstation安装Windows 11虚拟机完整指南与踩坑排查
  • HyperMesh与Inspire协同:拓扑优化到尺寸优化完整流程
  • PostgreSQL与MySQL语法差异详解:从建表到高级查询的实战对比
  • 当汽车电机控制器遇上工业液冷电源:热管理驱动的跨界机遇
  • CodeX、Ollama、Coze多智能体协作:企业级AI编码工作流实战
  • MATLAB fmincon非线性规划实战:从报错到收敛的完整指南
  • C语言宏定义括号规范:避免运算符优先级陷阱与副作用风险
  • 机器学习数据预处理:标准化、归一化与正则化的原理与应用
  • 超低功耗信号处理实战:从数据搬运到事件驱动的能效设计
  • Dinic算法性能飞跃:详解当前弧优化原理与实战代码
  • 共享单车调度优化建模实战:从问题解构到三层决策框架
  • 二进制速率乘法器(BRM)原理、Verilog实现与工程实战