Scaffold-GS 论文与代码笔记
论文: Scaffold-GS: Structured 3D Gaussians for View-Adaptive Rendering
代码仓库: city-super/Scaffold-GS
摘要 与 简介
We introduce Scaffold-GS, which uses anchor points to distribute local 3D Gaussians, and predicts their attributes on-the-fly based on viewing direction and distance within the view frustum.
简介提到,传统高斯往往会过度膨胀高速球以适应每一个视角,从而忽略了场景结构,导致了冗余与低泛化性。
Furthermore, view-dependent effects are baked into individual Gaussian parameters with little interpolation capabilities, making it less robust to substantial view changes and lighting effects.
- 原版的3DGS对每个球储存了48个参数的球谐函数函数,使得可以精准表达反光效果,但是问题在于:球谐函数的值是固定的,如果遇到非常复杂的反射,原版 3DGS 单靠一个高斯球的球谐函数存不下这么复杂的视角变化,于是会在表面堆叠很多的微小高斯球,从左边看的时候,负责左边反光的高斯球显示出来;从右边看时,负责右边反光的高斯球显示出来,极大地破坏了场景的真实结构。
- 插值能力不足(每个参数由每个小球单独保存,没有插值可言)
相关工作三大基础路线对比总结
| 方法基础 (Base) | 代表性方法 | 优点 (Pros) | 缺点 (Cons) |
|---|---|---|---|
| 1. 基于 MLP 的神经场与渲染 (MLP-based Neural Fields…) | 早期的 NeRF 及其各类变体 | 1. 具有体积特性和 MLP 的归纳偏置,平滑插值能力极强。 2. 新视角合成的渲染质量极高 (SOTA)。 | 1. 需沿光线进行密集的采样计算,渲染速度极其缓慢。 2. 向复杂大规模场景的扩展与泛化能力有限。 3. 大多加速算法不是变慢就是牺牲画质。 |
| 2. 基于网格的神经场与渲染 (Grid-based Neural Fields…) | Plenoxel, K-planes, InstantNGP 等 | 1. 利用空间数据结构(如体素/哈希网格)存储局部特征,大幅加速特征查询。 2. 训练和推理速度显著提升(可达实时)。 3. 计算效率远高于纯全局的 MLP 模型。 | 1. 依然没有摆脱体积射线步进,仍需查询大量采样点来渲染一个像素。 2. 难以有效地表示和跳过“空白空间 ,存在无效计算冗余。 |
| 3. 基于点的神经场与渲染 (Point-based Neural Fields…) | 传统点云渲染, Point-NeRF, 3D-GS | 1. 利用几何图元直接进行光栅化,渲染速度极快(发挥 GPU 极限)。 2. 能够灵活地处理和表达拓扑变化。 3. 结合体渲染积分机制(如 3D-GS),可实现超高细节的实时渲染。 | 1. 传统点渲染容易受到孔洞和异常点困扰,产生伪影。 2. Point-NeRF 仍依赖体渲染步进,帧率受限。 3. 原版 3D-GS 缺乏几何结构约束,通过无脑堆叠产生严重冗余,且对复杂光线的泛化较弱。 |
To address this issue, we propose a hierarchical 3D Gaussian scene representation that respects the scene geometric structure, with anchor points initialized from SfM to encode local scene information and spawn local neural Gaussians. ----§ 3
Method Of Scaffold-GS
①体素化去冗余 (Voxelization):
将COLMAP 生成的稀疏 SfM 点云所在的空间划分为大小为 的体素(Voxel)网格,这里直接做除法并且向下取整,去重复后,将点云坐标离散化到对应的体素中心,建立一个结构化的稀疏网格。
② 构造锚点属性与特征增强 (Anchor Attributes & Feature Enhancement)
在确定了锚点的位置 后,每个锚点 都被赋予了一套核心属性,使其能够感知周围环境:
-
核心属性定义:
- :局部上下文特征(Local context feature),存储该区域的基础几何和外观信息。
- :缩放因子(Scaling factor),控制该锚点生成的神经高斯球的覆盖范围。
- : 个可学习的偏移量(Learnable offsets),定义了 个神经高斯相对于锚点中心的具体位置。
-
多分辨率特征库 (Multi-resolution Feature Bank): 为了让模型能够处理不同距离下的细节(LOD),作者为每个锚点创建了一个特征库:
其中 表示对特征进行 倍率的下采样,分别代表高、中、低三种分辨率。
-
视角相关的动态权重融合: 模型会根据相机位置 和锚点位置 计算:
- 相对距离:
- 观察方向:
然后,将距离和方向输入一个微型 MLP : 来预测三个权重:
在代码中,MLP被定义为
self.mlp_feature_bank = nn.Sequential(
nn.Linear(3+1, feat_dim),
nn.ReLU(True),
nn.Linear(feat_dim, 3),
nn.Softmax(dim=1)
).cuda()
最后,通过加权求和得到该视角下最终的综合锚点特征 :
1 if pc.use_feat_bank:
2 cat_view = torch.cat([ob_view, ob_dist], dim=1)
3 # 先拼接参数,然后调用微型 MLP 预测权重 {w, w1, w2}
4 bank_weight = pc.get_featurebank_mlp(cat_view).unsqueeze(dim=1)
5 # 对高、中、低三种分辨率特征进行加权求和,得到融合特征 feat
意义:通过这种设计,同一个锚点在近看时会倾向于使用高分辨率特征,远看时则切换到低分辨率特征,实现了平滑的 LOD 切换。
③ 神经高斯派生 (Neural Gaussian Derivation)
1. 神经高斯分布的初始化和推导 视锥体内的每个锚点会“孵化”出 个神经高斯。使用参数:
- 位置 :
- 不透明度:
- 旋转四元数 :
- 伸展 :
- 颜色 : 它们的位置 由锚点中心坐标 、学习到的偏移量 和缩放因子 共同决定:
其中,偏移量是被初始化为0的可学习参数。
1 # 偏移量 * 锚点缩放 (O * l_v)
2 offsets = offsets * scaling_repeat[:,:3]
3 # 最终球心位置 μ = 锚点位置 + 偏移量
4 xyz = repeat_anchor + offsets
2. 动态预测高斯属性 高斯球的透明度 、颜色 、旋转四元数 、缩放比例 等属性是根据当前的综合特征 以及相对相机视角 动态解码计算出来的,例如:
这里的 分别是控制这四个参数的神经网络MLP,定义的方法与上面的差不多
在代码中:
在调用 MLP 之前,代码先将 拼接在一起:
1. cat_local_view = torch.cat([feat, ob_view, ob_dist], dim=1) # [N, c+3+1]
之后再把各个特征送进MLP,但是预测缩放和旋转的参数使用了同一个MLP
1 # 调用 mlp_cov 一次性得到缩放和旋转的原始数据
2 scale_rot = pc.get_cov_mlp(cat_local_view)
3 # 从网络输出中拆分出缩放 s
4 scaling = scaling_repeat[:,3:] * torch.sigmoid(scale_rot[:,:3])
5 # 从网络输出中拆分出旋转 q 并进行归一化激活
6 rot = pc.rotation_activation(scale_rot[:,3:7])
3. 高效过滤剔除
- 只有视锥体内可见的锚点才会被激活运算。
- 为了进一步维持极限渲染速度,模型设置了一个不透明度阈值 。如果 MLP 预测出来的某个神经高斯的透明度低于阈值(),它在进入光栅化管线前就会被直接剔除。

