计算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提供了自动微分工具,可以直接计算输出对输入的梯度。针对单个输出像素的梯度计算(内存效率更高),可以按以下步骤实现:
完整代码示例
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

