参数分裂
在这篇文章中,我们讲解模型优化过程中的参数分裂(Splitting)环节。为了捕捉更丰富的细节,我们需要将一些“过大”的高斯球一分为二。
想象一下细胞分裂:大部分基因会完美继承,但体型和位置会发生改变。高斯球的分裂也是如此。
当我们决定要把某个高斯球分裂成两个子节点时:
- 保持不变的:颜色特征(如球谐系数 fdc, frest)、不透明度(opacity)等。子节点 A 和 B 会完美复刻父节点的这些参数。因此代码中,我们不需要为 A 和 B 分别计算颜色,直接沿用即可。
- 发生改变的:尺度(Scale)和 位置(Position)。
生成新高斯的位置(重参数化采样)
一个 3D 高斯由中心位置
要从原本的大高斯中分裂出两个小高斯,论文里说明的做法是从原高斯的 3D 概率密度函数(PDF)中随机采样出新的中心点位置。
根据多元高斯分布的重参数化技巧(Reparameterization Trick),采样点
其中:
\epsilon \sim \mathcal{N}(0, I) 是从标准正态分布中采样的三维随机向量。S \epsilon 将标准分布拉伸到原高斯的尺度。R (S \epsilon) 将拉伸后的向量旋转到原高斯的方向。- 加上
\mu 最终将偏移量平移到原高斯的中心。
更新新高斯的缩放(Scale)
为了保持分裂前后的总体积/能量基本一致,新生高斯的体积必须减小。论文里说明将每个轴的缩放尺度缩小为原来的 1.6 倍。
这边稍微有点绕,需要注意。这里的
我们可以把它理解为这个高斯椭球体在 X、Y、Z 三个轴上的半轴长度(标准差)。因为它是代表物理长度或大小的量,所以它必须大于 0(
我们之前的初始化代码:
- scale_raw = torch.log(mean_dists.clamp(min=1e-6)).repeat(1, 3)
这里的 mean_dists(点云中点与点之间的平均距离)就是初始状态下的真实物理大小
那为什么要取 log?深度学习的优化器(比如 Adam)最喜欢优化的变量范围是
为了解决这个问题,我们在对数域进行优化:
- 我们不直接优化
S 。 - 我们定义一个可以在
(-\infty, +\infty) 随便跑的变量,叫做 scale_raw。 - 它们之间的关系是:
S = \exp(\text{scale\_raw}) ,反过来也就是\text{scale\_raw} = \log(S) 。
scale_raw 是已经 log 过的,它本身就代表了真实大小
代码实现
首先我们提取参数并构造旋转矩阵。以下提取了待分裂的
- pos = parameters["pos"]
- scale_raw = parameters["scale_raw"]
- rot_raw = parameters["rot_raw"]
- M = mask_split.sum().item()
- device = pos.device
- dtype = pos.dtype
- # 1. Sample N new positions from the 3D Gaussian PDF of split elements
- mu = pos[mask_split] # (M, 3)
- scale = torch.exp(scale_raw[mask_split]) # (M, 3)
- q_norm = torch.nn.functional.normalize(rot_raw[mask_split], dim=-1) # (M, 4)
- # rot_raw has [w, x, y, z] format, needs conversion to [x, y, z, w] for quat_to_rotmat
- rot_xyzw = torch.cat([q_norm[..., 1:], q_norm[..., :1]], dim=-1)
- rot_mat = quat_to_rotmat(rot_xyzw) # (M, 3, 3)
接着我们按上述公式采样生成新的位置。
- new_positions = []
- for _ in range(N):
- epsilon = torch.randn((M, 3), device=device, dtype=dtype) # (M, 3)
- offset = scale * epsilon # (M, 3)
- # (M, 3, 3) @ (M, 3, 1) -> (M, 3, 1) -> (M, 3)
- rotated_offset = (rot_mat @ offset.unsqueeze(-1)).squeeze(-1) # (M, 3)
- pos_child = mu + rotated_offset # (M, 3)
- new_positions.append(pos_child)
构造新的缩放。
- scale_new = scale_raw[mask_split] - torch.log(torch.tensor(1.6, device=device, dtype=dtype)) # (M, 3)
完整的新参数构造逻辑:
- # 2. Build new parameters
- new_parameters = {}
- for name, param in parameters.items():
- old_val = param.detach()
- remaining_val = old_val[~mask_split]
- if name == "pos":
- new_part = torch.cat(new_positions, dim=0) # (M * N, 3)
- elif name == "scale_raw":
- scale_new = scale_raw[mask_split] - torch.log(torch.tensor(1.6, device=device, dtype=dtype)) # (M, 3)
- new_part = scale_new.repeat(N, 1) # (M * N, 3)
- else:
- if old_val.dim() == 1:
- new_part = old_val[mask_split].repeat(N) # (M * N,)
- else:
- repeat_dims = (N,) + (1,) * (old_val.dim() - 1) # e.g., (N, 1)
- new_part = old_val[mask_split].repeat(*repeat_dims) # (M * N, ...)
- new_val = torch.cat([remaining_val, new_part], dim=0) # (N_old - M + M * N, ...)
- new_parameters[name] = torch.nn.Parameter(new_val, requires_grad=True)