You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用JAX实现牛顿法求解优化问题时内核崩溃求助

牛顿法求解优化问题时Jupyter内核崩溃排查

问题背景

我有一个优化问题,尝试用牛顿法求解,用jax.jacobian计算雅可比矩阵。目标函数名为calc_ms,实现的牛顿法找根函数如下:

def newton(f, x_0, tol=1e-5, max_iter=10):
    x = x_0
    q = lambda x: x - jnp.linalg.solve(jax.jacobian(f)(x), f(x))
    error = tol + 1
    n = 0
    while error > tol:
        n+=1
        if(n > max_iter):
            raise Exception('Max iteration reached without convergence')
        y = q(x)
        if(any(jnp.isnan(y))):
            raise Exception('Solution not found with NaN generated')
        error = jnp.linalg.norm(x - y)
        x = y
        print(f'iteration {n}, error = {error:.5f}')
    print('\n' + f'Result = {x} \n')
    return x

调用代码:

newton(lambda delta: calc_ms(delta, dist_util, food_exp, market_shares, corr_mat, n_stores), delta0).block_until_ready()

运行后内核崩溃,报错:

Canceled future for execute_request message before replies were done
The Kernel crashed while executing code in the the current cell or a previous cell. Please review the code in the cell(s) to identify a possible cause of the failure. Click here for more info. View Jupyter log for further details.

可能的崩溃原因分析

1. 雅可比矩阵维度/计算量过大导致内存溢出

如果delta是高维向量(比如维度上千),jax.jacobian(f)(x)会生成(N, N)形状的矩阵,当N很大时,这个矩阵会占用极大内存。后续jnp.linalg.solve对大矩阵的操作也会消耗大量内存和计算资源,直接导致内核因内存不足崩溃。

2. JAX即时编译与原生循环的冲突

牛顿法使用Python原生while循环,而JAX核心操作依赖即时编译(JIT)。每次迭代调用jax.jacobian(f)(x)都会触发一次编译,多次迭代会累积大量编译开销,甚至引发资源耗尽。另外,lambda delta: calc_ms(...)这类匿名函数在JAX编译时可能存在追踪问题,导致编译过程出错。

3. 数值不稳定与类型操作错误

  • 若雅可比矩阵接近奇异,jnp.linalg.solve会产生极大数值甚至inf/nan,如果计算雅可比的过程中出现数值溢出,可能直接导致内核崩溃,无法执行后续的nan检查。
  • 代码中使用Python内置any检查JAX数组的nan,跨类型操作可能触发未定义行为,引发崩溃,应改用JAX原生的jnp.any。

4. Jupyter内核资源限制

Jupyter默认内核的内存、CPU资源有限,若计算任务超出这些限制,内核会被系统强制终止,出现上述报错。

排查与修复方案

  • 内存优化:查看delta0的维度,若维度超过几百,改用拟牛顿法(如BFGS)或稀疏雅可比矩阵计算(若雅可比为稀疏矩阵),避免生成全量大矩阵。
  • 循环JIT优化:用JAX的jax.lax.while_loop替代Python原生while循环,同时对整个牛顿法逻辑做JIT编译,减少编译开销。示例框架:
    from jax import jit, lax
    
    def newton_step(x, f):
        jac = jax.jacobian(f)(x)
        return x - jnp.linalg.solve(jac, f(x))
    
    @jit
    def jit_newton(f, x0, tol=1e-5, max_iter=10):
        def cond_fun(carry):
            x, error, n = carry
            return (error > tol) & (n < max_iter)
    
        def body_fun(carry):
            x, _, n = carry
            y = newton_step(x, f)
            error = jnp.linalg.norm(x - y)
            return (y, error, n + 1)
    
        init_carry = (x0, tol + 1, 0)
        final_x, final_error, final_n = lax.while_loop(cond_fun, body_fun, init_carry)
        return final_x
    
  • 数值稳定性修正:
    1. 替换any(jnp.isnan(y))为jnp.any(jnp.isnan(y))。
    2. 在jnp.linalg.solve前检查雅可比矩阵的条件数,或用jnp.linalg.lstsq替代,增强鲁棒性。
  • 资源监控:运行时用系统工具(如top、htop)监控内存和CPU占用,确认是否因资源耗尽导致崩溃。

内容的提问来源于stack exchange,提问作者Alina Ozhegova

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.28 13:17:38