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

计算JAX卷积的雅可比矩阵:输入像素对输出的梯度提取

卷积关于输入像素的导数(数学定义)

假设输入图像为 ( I ),输出图像为 ( O ),卷积核为 ( K )(尺寸为 ( 2k+1 \times 2k+1 )),且使用same模式的卷积操作。输出像素 ( O(x_{\text{out}}, y_{\text{out}}) ) 的计算公式为:
[
O(x_{\text{out}}, y_{\text{out}}) = \sum_{i=-k}^{k} \sum_{j=-k}^{k} I(x_{\text{out}}+i, y_{\text{out}}+j) \cdot K(i, j)
]

对于输入像素 ( I(x_{\text{in}}, y_{\text{in}}) ),其对输出像素 ( O(x_{\text{out}}, y_{\text{out}}) ) 的导数形式为:
[
\frac{\partial O(x_{\text{out}}, y_{\text{out}})}{\partial I(x_{\text{in}}, y_{\text{in}})} =
\begin{cases}
K(x_{\text{in}} - x_{\text{out}}, y_{\text{in}} - y_{\text{out}}) & \text{若 } |x_{\text{in}} - x_{\text{out}}| \leq k \text{ 且 } |y_{\text{in}} - y_{\text{out}}| \leq k \
0 & \text{其他情况}
\end{cases}
]

简单来说:只有当输入像素位于输出像素对应的卷积核覆盖窗口内时,导数等于卷积核中对应偏移位置的数值;不在窗口内的输入像素对该输出像素的导数为0。

用JAX提取输出像素相对于输入像素的梯度

JAX提供了自动微分工具,可以直接计算输出对输入的梯度。针对单个输出像素的梯度计算(内存效率更高),可以按以下步骤实现:

完整代码示例

import jax
import jax.numpy as jnp
from jax.scipy.signal import convolve2d

def gaussian_kernel(size: int, std: float):
    """生成2D高斯卷积核"""
    x, y = jnp.mgrid[-size:size+1, -size:size+1]
    g = jnp.exp(-(x**2 + y**2) / (2 * std**2))
    return g / g.sum()
    
def gaussian_blur(image, kernel_size=5, sigma=1.0):
    """对2D图像应用高斯模糊"""
    kernel = gaussian_kernel(kernel_size, sigma)
    blurred_image = convolve2d(image, kernel, mode='same')
    return blurred_image 

# 定义函数:输入图像,返回指定位置的输出像素值
def get_single_output_pixel(image, output_x, output_y):
    blurred = gaussian_blur(image)
    return blurred[output_x, output_y]

# 示例:计算输出(10,10)位置像素对输入的梯度
input_img = jnp.random.rand(20, 20)  # 20x20的示例输入图像
# 生成梯度计算函数
grad_func = jax.grad(get_single_output_pixel, argnums=0)
# 计算梯度:结果是与输入同尺寸的数组,每个元素对应输入像素对输出(10,10)的导数
pixel_gradient = grad_func(input_img, output_x=10, output_y=10)

关键说明

  • jax.grad会生成一个函数,用于计算目标函数(这里是get_single_output_pixel)对指定输入参数(argnums=0表示第一个参数,即输入图像)的梯度。
  • 得到的pixel_gradient数组中,非零区域恰好是输出(10,10)对应的输入卷积窗口,每个非零值等于高斯核中对应位置的数值,完全符合前面的数学定义。
  • 如果需要计算整个输出图像对输入的梯度(所有输出像素对所有输入像素的导数),可以使用jax.jacrev,但注意:对于尺寸为 ( H \times W ) 的图像,雅可比矩阵的尺寸为 ( H \times W \times H \times W ),大图像会占用大量内存,因此仅推荐小图像使用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 21:08:16