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

JAX实现可扩展自治三对角系统雅可比的高效方案问询

如何使用JAX最高效实现可扩展的自治三对角系统?

以下是问题对应的基础测试代码:

import functools as ft
import jax as jx
import jax.numpy as jnp
import jax.random as jrn
import jax.lax as jlx

def make_T(m):
    # 生成伪随机三对角雅可比矩阵,以带状格式存储
    T = jnp.zeros((3,m), dtype='f8')
    T = T.at[0, 1:  ].set(jrn.normal(jrn.PRNGKey(0), shape=(m-1,)))
    T = T.at[1,  :  ].set(jrn.normal(jrn.PRNGKey(1), shape=(m  ,)))
    T = T.at[2,  :-1].set(jrn.normal(jrn.PRNGKey(2), shape=(m-1,)))
    return T

def make_y(m):
    # 生成伪随机状态数组
    y = jrn.normal(jrn.PRNGKey(3), shape=(m  ,))
    return y

def calc_f_base(y, T):
    # 根据当前状态计算变化率
    f = T[1,:]*y
    f = f.at[ 1:  ].set(f[ 1:  ]+T[0, 1:  ]*y[  :-1])
    f = f.at[  :-1].set(f[  :-1]+T[2,  :-1]*y[ 1:  ])
    return f

m = 2**22 # 该规模下常规雅可比计算方法可能耗尽硬件资源
T = make_T(m)
y = make_y(m)

calc_f = ft.partial(calc_f_base, T=T)

直接调用jax.jacrev或jax.jacfwd会生成完整的稠密雅可比矩阵,会严重限制系统可支持的最大规模。

现有突破规模限制的尝试实现

以下是基于前向自动微分、仅提取三对角带状结构的实现,避免生成稠密矩阵:

@ft.partial(jx.jit, static_argnums=(0,))
def calc_jacfwd_trid(calc_f, y):
    # 前向模式计算雅可比的三对角带

    def scan_body(carry, i):
        t, T = carry
        t = t.at[i  ].set(1.0)
        
        f, dfy = jx.jvp(calc_f, (y,), (t,))
        
        T = T.at[2,i-1].set(dfy[i-1])
        T = T.at[1,i  ].set(dfy[i  ])
        T = T.at[0,i+1].set(dfy[i+1])

        t = t.at[i-1].set(0.0)

        return (t, T), None

    # 初始化存储
    m = y.size
    t = jnp.zeros_like(y)
    T = jnp.zeros((3,m), dtype=y.dtype)

    # 对y[0]求导
    t = t.at[0].set(1.0)
    f, dfy = jx.jvp(calc_f, (y,), (t,))
    idxs = jnp.array([1,0]), jnp.array([0,1])
    T = T.at[idxs].set(dfy[0:2])

    # 对中间节点y[1:-1]批量求导
    (t, T), empty = jlx.scan(scan_body, (t,T), jnp.arange(1,m-1))

    # 对最后一个节点y[-1]求导
    t = t.at[m-2:].set(jnp.array([0.0,1.0]))
    f, dfy = jx.jvp(calc_f, (y,), (t,))
    idxs = jnp.array([2,1]), jnp.array([m-2,m-1])
    T = T.at[idxs].set(dfy[-2:])

    return T

该实现可支撑如下大规模三对角线性系统求解流程:

T = jacfwd_trid(calc_f, y)

df = jrn.normal(jrn.PRNGKey(4), shape=y.shape)
dx = jlx.linalg.tridiagonal_solve(*T,df[:,None]).flatten()

当前待解决的问题:

  • 是否存在更优的实现方案?
  • 是否可以进一步降低calc_jacfwd_trid的时间复杂度?

补充说明

以下实现写法更紧凑,但实际运行耗时略高于上述scan版本:

@ft.partial(jx.jit, static_argnums=(0,))
def calc_jacfwd_trid_map(calc_f, y):
    # 基于lax.map的前向模式三对角雅可比计算

    def map_body(i, t):

        t = t.at[i-1].set(0.0)
        
        f, dfy = jx.jvp(calc_f, (y,), (t,))

        im1 = jnp.where(i > 0, i-1, 0)
        Ti = jlx.dynamic_slice(dfy, (im1,), (3,))
        Ti = jnp.where(i >   0, Ti, jnp.roll(Ti, shift=+1))
        Ti = jnp.where(i < m-1, Ti, jnp.roll(Ti, shift=-1))

        t = t.at[i  ].set(1.0)

        return Ti

    # 初始化
    m = y.size
    t = jnp.zeros_like(y)

    # 对所有状态维度求导
    T = jlx.map(lambda i : map_body(i, t=t), jnp.arange(m))

    # 修正带状矩阵的存储顺序对齐接口要求
    T = T.transpose()
    T = jnp.flip(T, axis=0)
    T = T.at[0,:].set(jnp.roll(T[0,:], shift=+1))
    T = T.at[2,:].set(jnp.roll(T[2,:], shift=-1))

    return T

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 14:54:19