分块优化:Tile 与深度排序
在 3D Gaussian Splatting 的渲染架构中,为了提升渲染效率,我们通常不会选择在像素级别进行并行化,而是选择在 Tile 之间进行并行化。
要实现这一点,我们需要确保数据在进入渲染管线之前,满足一个非常严格的顺序要求。
目前,我们手中的数据(高斯 ID、扁平化 Tile ID 等)是按照“高斯顺序”排列,甚至可以说是完全随机的。但为了让 GPU 能够一个 Tile 接一个 Tile 地高效渲染,我们需要达成以下两个目标:
1. 宏观顺序:按 Tile ID 从小到大排列(例如:第0块、第1块、第2块……)。
2. 微观顺序:在同一个 Tile 内部,关联的高斯必须严格按照深度(Depth)进行排序。
这种多条件的排序被称为字典序排序。现在我们有两个排序键(Keys):
- 键 A:Tile ID(第一优先级)
- 键 B:高斯的深度值(第二优先级)
比较规则很简单:对于两组数据
代码实现
原生的 PyTorch 并没有直接提供针对多个数组的无缝字典序排序函数(例如 NumPy 中的 lexsort)。如果我们直接写循环,效率会极低。
为了解决这个问题,我们采用了一种非常经典的“位图压缩/数值合成”技巧:将两个排序变量合并为一个单一的数字,然后对这个单一数字进行标准排序。
我们将 Tile ID 和深度索引合成一个新的变量 comp。
合成公式为:
这里最关键的变量是
- # 双重排序:首先按 Tile ID 排序,同一个 Tile 内按深度(Gaussian ID)排序
- M = num_gaussians + 1
- comp = flat_tile_id * M + gaussian_ids
现在,我们只需要对这个合成后的 1D 数组进行一次普通的排序即可,PyTorch 将同时返回排序后的值以及对应的索引排列(permutation)。
- comp_sorted, permutation = torch.sort(comp)
排序完成后,我们不再关心 comp 本身,而是利用它的产物把我们需要的数据提取出来。
直接使用排序产生的排列(permutation)对原始高斯 ID 数组重新索引。
- gaussian_ids_sorted = gaussian_ids[permutation]
对排好序的合成值进行向下取整除法(除以
- tile_ids_1d = torch.div(comp_sorted, M, rounding_mode='floor')