分块优化:索引计算

在 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 到底“分到”了几个高斯球。

  1. 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 的高斯数量之和。用公式表达即为:start[i] = \sum_{k=0}^{i-1} counts[k]

  1. unique_starts = torch.zeros_like(unique_tile_ids)
  2. unique_starts[1:] = torch.cumsum(counts[:-1], dim=0)

计算 End 索引

有了精准的 start 数组,end 数组的计算就变得很自然了:起点 + 当前 Tile 拥有的高斯数量 = 终点。

  1. unique_ends = unique_starts + counts

遍历 Tile

至此 Tile 相关的数据结构均已准备完毕。我们完善之前遗留的 Tile 遍历循环逻辑。

  1. for tile_id, start, end in zip(unique_tile_ids.tolist(), unique_starts.tolist(), unique_ends.tolist()):
  2.     txi = tile_id % num_tiles_u
  3.     tyi = tile_id // num_tiles_u

因为高斯点也排序过了,我们直接通过 startend 进行获取。

  1. # 获取当前 Tile 包含的高斯范围
  2. ids_tile = gaussian_ids_sorted[start:end]

可视化验证

完整的可视化代码可查看 Commit d96198f:visualize_bonsai_gaussian

选了两个视角,可视化出来有很多毛刺。不知道是拿到的预训练样本的问题,还是代码的问题。

样本问题的话,我们要自己解析没有毛刺的参照点云文件。但这个流程目前还没有。

代码问题的话,代码没有可以完全对照的版本进行参考,只能后续再参考官方的实现进行对比参照。

但总体的效果是对的,我们先继续后面的流程。持续关注这个问题,之后再回过头来看。