
深度学习科学计算【免费下载链接】torchdiffeqDifferentiable ODE solvers with full GPU support and O(1)-memory backpropagation.项目地址https://gitcode.com/gh_mirrors/to/torchdiffeq点击查看免费下载torchdiffeq 是一个完全基于 PyTorch 实现的常微分方程ODE求解器库提供统一入口odeint求解初值问题IVP并支持通过伴随方法adjoint method以 O(1) 常量内存开销完成反向传播。本文以官方 README 为骨架结合 FURTHER_DOCUMENTATION.md、FAQ.md 与仓库源码系统讲解安装、基本用法、求解器选型、全部关键字参数、事件处理及示例代码帮助你在 Neural ODE、物理模拟等场景中直接落地使用。一、库是什么核心能力一览torchdiffeq 解决了深度学习中的一个基础问题如何把求解微分方程这一过程当作一个可微层differentiable layer嵌入神经网络训练流程。其核心能力包括可微的 ODE 求解对初值问题dy/dt f(t, y)、y(t₀) y₀进行数值积分并对所有主要输入参数实现梯度GPU 全支持由于所有求解器均用 PyTorch 实现算法可完全运行在 GPU 上见 README 声明O(1) 内存反向传播通过伴随方法odeint_adjoint把反向传播化为一次额外的伴随 ODE 求解内存占用与步数无关事件处理支持odeint_event允许根据事件函数如碰撞检测提前终止求解并支持对事件时间与终态的微分丰富的求解器家族自适应步长与固定步长共十余种算法还封装了 SciPy 的全部求解器。其库入口定义在 torchdiffeq/init.py公开接口为odeint、odeint_adjoint、odeint_event与odeint_dense当前版本号为0.2.5。二、安装安装最新稳定版直接使用 pippip install torchdiffeq如需安装 GitHub 上的最新开发版pip install githttps://github.com/rtqichen/torchdiffeq在 Windows 上若安装遇到问题可以下载代码后直接执行python setup.py install仓库内提供 setup.py。项目本身依赖 PyTorch请确保环境中已安装可用且版本兼容的 PyTorch。三、基本用法odeint 求解初值问题3.1 概念什么是初值问题一个初值问题Initial Value Problem, IVP由一个 ODE 和一个初始值组成dy/dt f(t, y) y(t_0) y_0ODE 求解器的目标是找到一条满足该 ODE 且穿过初始条件的连续轨迹。torchdiffeq 将这一过程封装为一次函数调用。3.2 最小调用示例使用默认求解器求解 IVPfrom torchdiffeq import odeint odeint(func, y0, t)三个核心参数的含义func任意可调用对象实现微分方程f(t, x)。接收标量时间t与状态张量y返回状态对时间的导数y也可以是一个张量元组见源码 torchdiffeq/_impl/odeint.pyy0任意维度any-D的 Tensor表示初始值t1-D Tensor包含求值时刻点初始时间取t[0]时间序列可以递增也可以递减。返回的张量y的第一维对应t中的各个时刻点y0作为第一维的首个元素被包含在结果中对应integrate循环中solution[0] self.y0的实现见 torchdiffeq/_impl/solvers.py。3.3 参数类型与规范性检查源码 torchdiffeq/_impl/misc.py 中的_check_inputs对输入做了严格的规范化与校验理解这些规则可以避免踩坑t必须是 1-D 浮点 Tensor且严格递增或严格递减若y0是元组tuple则rtol/atol也须是等长的元组内部会做扁平化与形状还原_flat_to_shapemethod缺省时为dopri5若传入不在SOLVERS字典中的方法名会抛出ValueErrort与y0不在同一设备时t会被自动转换到y0.device并给出兼容性警告。3.4 前向求导的注意事项直接对odeint求反向传播会穿过求解器内部实现这对所有求解器并非数值上都稳定不过对默认的dopri5通常没问题。为此官方推荐使用伴随方法见下一节。四、伴随方法 odeint_adjointO(1) 内存反向传播4.1 用法from torchdiffeq import odeint_adjoint as odeint odeint(func, y0, t)odeint_adjoint只是对odeint的封装前向过程与普通odeint完全一致但在反向调用中额外求解一个伴随 ODEadjoint ODE从而把内存占用压缩到 O(1)与求解步数无关代价是多解一次 ODE。实现位于 torchdiffeq/_impl/adjoint.py核心反向逻辑封装在自定义 autograd FunctionOdeintAdjointMethod同文件 L8-L153中。4.2 最大的坑func 必须是 nn.Module使用伴随方法时func必须是一个torch.nn.Module。这是因为伴随反向传播需要收集微分方程的参数当未显式指定adjoint_params时库通过func.parameters()find_parameters见 adjoint.py自动获取参数集合。源码中对应检查为if adjoint_params is None and not isinstance(func, nn.Module): raise ValueError(func must be an instance of nn.Module to specify the adjoint parameters; ...)两条例外规则详见 FURTHER_DOCUMENTATION.md通过adjoint_params显式传入参数元组后func不必是nn.Module若func没有任何参数则必须显式指定adjoint_params()。4.3 伴随反向的工作原理源码视角反向传播构建了一个增广状态[vjp_t, y, vjp_y, vjp_params]时间方向的 vjp、状态、状态梯度与参数梯度并定义增广动力学augmented_dynamics在原系统动力学上叠加对y的向量-雅可比积VJP和对时间、参数的积分器见 adjoint.py。随后从t[-1]到t[0]逐段反向求解该增广 ODE并在每段起点用前向保存的y[i-1]覆盖求解结果以保证数值一致性L124-L141。这也是反向传递能够只用 O(1) 内存、可随意增加步数的本质原因。需要指出前向过程本身仍需按求解器保存若干中间状态例如dopri5需要保存 6 个中间状态这是伴随内存 O(1)在梯度反向层面成立的原因——它不保存前向的每一步而是重新反向积分一次。五、odeint(_adjoint) 的关键字参数与求解器清单5.1 通用关键字参数odeint与odeint_adjoint共享以下参数签名见 odeint.py参数说明默认值rtol相对容差relative tolerance自适应求解器用于接受/拒绝步长1e-7atol绝对容差absolute tolerance1e-9method求解器名称见下方清单dopri5options求解器专属选项字典详见 FURTHER_DOCUMENTATION.md仅在显式指定method时才能使用None关于容差的底层语义误差容限按atol rtol * max(|y0|, |y1|)计算见 misc.py 的_compute_error_ratio即绝对容差 相对容差 × 当前状态范数其中范数默认是对单张量输入取 RMS 范数、对元组输入取混合 L-infinity/RMS 范数misc.py。自适应步长求解器的第 k 次误差估计会换算为误差比率再结合safety、ifactor、dfactor计算下一步最优步长_optimal_step_sizemisc.py。5.2 自适应步长Adaptive-step求解器method 名称算法说明dopri8Dormand-Prince-Shampine 8 阶 Runge-Kutta最高阶的自适应算法dopri5Dormand-Prince-Shampine 5(4) 阶 Runge-Kutta[默认]对大多数问题推荐bosh3Bogacki-Shampine 3 阶 Runge-Kutta低阶快速fehlberg2Runge-Kutta-Fehlberg 2 阶最低阶自适应adaptive_heun2 阶 Runge-Kutta自适应 Heun低阶快速实现文件对应为 dopri8.py、dopri5.py、bosh3.py、fehlberg2.py、adaptive_heun.py。所有自适应 RK 求解器共享同一套步长控制骨架 rk_common.py。5.3 固定步长Fixed-step求解器method 名称算法说明euler欧拉法一阶midpoint中点法二阶rk4四阶 Runge-Kutta3/8 规则精度/速度均衡需配合optionsdict(step_size...)explicit_adams显式 Adams-Bashforth多步法implicit_adams隐式 Adams-Bashforth-Moulton多步预测-校正实现位于 fixed_grid.pyeuler/midpoint/rk4 等与 fixed_adams.py显式/隐式 Adams。另外SOLVERS注册表中还包含heun2、heun3、implicit_euler、implicit_midpoint、trapezoid、radauIIA3、radauIIA5、gl4、gl6、sdirk2、trbdf2等额外隐式/高阶方法见 odeint.py以及为向后兼容保留的fixed_adams别名。5.4 scipy_solver桥接 SciPy此外所有 SciPy 求解器都被封装可通过scipy_solver使用实现见 scipy_wrapper.py。其专属选项solver对应scipy.integrate.solve_ivp的method参数。5.5 工程选型建议对于大多数问题好的选择是默认的dopri5或者使用rk4并搭配optionsdict(step_size...)将步长设到足够小。调整容差自适应求解器或步长固定求解器可以在求解速度与精度之间做权衡# 自适应求解器收紧容差换取精度 odeint(func, y0, t, methoddopri5, rtol1e-9, atol1e-12) # 固定步长求解器必须显式指定 step_size odeint(func, y0, t, methodrk4, optionsdict(step_size0.05))注意固定步长求解器若不传step_size默认把步长设为t中相邻取值点的间隔即网格构造器直接返回t见 solvers.py。如果t只用于指定积分区间的起止那么显式传入step_size至关重要否则只会得到一次大步长积分。六、求解器专属选项options 字典详解以下选项均通过optionsdict(...)传入默认值以 FURTHER_DOCUMENTATION.md 为准。6.1 自适应求解器通用选项dopri8、dopri5、bosh3、adaptive_heunfirst_stepNone求解器第一步的步长默认通过经验算法自动选取。源码中的_select_initial_stepmisc.py根据初始导数与误差尺度估计初始步长算法出自 Hairer 等《Solving Ordinary Differential Equations I: Nonstiff Problems》Sec. II.4safety0.9、ifactor10.0、dfactor0.2控制下一步最优步长的计算方式。粗略地说safety会把步长略微收缩该比例ifactor是步长最多可增长的上限倍率dfactor是最多可收缩的下限倍率max_num_steps2**31 - 1求解器允许的最大步数上限dtypetorch.float64时间类量使用的数据类型。设为torch.float32可提升速度但更容易产生下溢错误step_tNone必须踩到的时刻序列应为torch.Tensor。当func在这些时刻存在导数间断kink时尤其有用——求解器不必再缓慢地自行发现这些间断点jump_tNone必须踩到并重新求值func的时刻序列。当func在这些时刻存在不连续时前一步的最后一次函数求值不等于下一步的首次求值即 FSAL 性质在该点不成立norm计算接受/拒绝判据所用的范数。对张量输入默认使用 RMS 范数对元组输入默认对每个张量计算 RMS 后取最大值得到混合 L-infinity/RMS 范数。若传入自定义范数其签名为接收与y0同形状的张量/元组返回标量。当作为adjoint_options的一部分传入时可使用特殊值seminorm以剔除参数项对范数的贡献源自 Hey, thats not an ODE 论文。6.2 固定步长求解器通用选项euler、midpoint、rk4、explicit_adams、implicit_adamsstep_sizeNone每个离散步的大小。未指定时默认在t的取值点之间步进见 5.5 节警告grid_constructorNone更细粒度地设置步进位置。应为可调用对象func, y0, t - grid把odeint的func, y0, t参数变换成期望的网格1-D 张量。它与step_size互斥——同时传入会抛出ValueErrorsolvers.pyperturbFalse若为True自动在每个步的起止处加上微小扰动通过nextafter取相邻可表示浮点数见 misc.py 的_PerturbFunc使步进可以精确到达不连续点。6.3 各求解器独有选项explicit_adams显式 Adams-Bashforthmax_orderAdams-Bashforth 预测器的最大阶数注意此求解器忽略rtol和atol。implicit_adams隐式 Adams-Bashforth-Moultonmax_order预测-校正器的最大阶数max_itersAdams-Moulton 校正器的最大迭代次数注意此求解器的rtol/atol对应校正器收敛判据。scipy_solversolver使用的 SciPy 求解器名称对应scipy.integrate.solve_ivp的method参数。七、伴随方法专属选项odeint_adjoint额外支持以下参数FURTHER_DOCUMENTATION.md 与 adjoint.pyadjoint_rtol、adjoint_atol、adjoint_method、adjoint_options反向传播伴随 ODE 求解使用的容差、方法与选项默认继承前向的值。若前向method与adjoint_method不同则必须显式给出adjoint_options否则会抛出ValueErroradjoint_options支持特殊键值对{norm: seminorm}对自适应步长求解器可提供更高效的伴随求解Hey, thats not an ODE 论文方法。其实现为自定义伴随范数adjoint_seminorm——只取max(|t|, norm(y), norm(adj_y))而忽略参数项见 adjoint.pyadjoint_params反向传播中需要计算梯度的参数元组默认取tuple(func.parameters())。注意会自动过滤掉requires_gradFalse的参数若同时使用自定义 norm 会给出警告。八、求解过程回调Callbacks回调以func的方法形式定义在求解过程中被触发机制见 misc.py回调在_check_inputs中被挂接到包装后的func上。目前支持三种均有(t0, y0, dt)签名callback_step(self, t0, y0, dt)在步长dt的步进之前、时刻t0、当前解y0处调用。所有求解器除scipy_solver均支持callback_accept_step(self, t0, y0, dt)在接受一个步长为dt的步时调用仅自适应求解器dopri8、dopri5、bosh3、adaptive_heun支持callback_reject_step(self, t0, y0, dt)同callback_accept_step但在拒绝步长时调用。在伴随反向传播阶段可给回调名加上_adjoint后缀启用对应回调如callback_step_adjoint。从 solvers.py 可看到自适应求解器的valid_callbacks默认返回空集、固定步长求解器返回{callback_step}传入不支持的回调会触发警告而非错误。九、事件处理odeint_event9.1 什么是事件处理事件处理允许基于事件函数提前终止 ODE 求解并且对大多数求解器支持反向传播。典型应用是碰撞检测弹跳球或学习何时发生事件的可微模型。相关论文见参考文献 [2]。9.2 调用签名from torchdiffeq import odeint_event odeint_event(func, y0, t0, *, event_fn, reverse_timeFalse, odeint_interfaceodeint, **kwargs)参数说明摘自 README 与 odeint.pyfunc、y0与odeint相同t0标量表示初始时刻event_fn(t, y)必填关键字参数返回一个张量reverse_time布尔值是否反向求解默认Falseodeint_interfaceodeint或odeint_adjoint之一指定用哪种方式对 ODE 解求微分默认odeint**kwargs其余关键字参数透传给odeint_interface如rtol、atol、method、options。9.3 终止条件与返回值当event_fn(t, y)的任一元素等于零时求解终止于事件时刻t与状态y。event_fn可以返回多个输出以指定多个事件函数最先触发者终止求解。odeint_event同时返回事件时间与终态两者都可被求导梯度会穿过事件函数回传。注意要获得事件函数参数的梯度这些参数必须位于状态state本身之中。事件时间的数值精度由atol参数决定。从源码看事件求解通过find_event在步长内用插值线性或三次 Hermite见 solvers.py定位符号变化的零点并以atol作为定位精度固定步长求解器的事件处理要求必须在options中提供step_sizesolvers.py且最多迭代 20000 次超出抛RuntimeError。9.4 事件梯度的自动衔接odeint_event的梯度通过自定义 autograd FunctionImplicitFnGradientReroutingodeint.py实现隐式函数梯度重路由反向时对事件函数求向量-雅可比积用dcdt par_dt sum(dstate * f_val)计算事件函数对时间的全导数再据此把对事件时间的梯度折算进对状态的梯度中。这使得事件时间对参数可微这一需求被优雅地封装在库内部。9.5 示例弹跳球仿真仓库 examples/bouncing_ball.py 演示了用事件处理模拟并微分一个弹跳球默认 10 次弹跳支持--adjoint开关、含gradcheck数值梯度校验。核心思路状态为(pos, vel, log_radius)元组forward定义自由落体动力学dpos vel, dvel -gravityevent_fn返回pos - exp(log_radius)——球在空中的时候为正、陷入地面时为负每次odeint_event求得碰撞时刻后用state_update反转速度并乘上吸收系数(1 - absorption)把t0更新为event_t继续下一次求解event_t, solution odeint_event( self, state, t0, event_fnself.event_fn, reverse_timeFalse, atol1e-8, rtol1e-8, odeint_interfaceself.odeint, )其中state_update给位置加上1e-7微小偏移以避免立即再次触发事件函数。十、示例程序导览官方示例全部位于 examples 目录examples/README.md 有汇总说明。对新手最重要的是 examples/ode_demo.py——它展示了如何用torchdiffeq拟合一条螺旋 ODE 轨迹该示例的核心要素生成真值数据用Lambday**3乘一个常数矩阵true_A与methoddopri5前向生成螺旋轨迹定义可学习的 ODE 网络ODEFunc是一个两层Linear(2, 50) - Tanh - Linear(50, 2)的nn.Moduleforward输出net(y**3)命令行参数--methoddopri5/adams、--data_size默认 1000、--batch_time默认 10、--batch_size默认 20、--niters默认 2000、--test_freq默认 20、--viz可视化、--gpu、--adjoint切换为odeint_adjointpython examples/ode_demo.py --adjoint --viz训练循环随机采样时间窗批get_batchodeint(func, batch_y0, batch_t)前向mean(|pred - true|)作为损失RMSprop 优化。其余示例及主题examples/bouncing_ball.py事件驱动的弹跳球仿真与梯度校验上文 9.5 节examples/learn_physics.py学习简单事件函数的示例如学习跳台/障碍触发examples/latent_ode.py潜在变量 ODELatent ODE将高维时序映射到潜在空间后再做 ODE 演化examples/cnf.py连续归一化流Continuous Normalizing Flows利用 ODE 做概率密度变换examples/odenet_mnist.py在 MNIST 上训练 ODE-Net把残差网络块替换为 ODE 层。十一、常见问题FAQ 速查以下要点提炼自 FAQ.md帮助排查实际使用中的高频问题1. NFE-F / NFE-B 是什么前向与反向传播中的函数求值次数Number of Function Evaluations。自适应求解器每次步进会做多次函数求值例如dopri5至少保存 6 次 ODE 求值再用它们的线性组合完成一步其中前两次用于选取初始步长。2. rtol / atol 的作用相对与绝对误差容限。自适应求解器每步产生误差估计若误差超过容限则缩小步长重算直到误差小于容限为止。误差容限的计算式为atol rtol * 当前状态范数范数为混合 L-infinity/RMS 范数misc.py。3. 如何获得自适应求解器在估计路径上的取值odeint的参数t指定输出时刻点例如odeint(func, x0, ttorch.linspace(0, 1, 50))。注意求解器总是从t的最小值积分到最大值中间的t取值不影响求解本身只是用多项式插值在这些时刻点取值成本极低。4. Neural ODE 中应该用什么非线性避免 ReLU、LeakyReLU 这类非光滑非线性优先使用理论上具有唯一伴随/梯度的 Softplus 等光滑函数。5. 训练时最耗内存的操作是什么使用伴随方法时最耗内存的是反向传播中对网络的单次backward调用adjoint.py 附近的torch.autograd.grad计算。6. 数值解比初始值离目标更远残差网络的初始化技巧如把最后一层权重置零对 ODE 同样有效——这会把 ODE 初始化为恒等映射。7. 训练太慢大概率是在 CPU 上运行。训练需要反复对网络求值CPU 上极度缓慢是预期行为请使用 GPU。8. 自适应求解器如 dopri5出现 dt 下溢这是 ODE 变刚硬stiff的信号——某个区域动力学过于剧烈步长被压到接近零而无法推进。缓解办法使用权重衰减等正则化、选用温和的激活函数、放宽atol/rtol接受更大误差或改用固定步长求解器。9. 关于数值方法的深入资料库采用的 RK 系数表Butcher tableau源自 Dormand Prince 等经典文献自适步长控制算法参考 Hairer, Norsett, Wanner《Solving Ordinary Differential Equations I: Nonstiff Problems》Sec. II.4可进一步阅读该专著理解步长控制细节。十二、进一步阅读与参考文献FURTHER_DOCUMENTATION.md求解器选项与伴随选项的完整默认值清单、回调说明FAQ.md常见问题与求解器机制详解示例目录 examples从螺旋拟合到事件学习、Latent ODE、CNF、ODE-Net 的完整可运行代码测试目录 testsodeint_tests.py、gradient_tests.py、event_tests.py、api_tests.py等覆盖了本文所述接口的数值正确性与梯度校验。本文内容对应的两篇核心论文应用场景分别对应 ODE 求解器与事件处理Ricky T. Q. Chen, Yulia Rubanova, Jesse Bettencourt, David Duvenaud. Neural Ordinary Differential Equations.Advances in Neural Information Processing Systems.2018.arXiv:1806.07366Ricky T. Q. Chen, Brandon Amos, Maximilian Nickel. Learning Neural Event Functions for Ordinary Differential Equations.International Conference on Learning Representations.2021.arXiv:2011.03902Patrick Kidger, Ricky T. Q. Chen, Terry Lyons. Hey, thats not an ODE: Faster ODE Adjoints via Seminorms.International Conference on Machine Learning.2021.arXiv:2009.09457seminorm选项的理论依据若你的研究使用了本库官方建议引用misc{torchdiffeq, author{Chen, Ricky T. Q.}, title{torchdiffeq}, year{2018}, url{https://github.com/rtqichen/torchdiffeq}, }赞分享深度学习科学计算【免费下载链接】torchdiffeqDifferentiable ODE solvers with full GPU support and O(1)-memory backpropagation.项目地址https://gitcode.com/gh_mirrors/to/torchdiffeq点击查看免费下载相关推荐终极指南使用torchdiffeq掌握PyTorch可微ODE求解技术终极指南使用torchdiffeq掌握PyTorch可微ODE求解技术 torchdiffeq是PyTorch生态系统中革命性的常微分方程 ODE 求解器库深度学习科学计算终极实战指南如何用torchdiffeq构建可微分ODE求解应用终极实战指南如何用torchdiffeq构建可微分ODE求解应用 欢迎来到torchdiffeq的实战指南 这是一个基于PyTorch的常微分方程OD深度学习科学计算如何使用PyTorch微分方程求解器torchdiffeq深度学习中的ODE完整解决方案如何使用PyTorch微分方程求解器torchdiffeq深度学习中的ODE完整解决方案 torchdiffeq是一个基于PyTorch的微分方程求解器库提深度学习科学计算上一篇TradeMaster多智能体协作构建多策略组合的量化交易系统下一篇如何快速掌握Mongoku面向开发者的终极MongoDB Web管理工具指南创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
锦
锦皓数字建站
深耕本土企业品牌数字化升级,专注原创端正雅致商务官网,从视觉设计到稳定运维全程保驾护航。