协方差(相机空间)梯度计算

还记得我们之前介绍的 3D 协方差到 2D 协方差的公式吗?

\Sigma^{\prime}=JR_{cw}\Sigma R_{cw}^{\top}J^{\top}

我们要优化的最终目标是 \Sigma,但屏幕渲染用的是 \Sigma^{\prime}。在这篇文章中,我们先计算 \Sigma^{\prime} 的梯度。

我们的最终目标是求损失函数 \mathcal{L} 对二维协方差矩阵 \Sigma'_n 的梯度,也就是 \frac{\partial \mathcal{L}}{\partial \Sigma'_n}

根据链式法则,我们需要一个“中间桥梁”。在前面计算像素颜色时,我们接触过一个叫 \sigma_n 的标量参数。在文档中,\sigma_n 的定义是高斯分布指数部分的核心项:

\sigma_n = \frac{1}{2}\Delta_n^\top \Sigma'^{-1} \Delta_n

这里的 \Delta_n \in \mathbb{R}^2 是像素中心与二维高斯中心 \mu' 的坐标偏移量 。

\frac{\partial \sigma_n}{\partial \Sigma'_n} = -\frac{1}{2} \Sigma_n^{\prime -1} \Delta_n \Delta_n^\top \Sigma_n^{\prime -1}

这个推导有点复杂,先直接拿来用。留作问题。

根据链式法则,损失 \mathcal{L} 必须先经过 \sigma_n,才能传递到 \Sigma'_n

\frac{\partial \mathcal{L}}{\partial \Sigma'_n} = \frac{\partial \mathcal{L}}{\partial \sigma_n} \cdot \frac{\partial \sigma_n}{\partial \Sigma'_n}

等式右边的第二项 \frac{\partial \sigma_n}{\partial \Sigma'_n} 我们已经有了。所以,现在的唯一任务就是求出等式右边的第一项:损失函数对标量 \sigma_n 的梯度 \frac{\partial \mathcal{L}}{\partial \sigma_n}

\sigma_n 是如何影响最终画面的?在渲染时,\alpha_n 是由它的基础不透明度 o_n 和空间衰减项 \sigma_n 共同决定的。

\alpha_n = o_n \cdot \exp(-\sigma_n)

\alpha_n\sigma_n 的导数非常直接:

\frac{\partial \alpha_n}{\partial \sigma_n} = -o_n \cdot \exp(-\sigma_n)

o_n \cdot \exp(-\sigma_n) 其实就是 \alpha_n 本身。所以我们可以等价地写成:

\frac{\partial \alpha_n}{\partial \sigma_n} = -\alpha_n

反向传播可以只进行到这一步,因为在之前的文章中,我们已经得到了损失对 \alpha_n 的梯度 \frac{\partial \mathcal{L}}{\partial \alpha_n}。利用链式法则,我们就能得到损失对 \sigma_n 的梯度:

\frac{\partial \mathcal{L}}{\partial \sigma_n} = \frac{\partial \mathcal{L}}{\partial \alpha_n} \cdot \frac{\partial \alpha_n}{\partial \sigma_n} = -\alpha_n \frac{\partial \mathcal{L}}{\partial \alpha_n}

所以,最终公式为:

\frac{\partial \mathcal{L}}{\partial \Sigma'_n} = \left( -\alpha_n \frac{\partial \mathcal{L}}{\partial \alpha_n} \right) \cdot \left( -\frac{1}{2} \Sigma_n^{\prime -1} \Delta_n \Delta_n^\top \Sigma_n^{\prime -1} \right)


代码实现

因为协方差矩阵比较小,写成矩阵形式没有手动展开效率高。下面我们对其进行展开。

我们之前已经计算得到逆协方差矩阵,是一个 2x2 的对称矩阵:

\Sigma_n^{\prime -1} = \begin{bmatrix} a_{11} & a_{12} \\ a_{12} & a_{22} \end{bmatrix}

像素到高斯中心的 2D 偏移向量为:

\Delta_n = \begin{bmatrix} du \\ dv \end{bmatrix}

根据矩阵乘法结合律,我们可以将公式 \Sigma_n^{\prime -1} \Delta_n \Delta_n^\top \Sigma_n^{\prime -1} 改写为两部分向量点积的矩阵外积:

\Sigma_n^{\prime -1} \Delta_n \Delta_n^\top \Sigma_n^{\prime -1} = (\Sigma_n^{\prime -1} \Delta_n) (\Sigma_n^{\prime -1} \Delta_n)^\top

计算中间向量 x = \Sigma_n^{\prime -1} \Delta_n

x = \begin{bmatrix} x_1 \\ x_2 \end{bmatrix} = \begin{bmatrix} a_{11} & a_{12} \\ a_{12} & a_{22} \end{bmatrix} \begin{bmatrix} du \\ dv \end{bmatrix} = \begin{bmatrix} a_{11} \cdot du + a_{12} \cdot dv \\ a_{12} \cdot du + a_{22} \cdot dv \end{bmatrix}

计算外积矩阵 x x^\top

x x^\top = \begin{bmatrix} x_1 \\ x_2 \end{bmatrix} \begin{bmatrix} x_1 & x_2 \end{bmatrix} = \begin{bmatrix} x_1^2 & x_1 x_2 \\ x_1 x_2 & x_2^2 \end{bmatrix}

代入完整的损失函数梯度公式:

\frac{\partial L}{\partial \Sigma'_n} = \frac{\partial \mathcal{L}}{\partial \alpha_n} \cdot \left( \frac{1}{2} \alpha \cdot x x^\top \right) = \left( 0.5 \cdot \alpha \cdot \frac{\partial \mathcal{L}}{\partial \alpha_n} \right) \cdot \begin{bmatrix} x_1^2 & x_1 x_2 \\ x_1 x_2 & x_2^2 \end{bmatrix}

以上对应的代码实现如下。

  1. # 6. 计算 2D 协方差在相机空间下的梯度并累加到 grad_sigma_camera_sorted 上
  2. # dL_da shape: (N_tile, P)
  3. dL_da = 0.5 * alpha * tile_grad_alpha
  4.  
  5. # 计算投影中间项 x = A * Delta
  6. x1 = a11 * du + a12 * dv  # (N_tile, P)
  7. x2 = a12 * du + a22 * dv  # (N_tile, P)
  8.  
  9. # 计算 2D 协方差梯度分量
  10. tile_grad_sig_11 = (dL_da * x1 * x1).sum(dim=1)  # (N_tile,)
  11. tile_grad_sig_12 = (dL_da * x1 * x2).sum(dim=1)  # (N_tile,)
  12. tile_grad_sig_22 = (dL_da * x2 * x2).sum(dim=1)  # (N_tile,)

接着我们重新组合成 2x2 的矩阵张量。

  1. # 堆叠成 (N_tile, 2, 2) 的 2D 协方差梯度
  2. tile_grad_sigma_camera = torch.stack([
  3.     tile_grad_sig_11, tile_grad_sig_12,
  4.     tile_grad_sig_12, tile_grad_sig_22
  5. ], dim=-1).view(-1, 2, 2)

最后我们将梯度累加回全局张量。

  1. # 原地原子累加到全局的相机空间协方差梯度张量上
  2. grad_sigma_camera_sorted.scatter_add_(
  3.     0,
  4.     ids_tile.unsqueeze(1).unsqueeze(2).expand(-1, 2, 2),
  5.     tile_grad_sigma_camera
  6. )

以上操作等价于:

  1. global_idx = ids_tile[i]
  2. grad_sigma_camera_sorted[global_idx, m, n] += tile_grad_sigma_camera[i, m, n]