参数分裂

在这篇文章中,我们讲解模型优化过程中的参数分裂(Splitting)环节。为了捕捉更丰富的细节,我们需要将一些“过大”的高斯球一分为二。

想象一下细胞分裂:大部分基因会完美继承,但体型和位置会发生改变。高斯球的分裂也是如此。

当我们决定要把某个高斯球分裂成两个子节点时:

  • 保持不变的:颜色特征(如球谐系数 fdc, frest)、不透明度(opacity)等。子节点 A 和 B 会完美复刻父节点的这些参数。因此代码中,我们不需要为 A 和 B 分别计算颜色,直接沿用即可。
  • 发生改变的:尺度(Scale)和 位置(Position)。

生成新高斯的位置(重参数化采样)

一个 3D 高斯由中心位置 \mu \in \mathbb{R}^3 和协方差矩阵 \Sigma \in \mathbb{R}^{3 \times 3} 描述。协方差矩阵可以分解为旋转矩阵 R(由四元数 q 计算得出)和缩放矩阵 S(对角矩阵):

\Sigma = R S S^T R^T

要从原本的大高斯中分裂出两个小高斯,论文里说明的做法是从原高斯的 3D 概率密度函数(PDF)中随机采样出新的中心点位置。

根据多元高斯分布的重参数化技巧(Reparameterization Trick),采样点 x 的计算公式为:

x = \mu + R S \epsilon

其中:

  • \epsilon \sim \mathcal{N}(0, I) 是从标准正态分布中采样的三维随机向量。
  • S \epsilon 将标准分布拉伸到原高斯的尺度。
  • R (S \epsilon) 将拉伸后的向量旋转到原高斯的方向。
  • 加上 \mu 最终将偏移量平移到原高斯的中心。

更新新高斯的缩放(Scale)

为了保持分裂前后的总体积/能量基本一致,新生高斯的体积必须减小。论文里说明将每个轴的缩放尺度缩小为原来的 1.6 倍。

\log(S_{new}) = \log\left(\frac{S}{1.6}\right) = \log(S) - \log(1.6)

这边稍微有点绕,需要注意。这里的 S 指的是高斯球在物理世界中的“真实缩放大小”(Physical Scale)。

我们可以把它理解为这个高斯椭球体在 X、Y、Z 三个轴上的半轴长度(标准差)。因为它是代表物理长度或大小的量,所以它必须大于 0(S > 0),不能是负数。

我们之前的初始化代码:

  1. scale_raw = torch.log(mean_dists.clamp(min=1e-6)).repeat(1, 3)

这里的 mean_dists(点云中点与点之间的平均距离)就是初始状态下的真实物理大小 S

那为什么要取 log?深度学习的优化器(比如 Adam)最喜欢优化的变量范围是 (-\infty, +\infty)。如果直接让优化器去更新真实大小 S,优化器很容易在梯度下降时把它减成负数,导致程序报错(高斯球的大小不能为负)。

为了解决这个问题,我们在对数域进行优化:

  • 我们不直接优化 S
  • 我们定义一个可以在 (-\infty, +\infty) 随便跑的变量,叫做 scale_raw
  • 它们之间的关系是:S = \exp(\text{scale\_raw}),反过来也就是 \text{scale\_raw} = \log(S)

scale_raw 是已经 log 过的,它本身就代表了真实大小 S 的对数值。我们希望真实大小 S 被除以 1.6:

S_{new} = \frac{S}{1.6}


代码实现

首先我们提取参数并构造旋转矩阵。以下提取了待分裂的 M 个高斯的 \muSR

  1. pos = parameters["pos"]
  2. scale_raw = parameters["scale_raw"]
  3. rot_raw = parameters["rot_raw"]
  4.  
  5. M = mask_split.sum().item()
  6. device = pos.device
  7. dtype = pos.dtype
  8.  
  9. # 1. Sample N new positions from the 3D Gaussian PDF of split elements
  10. mu = pos[mask_split]  # (M, 3)
  11. scale = torch.exp(scale_raw[mask_split])  # (M, 3)
  12. q_norm = torch.nn.functional.normalize(rot_raw[mask_split], dim=-1)  # (M, 4)
  13. # rot_raw has [w, x, y, z] format, needs conversion to [x, y, z, w] for quat_to_rotmat
  14. rot_xyzw = torch.cat([q_norm[..., 1:], q_norm[..., :1]], dim=-1)
  15. rot_mat = quat_to_rotmat(rot_xyzw)  # (M, 3, 3)

接着我们按上述公式采样生成新的位置。

  1. new_positions = []
  2. for _ in range(N):
  3.     epsilon = torch.randn((M, 3), device=device, dtype=dtype)  # (M, 3)
  4.     offset = scale * epsilon  # (M, 3)
  5.     # (M, 3, 3) @ (M, 3, 1) -> (M, 3, 1) -> (M, 3)
  6.     rotated_offset = (rot_mat @ offset.unsqueeze(-1)).squeeze(-1)  # (M, 3)
  7.     pos_child = mu + rotated_offset  # (M, 3)
  8.     new_positions.append(pos_child)

构造新的缩放。

  1. scale_new = scale_raw[mask_split] - torch.log(torch.tensor(1.6, device=device, dtype=dtype))  # (M, 3)

完整的新参数构造逻辑:

  1. # 2. Build new parameters
  2. new_parameters = {}
  3. for name, param in parameters.items():
  4.     old_val = param.detach()
  5.     remaining_val = old_val[~mask_split]
  6.  
  7.     if name == "pos":
  8.         new_part = torch.cat(new_positions, dim=0)  # (M * N, 3)
  9.     elif name == "scale_raw":
  10.         scale_new = scale_raw[mask_split] - torch.log(torch.tensor(1.6, device=device, dtype=dtype))  # (M, 3)
  11.         new_part = scale_new.repeat(N, 1)  # (M * N, 3)
  12.     else:
  13.         if old_val.dim() == 1:
  14.             new_part = old_val[mask_split].repeat(N)  # (M * N,)
  15.         else:
  16.             repeat_dims = (N,) + (1,) * (old_val.dim() - 1)  # e.g., (N, 1)
  17.             new_part = old_val[mask_split].repeat(*repeat_dims)  # (M * N, ...)
  18.  
  19.     new_val = torch.cat([remaining_val, new_part], dim=0)  # (N_old - M + M * N, ...)
  20.     new_parameters[name] = torch.nn.Parameter(new_val, requires_grad=True)