颜色梯度反向传播
高斯渲染的反向传播往往非常耗费内存,如果直接依赖 PyTorch 的自动微分框架(Autograd)去跟踪每个像素和每个高斯球的中间计算图,显存会轻易导致 CUDA Out of Memory。
为了打破这一内存瓶颈,3DGS 库通常采用重计算(Recomputation)与自定义 Autograd 机制:在前向传播时只缓存轻量级的“路标”,并在反向传播时重新执行基于 Tile 的循环重计算权重,从而以计算时间换取了宝贵的显存空间。
本篇文章将对颜色梯度(Color Gradients)进行反向传播。
颜色梯度的链式法则
我们要计算的是损失函数
在前向传播中,某一个像素
其中:
c_n 是第n 个高斯球的颜色。\alpha_{n,p} 是第n 个高斯球在像素p 处的不透明度。T_{n,p} = \prod_{j=1}^{n-1}(1 - \alpha_{j,p}) 是透射率(Transmittance),表示光线穿过前面所有高斯球后,还能剩下多少比例能够到达当前高斯球。
在代码实现中,我们将单个高斯对该像素的渲染权重定义为:
所以,前向像素的混合公式可以简化为:
在反向传播时,PyTorch 会向我们传入上一层传回的梯度 grad_out,即损失函数
根据多元微积分的链式法则,损失
对前向公式
代入可得最终颜色梯度的计算公式:
公式中的像素级梯度
公式中的高斯球渲染贡献权重
最终求得的全局梯度
前向变量缓存
为了避免在 backward 中重新计算费时的排序与投影,我们在 RasterizerFunction.forward 渲染中需要将用于 backward 重组的张量提前缓存起来。
- # [关键准备] 将所有在 backward 阶段需要的 tensor 变量保存到 save_for_backward 中
- ctx.save_for_backward(
- pos, colors, opacity_raw, sigma,
- u_sorted, v_sorted, colors_sorted, opacity_sorted, inv_cov,
- gaussian_ids_sorted, unique_tile_ids, unique_starts, unique_ends, indices_onscreen
- )
- ctx.meta = {
- 'height': height,
- 'width': width,
- 'num_tiles_u': num_tiles_u,
- 'tile_size': tile_size,
- 'min_conic': min_conic,
- 'chi_square_clip': chi_square_clip,
- 'alpha_max': alpha_max,
- 'alpha_cutoff': alpha_cutoff
- }
张量维度对齐与广播
在 backward 的 Tile 循环中,我们会拿到以下局部的关键张量:
1. 当前 Tile 对应的上游梯度:grad_out_tile (形状为 [P, 3]),其中
2. 当前 Tile 中重计算出的各高斯在像素上的权重:w (形状为 [N_tile, P]),其中
我们需要将它们相乘,这需要做形状对齐与升维。
通过 PyTorch 广播:
对中间维度(dim=1,即像素维度
- # 3. 计算颜色梯度:链式法则 dl/dcolor = grad_out_tile * w
- # w.unsqueeze(-1) shape: (N_tile, P, 1)
- # grad_out_tile.unsqueeze(0) shape: (1, P, 3)
- # 两者相乘后在 P (dim=1) 维度求和,得到 tile_grad_colors shape: (N_tile, 3)
- 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. 初始化原始索引
- orig_indices = torch.arange(N, device=pos.device)
2. 软边缘视锥剔除(pixelGuard)过滤
- orig_indices_v = orig_indices[mask]
3. 数值稳定过滤(协方差有限值判定)
- orig_indices_keep = orig_indices_v[keep]
4. 全局深度排序
- orig_indices_sorted = orig_indices_keep[order]
5. 屏幕可见边界裁切
- indices_onscreen = orig_indices_sorted[onscreen]
最终的 indices_onscreen(大小为 N_onscreen)中,位置 j 存储的值就是:在排序且可见的高斯列表中排在第 j 位的高斯球,在最初输入的
局部到全局的累加:torch.scatter_add_
计算出当前 Tile 内高斯的局部颜色梯度 tile_grad_colors(形状为 [N_tile, 3])后,我们需要将它们累加到全局的高斯梯度 grad_colors 中去。
- # 4. 使用 scatter_add_ 将梯度累加回排序的高斯球上
- # grad_colors_sorted shape: (N_onscreen, 3)
- grad_colors_sorted.scatter_add_(0, ids_tile.unsqueeze(1).expand(-1, 3), tile_grad_colors)