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
相关产品推荐
相关产品推荐

