颜色梯度反向传播

高斯渲染的反向传播往往非常耗费内存,如果直接依赖 PyTorch 的自动微分框架(Autograd)去跟踪每个像素和每个高斯球的中间计算图,显存会轻易导致 CUDA Out of Memory。

为了打破这一内存瓶颈,3DGS 库通常采用重计算(Recomputation)与自定义 Autograd 机制:在前向传播时只缓存轻量级的“路标”,并在反向传播时重新执行基于 Tile 的循环重计算权重,从而以计算时间换取了宝贵的显存空间。

本篇文章将对颜色梯度(Color Gradients)进行反向传播。


颜色梯度的链式法则

我们要计算的是损失函数 L 对某个高斯球输入颜色 c_n 的梯度:\frac{\partial L}{\partial c_n}。我们以单通道的颜色为例(RGB 三通道的推导完全一致)。

在前向传播中,某一个像素 p 的最终颜色 C_p 是由覆盖在该像素上的所有高斯球(按深度从前到后排序)混合而成的:

C_p = \sum_{n=1}^{N} c_n \cdot \alpha_{n,p} \cdot T_{n,p}

其中:

  • c_n 是第 n 个高斯球的颜色。
  • \alpha_{n,p} 是第 n 个高斯球在像素 p 处的不透明度。
  • T_{n,p} = \prod_{j=1}^{n-1}(1 - \alpha_{j,p}) 是透射率(Transmittance),表示光线穿过前面所有高斯球后,还能剩下多少比例能够到达当前高斯球。

在代码实现中,我们将单个高斯对该像素的渲染权重定义为:

w_{n,p} = \alpha_{n,p} \cdot T_{n,p}

所以,前向像素的混合公式可以简化为:

C_p = \sum_{n=1}^{N} c_n \cdot w_{n,p}

在反向传播时,PyTorch 会向我们传入上一层传回的梯度 grad_out,即损失函数 L 对像素颜色的导数:\frac{\partial L}{\partial C_p}

根据多元微积分的链式法则,损失 L 对高斯球颜色 c_n 的导数等于它在所有它覆盖的像素上产生的梯度贡献之和:

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

对前向公式 C_p = \sum c_k w_{k,p} 求偏导可得:

\frac{\partial C_p}{\partial c_n} = w_{n,p} = \alpha_{n,p} \cdot T_{n,p}

代入可得最终颜色梯度的计算公式:

\frac{\partial L}{\partial c_n} = \sum_{p \in \text{Pixels}} \frac{\partial L}{\partial C_p} \cdot w_{n,p}

公式中的像素级梯度 \frac{\partial L}{\partial C_p} 对应代码中的输入上游梯度 grad_out

公式中的高斯球渲染贡献权重 \alpha_{n} \cdot T_{n} 对应代码中重计算出来的局部权重 w

最终求得的全局梯度 \frac{\partial L}{\partial c_n} 对应代码中最终返回的 grad_colors


前向变量缓存

为了避免在 backward 中重新计算费时的排序与投影,我们在 RasterizerFunction.forward 渲染中需要将用于 backward 重组的张量提前缓存起来。

  1. # [关键准备] 将所有在 backward 阶段需要的 tensor 变量保存到 save_for_backward 中
  2. ctx.save_for_backward(
  3.     pos, colors, opacity_raw, sigma,
  4.     u_sorted, v_sorted, colors_sorted, opacity_sorted, inv_cov,
  5.     gaussian_ids_sorted, unique_tile_ids, unique_starts, unique_ends, indices_onscreen
  6. )
  7. ctx.meta = {
  8.     'height': height,
  9.     'width': width,
  10.     'num_tiles_u': num_tiles_u,
  11.     'tile_size': tile_size,
  12.     'min_conic': min_conic,
  13.     'chi_square_clip': chi_square_clip,
  14.     'alpha_max': alpha_max,
  15.     'alpha_cutoff': alpha_cutoff
  16. }

张量维度对齐与广播

在 backward 的 Tile 循环中,我们会拿到以下局部的关键张量:

1. 当前 Tile 对应的上游梯度:grad_out_tile (形状为 [P, 3]),其中 P 是当前 Tile 内包含的像素总数。

2. 当前 Tile 中重计算出的各高斯在像素上的权重:w (形状为 [N_tile, P]),其中 N_{tile} 是当前 Tile 内覆盖的高斯球数量。

我们需要将它们相乘,这需要做形状对齐与升维。

通过 PyTorch 广播:

(N_{tile}, P, 1) \times (1, P, 3) \to (N_{tile}, P, 3)

对中间维度(dim=1,即像素维度 P)进行 .sum(dim=1),最终我们算出了 Tile 内的高斯颜色梯度 (N_tile, 3)。这与我们推导的数学公式相契合。

  1. # 3. 计算颜色梯度:链式法则 dl/dcolor = grad_out_tile * w
  2. # w.unsqueeze(-1) shape: (N_tile, P, 1)
  3. # grad_out_tile.unsqueeze(0) shape: (1, P, 3)
  4. # 两者相乘后在 P (dim=1) 维度求和,得到 tile_grad_colors shape: (N_tile, 3)
  5. tile_grad_colors = (w.unsqueeze(-1) * grad_out_tile.unsqueeze(0)).sum(dim=1)

索引映射

在前向传播中,为了正确处理遮挡关系(从前向后渲染),所有高斯球经过了多轮过滤和基于深度的重排序。因此,我们在 backward 的 Tile 循环中拿到的 ids_tile 是相对于最终投影到屏幕上并排序后的高斯列表的局部索引(范围为 [0, N_onscreen-1])。

为了在反向传播时将排序后的梯度“原序归位”,我们在 forward 阶段通过对索引进行与数据流同步的操作,一步步构建了 indices_onscreen 映射。

1. 初始化原始索引

  1. orig_indices = torch.arange(N, device=pos.device)

2. 软边缘视锥剔除(pixelGuard)过滤

  1. orig_indices_v = orig_indices[mask]

3. 数值稳定过滤(协方差有限值判定)

  1. orig_indices_keep = orig_indices_v[keep]

4. 全局深度排序

  1. orig_indices_sorted = orig_indices_keep[order]

5. 屏幕可见边界裁切

  1. indices_onscreen = orig_indices_sorted[onscreen]

最终的 indices_onscreen(大小为 N_onscreen)中,位置 j 存储的值就是:在排序且可见的高斯列表中排在第 j 位的高斯球,在最初输入的 N 个高斯列表中的原始序号。


局部到全局的累加:torch.scatter_add_

计算出当前 Tile 内高斯的局部颜色梯度 tile_grad_colors(形状为 [N_tile, 3])后,我们需要将它们累加到全局的高斯梯度 grad_colors 中去。

  1. # 4. 使用 scatter_add_ 将梯度累加回排序的高斯球上
  2. # grad_colors_sorted shape: (N_onscreen, 3)
  3. grad_colors_sorted.scatter_add_(0, ids_tile.unsqueeze(1).expand(-1, 3), tile_grad_colors)