稠密化与剪枝
截止提交 Commit 7302fa0,我们稍微停下脚步,让模型训练 600 次迭代。我们画出训练损失(Training Loss)和峰值信噪比(PSNR)的曲线,会发现它们都在稳步向好——损失是下降趋势,图像质量是上升趋势。
我们再提取一张预测图像,会发现仅仅 600 步,完整的图像分辨率就已经呈现,虽然还有一些区域需要微调,但整体效果已经比之前好多了。
但这还不够。为了让高斯点云完美贴合复杂的 3D 场景,我们需要引入这篇论文的核心操作:稠密化(Densification)与剪枝(Pruning)。
稠密化
稠密化并不是从头到尾都在进行的。论文中明确指出,在最初的预热阶段(Warm-up stage)结束后,模型才会开始考虑增加或修改高斯点的数量。
具体的规则是:在迭代 500 次之后,每隔 100 次迭代进行一次密实化。
- # 自适应密度控制与剪枝 (Densification & Pruning)
- # 起始步:500 步,结束步:3000 步,每 100 步触发一次
- if iteration > 500 and iteration <= 3000 and iteration % 100 == 0:
稠密化的目的是解决“重建过度”或“重建不足”的问题。作者通过经验发现,那些“问题高斯”通常具有一个共同点:较大的视空间位置梯度(View-space positional gradient)。
论文中给出了一个硬性阈值:
- tau_pos = 0.0002 # 触发分裂/克隆的位置梯度阈值 (2e-4)
被我们“盯上”的高斯点,到底该怎么处理呢?这取决于它的体型(Scale)。
- 克隆(Clone)针对的是那些体积小但位置梯度高的高斯点(通常意味着这里重建不足,需要多来几个小高斯帮忙)。
- 分裂(Split)针对的是那些体积大且位置梯度高的高斯点(通常意味着这个大高斯覆盖了太多细节,需要被拆分成更细粒度的小高斯)。
那么,如何计算高斯的“体型”呢?我们需要提取尺度参数的指数,并在 3D 维度中找到最大的那个轴。
- # 2. 计算高斯缩放
- scales = torch.exp(scale_raw) # scale_raw shape: (N, 3)
- max_scales = torch.max(scales, dim=1).values
- is_big = max_scales > tau_scale
- is_small = ~is_big
- mask_clone = is_high_grad & is_small
- mask_split = is_high_grad & is_big
有了掩码,我们就可以调用克隆和分裂处理函数了:
- # 3. 执行克隆
- if mask_clone.any():
- opt_params, optimizer = clone_gaussians(mask_clone, opt_params, optimizer)
- # 4. 执行分裂
- if mask_split.any():
- opt_params, optimizer = split_gaussians(mask_split, opt_params, optimizer)
剪枝
为了防止高斯点无限膨胀导致显存爆炸和渲染变慢,我们需要定期清理那些“没用”的高斯。
什么叫没用?就是不透明度(Opacity/Alpha)太低的高斯点,它们对最终的画面几乎没有贡献。
在这里,我们追求画质,选择阈值为 0.005。
- epsilon_alpha = 0.005 # 剪枝时的低透明度截断阈值
- # 5. 执行剪枝(剔除透明高斯点)
- # 重新读取当前最新的 alpha_raw 形状
- alpha_raw_latest = opt_params["alpha_raw"]
- mask_prune = torch.sigmoid(alpha_raw_latest) < epsilon_alpha
- # 为了防止剪掉所有的高斯,至少保留一个
- if mask_prune.all():
- mask_prune[0] = False
- if mask_prune.any():
- opt_params, optimizer = prune_gaussians(mask_prune, opt_params, optimizer)