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

基于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))

优化方案

核心思路

  1. 重构核函数为向量输入形式,避免拆分变量带来的冗余
  2. 利用JAX的批量操作和自动微分特性,一次性计算所有所需的二阶导数
  3. 通过转置和reshape直接将结果转换为8×8矩阵,无需额外处理
  4. 利用平方指数核的结构特性,批量计算差异项,提升计算效率

优化后代码

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.15 06:25:55