分块优化:索引计算
在 3D Gaussian Splatting 的渲染管线中,为了实现高效的光栅化,我们需要将屏幕划分成一个个网格区块(Tiles)。前面的步骤中,我们已经成功将高斯球(Gaussians)分配到了它们覆盖的 Tiles 中,并且得到了一份排序后的高斯球与 Tile 的对应关系。
这篇文章的核心任务是:在这个按 Tile 排序的一维数组中,我们如何快速定位每一个 Tile 对应的高斯球是从哪里开始,到哪里结束的?
例子说明
让我们从具体例子出发。假设我们在遍历处理过的 tile_ids 数组,它看起来像这样:[0, 0, 0, 1, 1, 2, 2, 2]。
这个数组直观地告诉了我们以下信息:
- Tile 0 被 3 个高斯球相交(占了索引 0, 1, 2)
- Tile 1 被 2 个高斯球相交(占了索引 3, 4)
- Tile 2 被 3 个高斯球相交(占了索引 5, 6, 7)
我们需要用代码推导出两个极其重要的数组:
1. start 数组:记录每个 Tile 在原数组中的起始索引(期望结果:[0, 3, 5])。
1. end 数组:记录每个 Tile 在原数组中的结束索引(期望结果:[3, 5, 8])。这里采用的是 Python 编程中常见的左闭右开区间 [start, end)。
统计每个 Tile 的高斯数量
要得到起止索引,我们首先得知道每个 Tile 到底“分到”了几个高斯球。
- unique_tile_ids, counts = torch.unique_consecutive(tile_ids_1d, return_counts=True)
torch.unique_consecutive 函数专门用来消除一维张量中连续出现的重复元素,并返回唯一元素构成的张量。当指定 return_counts=True 时,函数不仅返回去重后的 unique_tile_ids,还会额外返回一个 counts 数组,精准记录每个唯一元素在原张量中连续出现的次数。在上面的例子中,counts 是 [3, 2, 3]。
累积和计算 Start 索引
知道了每个 Tile 的高斯数量(counts),如何计算 start 索引?
这里的数学逻辑非常清晰:当前 Tile 的起始索引,等于前面所有 Tile 的高斯数量之和。用公式表达即为:
- unique_starts = torch.zeros_like(unique_tile_ids)
- unique_starts[1:] = torch.cumsum(counts[:-1], dim=0)
计算 End 索引
有了精准的 start 数组,end 数组的计算就变得很自然了:起点 + 当前 Tile 拥有的高斯数量 = 终点。
- unique_ends = unique_starts + counts
遍历 Tile
至此 Tile 相关的数据结构均已准备完毕。我们完善之前遗留的 Tile 遍历循环逻辑。
- for tile_id, start, end in zip(unique_tile_ids.tolist(), unique_starts.tolist(), unique_ends.tolist()):
- txi = tile_id % num_tiles_u
- tyi = tile_id // num_tiles_u
因为高斯点也排序过了,我们直接通过 start 和 end 进行获取。
- # 获取当前 Tile 包含的高斯范围
- ids_tile = gaussian_ids_sorted[start:end]
可视化验证
完整的可视化代码可查看 Commit d96198f:visualize_bonsai_gaussian。
选了两个视角,可视化出来有很多毛刺。不知道是拿到的预训练样本的问题,还是代码的问题。
样本问题的话,我们要自己解析没有毛刺的参照点云文件。但这个流程目前还没有。
代码问题的话,代码没有可以完全对照的版本进行参考,只能后续再参考官方的实现进行对比参照。
但总体的效果是对的,我们先继续后面的流程。持续关注这个问题,之后再回过头来看。