反向传播概述

在之前的文章中,我们用 200 行左右的核心代码,加上几个辅助函数,就实现了一个完全可用的 3DGS 前向传播(Forward Pass)栅格化器。

现在,既然图像已经能够成功渲染出来,下一步也是模型训练中最关键的一步:实现反向传播(Backward Pass)。


为什么我们要“手写”反向传播?

在 PyTorch 中,系统自带了非常强大的自动求导机制(AutoGrad)。通常情况下,我们只需要编写前向传播,PyTorch 就能自动帮我们算出梯度。那为什么在 3DGS 中,我们还要费时费力地去自己写反向传播代码呢?

原因很简单:计算成本。如果让 PyTorch 用默认的自动求导机制去追踪栅格化过程中的每一个像素和每一个高斯点的运算图,内存消耗和计算代价将是极其昂贵的。为了达到实时渲染和高效训练的目的,我们通过底层推导,手动实现自定义的反向传播。


搭建 torch.autograd.Function 基础框架

要让 PyTorch 识别并使用我们手写的反向传播逻辑,我们需要将之前写好的前向传播代码封装进一个 PyTorch 类中,具体来说,就是继承 torch.autograd.Function

  1. class RasterizerFunction(torch.autograd.Function):
  2.     @staticmethod
  3.     def forward(ctx, pos, colors, opacity_raw, height, width, fx, fy, cx, cy, camera2world, sigma, near=2e-3, far=100, pixelGuard=64, tile_size=16, min_conic=1e-6, chi_square_clip=9.21, alpha_max=0.99, alpha_cutoff=1.0/255.0):
  4.  
  5.         img = gaussian_rasterization(
  6.             pos, colors, opacity_raw, height, width, fx, fy, cx, cy, camera2world,
  7.             sigma=sigma, near=near, far=far, pixelGuard=pixelGuard, tile_size=tile_size,
  8.             min_conic=min_conic, chi_square_clip=chi_square_clip, alpha_max=alpha_max, alpha_cutoff=alpha_cutoff
  9.         )
  10.         return img
  11.  
  12.     @staticmethod
  13.     def backward(ctx, grad_out):
  14.         pass

在这个结构中,有两个参数绝对不容忽视,它们是连接前向与反向的桥梁。

ctx(Context 实例)是一个上下文对象。前向传播(Forward)是从输入到输出的过程;而反向传播(Backward)是从输出往回推导梯度的过程。在这个过程中,反向传播往往需要用到前向传播时的“中间计算结果”。

ctx 就是用来帮我们“跨时空记账”的。在前向函数中,我们可以使用 ctx.save_for_backward() 存下变量;在反向函数中,再用 ctx.saved_tensors 把它们取出来继续运算。

grad_out 代表的是损失函数(Loss)相对于当前函数输出(通常是渲染出的图像)的梯度。反向传播就像多米诺骨牌倒推,grad_out 就是推倒这部分代码的第一股力。我们的目标是利用链式法则(Chain Rule),结合 grad_out,计算出更上游的所有输入变量的梯度。


我们需要计算哪些梯度?

backward 函数有一个硬性规定:前向传播 forward 函数接收了多少个输入参数,backward 函数就必须原原本本地返回多少个梯度值。 顺序要一一对应。

但是,我们真的需要优化所有的输入参数吗?并不是。

我们不关心的梯度,直接返回 None。例如相机的朝向矩阵(Camera Matrix)、焦距(Focal Length)、近平面(Near Plane)、屏幕网格大小(Tile Size)。我们不是在做“相机标定”任务,相机的物理参数是固定死不需要模型去优化的。它们虽然在数学上对输出有影响,但在当前的训练任务中,我们对它们的梯度不感兴趣。

我们必须返回的梯度(模型真正需要学习的参数)是 3D 高斯点的四大核心属性:位置 (position)、颜色 (color)、不透明度 (opacity) 以及协方差矩阵 (sigma)。


代码实现

明确了目标后,我们在 backward 函数中先为这四个核心属性初始化占位梯度(暂时用零填充),以打通整个代码链路。

  1. @staticmethod
  2. def backward(ctx, grad_out):
  3.     grad_pos = torch.zeros_like(pos)
  4.     grad_colors = torch.zeros_like(colors)
  5.     grad_opacity_raw = torch.zeros_like(opacity_raw)
  6.     grad_sigma = torch.zeros_like(sigma)
  7.  
  8.     return (
  9.         grad_pos,
  10.         grad_colors,
  11.         grad_opacity_raw,
  12.         None# height
  13.         None# width
  14.         None# fx
  15.         None# fy
  16.         None# cx
  17.         None# cy
  18.         None# camera2world
  19.         grad_sigma,
  20.         None# near
  21.         None# far
  22.         None# pixelGuard
  23.         None# tile_size
  24.         None# min_conic
  25.         None# chi_square_clip
  26.         None# alpha_max
  27.         None# alpha_cutoff
  28.     )