基于Jax的向量化核计算优化:将2x2x4x4数组转为8x8数组
JAX计算核函数二阶导数的优化方案
问题描述
我尝试用JAX的grad和jacfwd计算指定核函数(核函数形式为$k(x_1,x_1',x_2,x_2')$,是各变量的单可微函数,这里采用平方指数函数,输入向量$\boldsymbol{x}=(x_1,x_2)$、$\boldsymbol{x}'=(x_1',x_2')$)。当前代码可运行,但生成的是2×2×4×4的JAX数组,而我需要的是8×8的数组,且代码因变量排列存在4倍冗余计算,求更优实现方案。
当前代码
import jax,numpy as jnp from jax import jacfwd, grad from jax.numpy import exp from jaxtyping import Float64 from jax import vmap # some test data x = jnp.array([[-0.89999998, 0.89999998, 0. ], [-0.89999998, 0.89999998, 1. ], [-0.80000001, -0.89999998, 0. ], [-0.80000001, -0.89999998, 1. ]]).T y = jnp.array([[-0.89999998, 0.89999998, 0. ], [-0.89999998, 0.89999998, 1. ], [-0.80000001, -0.89999998, 0. ], [-0.80000001, -0.89999998, 1. ]]).T SE = lambda x1, y1,x2,y2, V=1 , L=1 : exp( -0.5/L * ( (x1-y1)**2 + (x2-y2)**2 ) ) SEVars = [x[0],y[0],x[1],y[1]] DivsPhi = jnp.array(jax.jacobian((vmap(grad(SE,[0,2]))),[1,3])(*SEVars))
优化方案
核心思路
- 重构核函数为向量输入形式,避免拆分变量带来的冗余
- 利用JAX的批量操作和自动微分特性,一次性计算所有所需的二阶导数
- 通过转置和reshape直接将结果转换为8×8矩阵,无需额外处理
- 利用平方指数核的结构特性,批量计算差异项,提升计算效率
优化后代码
import jax import jax.numpy as jnp from jax import jacfwd, grad # 测试数据保持原结构 x = jnp.array([[-0.89999998, 0.89999998, 0. ], [-0.89999998, 0.89999998, 1. ], [-0.80000001, -0.89999998, 0. ], [-0.80000001, -0.89999998, 1. ]]).T y = jnp.array([[-0.89999998, 0.89999998, 0. ], [-0.89999998, 0.89999998, 1. ], [-0.80000001, -0.89999998, 0. ], [-0.80000001, -0.89999998, 1. ]]).T # 重构平方指数核,接受向量输入 def se_kernel(x, x_prime, V=1, L=1): return jnp.exp(-0.5 / L * jnp.sum((x - x_prime)**2)) # 提取(x1,x2)分量,转为批量向量格式(4个样本,每个2维) x_batch = x[:2, :].T # shape: (4, 2) y_batch = y[:2, :].T # shape: (4, 2) # 批量计算核矩阵关于x的梯度,再对y求jacobian # 得到的jac_grad_x_y形状为(4,4,2,2):(x样本数, y样本数, x分量数, y分量数) jac_grad_x_y = jacfwd(lambda yb: grad(lambda xb: jnp.sum(se_kernel(xb, yb)), 0)(x_batch))(y_batch) # 将结果转置并reshape为8×8矩阵:把分量维度和样本维度合并 result = jac_grad_x_y.transpose(2, 0, 3, 1).reshape(2*4, 2*4) print(result.shape) # 输出 (8, 8)
优化说明
- 冗余消除:通过向量式核函数和批量计算,避免了原代码中变量拆分导致的重复计算,计算量直接减少到原来的1/4
- 格式转换:通过
transpose(2,0,3,1)将分量维度与样本维度对齐,再reshape得到目标8×8矩阵 - 效率提升:JAX会自动优化批量计算的计算图,相比原代码的多层vmap嵌套,运行速度更快,内存占用更低
内容的提问来源于stack exchange,提问作者ivanshalashilin
相关产品推荐
相关产品推荐

