如何使用Google JAX实现对标scipy.curve_fit的一阶ODE曲线拟合
现有代码问题分析
- 优化器性能过差:你当前用的是固定小学习率(0.0001)的朴素SGD优化器,属于一阶优化方法中收敛最慢的一类,而
scipy.curve_fit默认使用Levenberg-Marquardt、信赖域这类拟牛顿二阶优化方法,收敛速度和精度远高于手写SGD。 - 未启用JIT编译:所有计算函数都没有做JAX的即时编译,JAX的算力优势完全没有发挥,运行效率远低于scipy预编译的C后端实现。
- ODE求解器精度不足:手写的欧拉法是一阶精度的固定步长求解器,数值误差大,还可能存在稳定性问题,误差会传导到损失和梯度计算中,进一步降低最终拟合精度。
- 训练循环开销大:迭代更新用的是Python原生for循环,存在大量Python层开销,迭代次数多的时候会进一步拖慢运行速度。
优化方案
- 启用JIT编译:给损失函数、梯度计算函数加
@jax.jit装饰器,将计算逻辑编译为静态计算图执行,运行速度可提升1~2个数量级。 - 替换高性能优化器:放弃手写SGD,使用
optax库提供的Adam、LBFGS等优化器,其中LBFGS属于拟牛顿法,适配小参数规模的曲线拟合场景,收敛速度和精度和scipy.curve_fit接近。 - 替换高精度ODE求解器:放弃手写欧拉法,使用
diffrax库提供的自适应步长高阶ODE求解器,数值精度高、稳定性好,梯度计算也更准确。 - 优化迭代循环:用
jax.lax.scan实现迭代逻辑,将整个训练过程编译进计算图,消除Python for循环的开销。
更优的JAX曲线拟合方案
直接使用jaxopt库的CurveFit接口,该接口完全对标scipy.optimize.curve_fit实现,原生支持JAX的自动微分、JIT编译特性,不需要手动写损失函数和优化循环,接口用法和scipy几乎一致,迁移成本极低,性能和精度都有官方优化。
如果是ODE类的拟合任务,只需要将ODE求解逻辑封装为待拟合函数,直接传入jaxopt.CurveFit即可调用。
优化后参考代码
import jax.numpy as jnp import jax import numpy as np import jaxopt import diffrax as dfx # 生成模拟数据 u = np.zeros(100) u[10:50] = 1 t = jnp.arange(len(u)) u = jnp.array(u) # 一阶ODE定义 def f(t, y, args): k, tau = args return (k * u[t.astype(int)] - y) / tau # ODE求解封装 def ode_solve(k, tau): solver = dfx.Tsit5() # 自适应步长五阶龙格库塔求解器,精度远高于欧拉法 t0 = t[0] t1 = t[-1] dt0 = 1.0 y0 = jnp.array([0.0]) saveat = dfx.SaveAt(ts=t) sol = dfx.diffeqsolve( dfx.ODETerm(f), solver, t0, t1, dt0, y0, args=(k, tau), saveat=saveat ) return sol.ys.reshape(-1) # 带噪声的观测数据 true_k, true_tau = 2.0, 5.0 pred = ode_solve(true_k, true_tau) pred_noise = pred + 0.05 * np.random.randn(len(pred)) # JIT编译损失函数 @jax.jit def loss(params): k, tau = params pred = ode_solve(k, tau) return jnp.sum((pred - pred_noise) ** 2) # 用LBFGS优化器迭代,收敛速度远快于SGD solver = jaxopt.LBFGS(fun=loss, maxiter=100) init_params = jnp.array([1.0, 2.0]) res = solver.run(init_params) print("拟合参数:", res.params)
内容的提问来源于stack exchange,提问作者Siva-Sg
相关产品推荐
相关产品推荐

