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

计算图像雅可比矩阵:如何正确重塑JAX输出的NumPy数组?

批量图像变换的散度计算方案

核心问题拆解

  • 批量输入[k, width, height, channel_count]下,jax.jacfwd默认计算全批量的雅可比矩阵(包含跨样本的导数),但我们仅需每个样本自身的雅可比(第i个样本输出仅对第i个样本输入求导)
  • 需要同时保留变换f的输出值与散度结果
  • 需避免计算完整雅可比矩阵以提升计算效率

解决方案

方法1:修剪完整雅可比(适合调试理解)

先计算全量雅可比,再提取每个样本对应的对角块,最后求迹得到散度:

import jax
import jax.numpy as jnp

# 模拟不可修改的神经网络输出函数f
def f(x):
    return jnp.sin(x) + x ** 2

# 构造测试输入
k, W, H, C = 2, 3, 3, 2
x = jnp.ones((k, W, H, C))

# 计算完整雅可比矩阵
jac_full = jax.jacfwd(f)(x)  # 形状: [k, W, H, C, k, W, H, C]

# 提取单样本雅可比:取每个样本i对应的jac_full[i, ..., i, ...]
jac_single = jac_full[jnp.arange(k), :, :, :, jnp.arange(k), :, :, :]
# 此时jac_single形状为[k, W, H, C, W, H, C]

# 计算散度:对每个样本的雅可比矩阵求迹
divergence = jnp.einsum('nwhcwhc->n', jac_single)

# 获取f的输出
f_output = f(x)

print("f输出形状:", f_output.shape)
print("散度形状:", divergence.shape)

方法2:单样本批量映射(高效避免跨样本计算)

用jax.vmap批量处理单样本的散度计算,跳过跨样本导数的计算:

import jax
import jax.numpy as jnp

def f(x):
    return jnp.sin(x) + x ** 2

k, W, H, C = 2, 3, 3, 2
x = jnp.ones((k, W, H, C))

# 定义单样本的散度+输出计算函数
def single_sample_process(x_i):
    f_i = f(x_i[None, ...])[0]
    jac_i = jax.jacfwd(lambda x: f(x[None, ...])[0])(x_i)
    div_i = jnp.trace(jac_i.reshape(-1, -1))  # 展平后求迹
    return f_i, div_i

# 批量映射处理所有样本
f_output, divergence = jax.vmap(single_sample_process)(x)

print("f输出形状:", f_output.shape)
print("散度形状:", divergence.shape)

方法3:最优效率方案(利用迹的微分性质)

利用散度=雅可比迹=输入各维度导数之和的性质,直接通过梯度求和计算,完全避免显式雅可比:

import jax
import jax.numpy as jnp

def f(x):
    return jnp.sin(x) + x ** 2

k, W, H, C = 2, 3, 3, 2
x = jnp.ones((k, W, H, C))

# 批量计算散度并保留f输出
def batch_div_with_output(x):
    f_out = f(x)
    # 单样本散度计算:对输入所有元素的导数求和
    def div_single(x_i):
        return jnp.sum(jax.grad(lambda x: jnp.sum(f(x[None, ...])))(x_i))
    divergence = jax.vmap(div_single)(x)
    return f_out, divergence

f_output, divergence = batch_div_with_output(x)

print("f输出形状:", f_output.shape)
print("散度形状:", divergence.shape)

方案对比

  • 方法1直观但效率低,仅适合小批量调试
  • 方法2跳过跨样本导数计算,内存占用大幅降低
  • 方法3利用微分性质直接计算,是计算散度的最优方案,内存与计算效率最高

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 03:12:43