分块优化:高斯 ID 与 Tile ID 映射
本篇教程将讲解如何建立每一个高斯点与其相交的 Tile 之间的映射关系,并通过一维扁平化编码(Flattening)和 PyTorch 向量化操作,为最终的按 Tile 深度排序做准备。
按需复制构建高斯 ID 数组 (repeat_interleave)
因为一个高斯点会覆盖多个 Tile,所以这个高斯点的 ID 在我们最终的相交列表中必须重复出现多次(覆盖了多少个 Tile,就出现多少次)。
若某个高斯球在图像中横向(U 轴)覆盖了
假设我们屏幕上有 3 个高斯球(ID 分别为 0, 1, 2):
- 0 号高斯覆盖了 2 个 Tile
- 1 号高斯覆盖了 1 个 Tile
- 2 号高斯覆盖了 3 个 Tile
我们希望最终生成的高斯 ID 数组为:[0, 0, 1, 2, 2, 2]。
PyTorch 中的 torch.repeat_interleave 函数能够极其高效地实现这一操作。
- # 构建高斯 ID 数组 (每个高斯按相交的 Tile 数量进行重复)
- num_tiles_per_gaussian = n_u * n_v # Shape: [N]
- num_gaussians = u_min_tile.shape[0]
- base_ids = torch.arange(num_gaussians, dtype=torch.int64, device=device) # Shape: [N]
- gaussian_ids = torch.repeat_interleave(base_ids, num_tiles_per_gaussian) # Shape: [M] (M 为所有高斯覆盖 Tile 的总数)
生成高斯覆盖的所有 Tile 网格并使用 Mask 过滤
对于每个高斯点,我们已经知道了它覆盖的 Tile 的最小值 u_min_tile / v_min_tile,以及它的跨度 span_indices_u / span_indices_v。
为了利用 PyTorch 向量化地生成所有高斯点所覆盖的每一个 Tile 的具体坐标,我们使用广播(Broadcasting)与掩码(Masking)。
- tile_u_grid = tile_u[:, :, None].expand(-1, -1, nv_max_item) # Shape: [N, nu_max_item, nv_max_item]
- tile_v_grid = tile_v[:, None, :].expand(-1, nu_max_item, -1) # Shape: [N, nu_max_item, nv_max_item]
- tile_u_flat = tile_u_grid[mask] # Shape: [M]
- tile_v_flat = tile_v_grid[mask] # Shape: [M]
tile_u 的形状是 [N, nu_max],我们通过 [:, :, None] 增加第三个维度,再用 .expand(-1, -1, nv_max_item) 将其在 v 维度上复制。
当我们使用 [mask](形状为 [N, nu_max, nv_max] 的布尔张量)对 3D 网格进行索引时,PyTorch 会自动将所有为 True 的元素收集并拉平成一个一维张量。
这种扁平化收集的顺序是行优先(Row-Major)的,刚好与我们用 repeat_interleave 生成的高斯 ID 顺序严格一一对应。
二维 Tile 坐标的“扁平化” (Flatten)
为了记录这些重复的高斯 ID 分别对应哪一个具体的 Tile,我们需要获取每个高斯与其相交的 Tile 坐标,并将其转换成一维索引。
在计算屏幕横向一共有多少个 Tile(num_tiles_u)时,我们不能使用简单的整除。
假设图像像素宽度
如果直接整除:17 // 8 = 2。程序会认为横向只有 2 个 Tile(覆盖了 0~15 像素)。第 16 像素对应的第 3 个 Tile 被截断抛弃,导致图像右侧出现黑边或空白。
正确的向上取整公式:
- num_tiles_u = (width + tile_size - 1) // tile_size
有了每个相交点的二维 Tile 坐标 (tile_u_flat, tile_v_flat),我们将其拍平为唯一的扁平化 Tile ID:
- flat_tile_id = tile_v_flat * num_tiles_u + tile_u_flat # Shape: [M]