alpha 梯度计算

在这篇文章中,我们计算 \alpha 的梯度。我们的目标是得到损失函数 L\alpha_n 的偏导数:

\frac{\partial L}{\partial \alpha_n}

利用链式法则,我们先计算输出像素颜色对 \alpha_n 的偏导,再结合损失函数传回的梯度:

\frac{\partial L}{\partial \alpha_n} = \sum_{p \in \text{Pixels}} \frac{\partial L}{\partial C_p} \cdot \frac{\partial C_p}{\partial \alpha_n}

在 3DGS 的渲染管线中,单个像素 p 的颜色 C_p 渲染公式为:

C_p = \sum_{j=1}^{N} c_j \cdot \alpha_{j} \cdot T_{j}

其中,透射率 T_j 代表光线穿过前 j-1 个高斯后,还剩下多少比例能够到达当前高斯:

T_j = \prod_{m=1}^{j-1} (1 - \alpha_m)

我们对第 n 个高斯的 \alpha_n 求偏导。将总和拆分为三部分来看:

1. 排在 n 前面的高斯 (j < n):它们的渲染不受 \alpha_n 的任何影响,偏导数为 0

2.n 个高斯自己 (j = n):其项为 c_n \alpha_n T_n,对 \alpha_n 求导的结果为 c_n T_n

3. 排在 n 后面的高斯 (j > n):对于任意 j > n,其透射率 T_j 的连乘积中必然包含了 (1 - \alpha_n) 因子。我们可以写成:

T_j = (1 - \alpha_n) \cdot \prod_{m \neq n}^{j-1} (1 - \alpha_m)

因此,对 \alpha_n 求导得到:

\frac{\partial T_j}{\partial \alpha_n} = - \frac{T_j}{1 - \alpha_n}

将这一项代入后面的求和式中,即可提取出共同的系数 - \frac{1}{1 - \alpha_n}

\sum_{j>n} c_j \alpha_j \frac{\partial T_j}{\partial \alpha_n} = - \frac{\sum_{j>n} c_j \alpha_j T_j}{1 - \alpha_n} = - \frac{S_n}{1 - \alpha_n}

其中我们定义 S_n = \sum_{j>n} c_j \alpha_j T_j,意为:排在第 n 个高斯之后的所有高斯对该像素渲染颜色的累积贡献和。

结合上述结果,我们得到了最终的公式:

\frac{\partial C_p}{\partial \alpha_n} = c_n T_n - \frac{S_n}{1 - \alpha_n}

这个公式有直观的理解:

  • c_n T_n 是正向激励,表示“自身显色”。如果增加不透明度 \alpha_n,高斯自身的颜色就会更多地显示出来。
  • - \frac{S_n}{1 - \alpha_n} 是反向抑制,表示“遮挡他人”。如果增加不透明度 \alpha_n,高斯会变厚,从而挡住排在后面的所有高斯 (S_n),使它们能透出的光变少。

代码实现

数学上的 S_n = \sum_{j>n} 要求计算“排在自己后面元素的累加和”(不含自身)。但是在 PyTorch 中,内置的 torch.cumsum 默认只能从前向后累加。

为了高效计算,我们采用了“翻转-累加-平移补零-再翻转回来”的标准张量操作范式。

  1. colors_tile = colors_sorted[ids_tile]  # (N_tile, 3)
  2. cw = colors_tile.unsqueeze(1) * w.unsqueeze(-1)  # (N_tile, P, 3)
  3. cw_flip = torch.flip(cw, dims=[0])  # (N_tile, P, 3)
  4. s_flip = torch.cumsum(cw_flip, dim=0)  # (N_tile, P, 3)
  5. s_shifted_flip = torch.cat([
  6.     torch.zeros((1, s_flip.shape[1], s_flip.shape[2]), device=s_flip.device, dtype=s_flip.dtype),
  7.     s_flip[:-1]
  8. ], dim=0)  # (N_tile, P, 3)
  9. s = torch.flip(s_shifted_flip, dims=[0])  # (N_tile, P, 3)

主要是 S_n 计算步骤比较多,后续就是变量带入公式,比较容易对照。

  1. denom = torch.clamp(1.0 - alpha, min=1e-8).unsqueeze(-1)  # (N_tile, P, 1)
  2. d_out_d_alpha = colors_tile.unsqueeze(1) * ti.unsqueeze(-1) - s / denom  # (N_tile, P, 3)
  3.  
  4. tile_grad_alpha = (d_out_d_alpha * grad_out_tile.unsqueeze(0)).sum(dim=-1)  # (N_tile, P)