④锚点优化
Growing Operation
把 3D 空间切分成边长为 的体素,在训练过程中,如果没渲染出来,就会在反向传播时给这个区域的高斯球一个很大的梯度,并且统计平均梯度 ,如果 ,说明在这个部分渲染很吃力,会在这里体素的中心放新锚点(如果本来没有锚点)。

从左至右,我们将神经高斯体在空间上量化为尺寸为 的多分辨率体素()。新的锚点将被添加到平均梯度大于 的体素中,以m作为“维度”:
Pruning Operation
计算N次训练迭代的与锚点关联的神经高斯球的不透明度,如果低于预期则移除
⑤损失函数设计
- 像素颜色损失:
- SSIM损失 :
- 体积正则化(结构相似性) :
- 总的损失函数: 其中体积正则化 为: 在这里, 表示场景中神经高斯体的数量, 是向量各个元素的乘积,这里指代每个神经高斯体的缩放向量 。
对于压缩的意义:体积正则化项鼓励神经高斯体保持较小的体积,并使它们之间的重叠最小化。机制在模型训练的源头就遏制了冗余表达,使得最终生成的几何结构更加紧凑和锐利,为后续进一步的量化和剪枝操作扫清了障碍。
完整训练过程
初始化模型
gaussians = GaussianModel(...)
scene = Scene(dataset, gaussians, ply_path=ply_path, shuffle=False)
# 创建场景,调用了/scene/gaussian_model.py 的create_from_pcd,创造点云
# 注册参数,调用了training_setup,把4个MLP的参数打包传递
训练主循环
- 动态训练学习率调整
gaussians.update_learning_rate(iteration)
锚点和 MLP 的学习率会随着训练不断衰减,保证后期模型能够收敛到精细的细节。
- 选择视角
1 # Pick a random Camera
2 if not viewpoint_stack:
3 viewpoint_stack = scene.getTrainCameras().copy()
4 viewpoint_cam = viewpoint_stack.pop(randint(0, len(viewpoint_stack)-1))
- 渲染
- 先计算哪些锚点在相机视野内部
voxel_visible_mask = prefilter_voxel(viewpoint_cam, gaussians, pipe,background)
- 调用render
render_pkg = render(...)
根据相机距离和方向,把锚点特征解码成一堆临时的神经高斯球。这些球只存在于显存中.
- 损失计算与梯度反向传播
1 Ll1 = l1_loss(image, gt_image)
2 ssim_loss = (1.0 - ssim(image, gt_image))
3 scaling_reg = scaling.prod(dim=1).mean()
4 loss = (1.0 - opt.lambda_dssim) * Ll1 + opt.lambda_dssim * ssim_loss + 0.01*scaling_reg
loss.backward()
- 结构优化
锚点梯度与不透明度的平均统计
# train.py
1 gaussians.training_statis(...)
1 #scene/gaussiam_model.py
2 def training_statis(...):
3 # 更新不透明度统计
4 ...
5 # 剔除负值
6
7 # 将对应的神经高斯的透明度求和,累加给锚点
8 temp_opacity = temp_opacity.view([-1, self.n_offsets])
9 self.opacity_accum[anchor_visible_mask] += temp_opacity.sum(dim=1, keepdim=True)
10 # 找出当前视野中激活的高斯球的掩码
...
11 temp_mask = combined_mask.clone()
12 combined_mask[temp_mask] = update_filter
13 # 计算当前高斯球受到的梯度范数
14 grad_norm = torch.norm(viewspace_point_tensor.grad[update_filter, :2], dim=-1, keepdim=True)
15
16 # 累加梯度误差到每个偏移量上
17 self.offset_gradient_accum[combined_mask] += grad_norm
18 self.offset_denom[combined_mask] += 1
- 锚点的Growing 和 Pruning
# train.py
gaussians.adjust_anchor(...)
def adjust_anchor(...):
# 1. 计算平均梯度
grads = self.offset_gradient_accum / self.offset_denom # [N*k, 1]
grads[grads.isnan()] = 0.0
# 2. 计算被观测到的锚点(迭代中)被观测超过一定次数
offset_mask = (self.offset_denom > check_interval*success_threshold*0.5).squeeze(dim=1)
# Grow
self.anchor_growing(grads_norm, grad_threshold, offset_mask)
# update offset_denom
...
# # prune anchors
# 计算掩码,只有出现次数足够多,并且不透明度低才能判定需要被修剪
prune_mask = (self.opacity_accum < min_opacity*self.anchor_demon).squeeze(dim=1)
anchors_mask = (self.anchor_demon > check_interval*success_threshold).squeeze(dim=1) # [N, 1]
prune_mask = torch.logical_and(prune_mask, anchors_mask) # [N]
...