分块优化:Tile 相交处理

在 3D Gaussian Splatting 中,当我们得到每个高斯的边界框(AABB:Axis-Aligned Bounding Box)后,下一步面临一个棘手的问题:如何知道每个高斯具体和哪些屏幕瓷砖(Tile)相交?在这篇文章中,我们就来拆解这段核心代码的推演过程。


构建基础坐标

假设我们有一个高斯(Gaussian 0),我们已经算出它覆盖的起始 Tile 坐标(最小坐标),即 u_min(比如是 5)。同时,我们知道在 U 方向上,整个系统中高斯能跨越的最大 Tile 数量是 max_u

为了得到这个高斯覆盖的所有 U 坐标,我们需要生成一个“跨度(Span)”。

  1. # 计算高斯与 Tile 的相交映射
  2. device = pos.device
  3. span_indices_u = torch.arange(nu_max_item, device=device, dtype=torch.int64)
  4. span_indices_v = torch.arange(nv_max_item, device=device, dtype=torch.int64)

如果我们传入 max_u = 4,它会返回 [0, 1, 2, 3]。在我们的场景中,它代表了偏移量:+0, +1, +2, +3

接下来,我们将起始点和跨度相加:

  1. # 起始点加上跨度,得到每个高斯在每个跨度上的 Tile 坐标 (未过滤,有过度填充)
  2. tile_u = u_min_tile[:, None] + span_indices_u[None, :]  # Shape: [N, nu_max_item]
  3. tile_v = v_min_tile[:, None] + span_indices_v[None, :]  # Shape: [N, nv_max_item]

如果 u_min = 5,那么 tile_u 就会变成 [5, 6, 7, 8]。看起来非常完美,对吧?

如果每个高斯跨越的 Tile 数量完全一样,上述逻辑毫无问题。但现实是骨感的:Gaussian 0 可能跨越了 3 个 Tile。Gaussian 1 可能跨越了 4 个 Tile。Gaussian 2 可能只跨越了 1 个 Tile。

为了让张量维度对齐(形成规整的 2D Tensor),我们刚才借用了全局的 max_u 来做计算。这就导致了一个严重的后果:过度填充(Overfilling)。


掩码(Mask)与维度扩展

既然多填了,我们就需要一个“过滤器”把它们剔除。过滤的标准很简单:偏移量必须小于该高斯真实的跨度 n_u

例如,对于 Gaussian 0,它的 n_u = 3。那么只有当 span_indices_u(即 [0, 1, 2, 3])中的值小于 3 时,才是我们要保留的有效数据(True),否则就是废弃数据(False)。

为了在 PyTorch 中同时对 N 个高斯进行这种比较,我们需要引入维度扩展(Unsqueezing)。

  1. # 创建掩码,以过滤掉超出该高斯实际跨度 (n_u, n_v) 的 Tile
  2. mask_u = span_indices_u[None, :] < n_u[:, None# Shape: [N, nu_max_item]
  3. mask_v = span_indices_v[None, :] < n_v[:, None# Shape: [N, nv_max_item]
  4. mask = mask_u[:, :, None] & mask_v[:, None, :]  # Shape: [N, nu_max_item, nv_max_item]

在 PyTorch 中,tensor[:, None] 等价于 tensor.unsqueeze(1)。它可以在指定位置强行插入一个长度为 1 的新维度。