球谐函数阶梯式激活
在构建 3D Gaussian Splatting 训练管线的旅程中,我们将进入非常关键的一步:逐个激活球谐函数。
在 3DGS 中,球谐函数用于表达高斯球在不同视角下颜色的变化(视角依赖外观)。球谐函数分为不同的阶数(Degree),阶数越高,能表达的颜色细节(高频信息)越丰富。
那为什么要“逐个”激活球谐函数?为什么不一开始就把所有阶数的球谐函数全打开?
论文作者的说法:如果在训练初期就让模型去学习复杂的高频反光细节,由于初期视角的覆盖范围和整体形状还没稳定,模型极大概率会陷入局部最小值。
建议的做法是:第 0~1000 次迭代只激活 Degree 0(基础底色,低频)。第 1000~2000 次迭代激活 Degree 1。第 2000~3000 次迭代激活 Degree 2。之后激活完整的 Degree 3。这就像画画,先铺大色块,再慢慢抠细节。
代码实现
为了实现上述的阶梯式激活,我们需要编写两个核心函数。思路是:根据当前的训练迭代次数,生成一个遮罩(Mask),把当前不需要的高阶球谐系数直接“抹零”。
1. 生成边界遮罩函数:get_sh_bound_mask
在 3DGS 中,除了基础颜色(Degree 0,通常叫 f_dc),剩余的球谐系数(Degree 1, 2, 3,统称为 f_rest)对于每个 RGB 通道共有 15 个值。
Degree 1 有 3 项;Degree 2 有 5 项;Degree 3 有 7 项。总计 3 + 5 + 7 = 15 项。RGB 3 个通道,共计 45 个系数(形状为 [N, 45])。
- def get_sh_bound_mask(bound, device="cuda"):
- if bound == 0:
- num_terms = 0
- elif bound == 1:
- num_terms = 3
- elif bound == 2:
- num_terms = 8
- else:
- num_terms = 15
- mask_15 = torch.zeros(15, device=device)
- if num_terms > 0:
- mask_15[:num_terms] = 1.0
- return torch.cat([mask_15, mask_15, mask_15], dim=0)
2. 动态应用遮罩函数:apply_sh_masking
有了 Mask 之后,我们如何让它在训练中动起来呢?我们需要一个函数,将它与迭代次数(Iteration)挂钩。
- def apply_sh_masking(f_rest, iteration):
- bound = min(iteration // 1000, 3)
- mask = get_sh_bound_mask(bound, device=f_rest.device)
- return f_rest * mask
现在,我们把做好的零件装配到引擎里。在计算颜色前,稍微做一点修改。
- # 动态评估当前视角下的高斯 RGB 颜色并应用球谐函数分阶激活遮罩
- f_rest_effective = apply_sh_masking(f_rest, iteration)
- colors = evaluate_sh(f_dc, f_rest_effective, pos, c2w, interleaved=False)