Tile 像素坐标计算

在进行复杂的图像渲染(例如高斯泼溅 Gaussian Splatting 的体积渲染)时,我们经常需要将整个图像划分为多个小块(Tiles)来进行并行计算或优化内存。在这篇文章中,我们将一步步拆解如何遍历这些 Tile,计算其内部的像素坐标,并最终求出像素点与渲染目标(如高斯中心)之间的距离。


从 Tile 索引到像素边界

我们的目标是确定每一个 Tile 在原图中的确切位置。假设图像被划分为多个大小为 T \times T 像素的 Tile。

首先,我们拥有两个关键变量:TX_iTY_i。它们代表当前 Tile 在水平和垂直方向上的网格索引(例如左上角第一个 Tile 是 (0, 0),右边紧挨着的是 (1, 0))。

接下来,我们需要将“Tile 索引”转换为具体的“像素坐标”:

  • 起始像素(左上角):X_0 = TX_i \times TY_0 = TY_i \times T
  • 结束像素(右下角):X_1 = (TX_i + 1) \times TY_1 = (TY_i + 1) \times T
  1. # 计算当前 Tile 的像素边界 (X0, Y0 为左上角,X1, Y1 为右下角)
  2. x0 = txi * tile_size
  3. y0 = tyi * tile_size
  4. x1 = min((txi + 1) * tile_size, width)
  5. y1 = min((tyi + 1) * tile_size, height)

在实际的图像处理中,图像的宽 (W) 和高 (H) 未必是 Tile 尺寸 T 的完美倍数。因此,在计算右下角坐标时,必须进行边界裁剪(Clamping)。


构建 Tile 内部的二维像素网格

知道了边界后,我们要获取这个 Tile 内每一个像素的坐标。这时,我们请出 PyTorch 中的两个重要函数:torch.arangetorch.meshgrid

  1. # 生成 X 和 Y 方向的一维坐标序列
  2. xs = torch.arange(x0, x1, dtype=pos.dtype, device=pos.device)
  3. ys = torch.arange(y0, y1, dtype=pos.dtype, device=pos.device)
  4.  
  5. # 生成 Tile 内的二维像素网格 (使用 xy 索引模式,以符合图像直觉)
  6. px, py = torch.meshgrid(xs, ys, indexing='xy')

torch.meshgrid 的作用是接收多个一维张量,并生成一个多维的坐标网格。

  • 输入:xs 包含了从 X_0X_1 的所有横坐标,ys 包含了纵坐标。
  • 输出:pxpy 都是二维张量。px 的每一个位置存放着该像素的 X 坐标,py 存放着 Y 坐标。将它们叠在一起,就得到了网格中每个点的完整 (X, Y) 坐标。

全局 1D 索引计算

为了后续运算更高效,我们将二维张量转换为一维张量。

  1. # 重塑为一维数组 (px_u 和 px_v 分别代表该 Tile 内部的横、纵坐标序列)
  2. px_u = px.reshape(-1)
  3. px_v = py.reshape(-1)
  4.  
  5. # 计算在全局一维图像数组中的像素索引 (Y * Width + X)
  6. pixel_idx_1D = (px_v * width + px_u).to(torch.int64)

计算像素与高斯中心的距离

这部分进入了高斯体积渲染的前置计算。我们需要计算当前 Tile 内的所有像素,与我们要渲染的高斯点(Gaussian Min/Mean,记作 \mu)之间的距离。

当前有 N 个高斯点,以及当前 Tile 拍平后的 T^2 个像素。

  1. # 计算像素点到高斯中心的距离 (du, dv)
  2. du = px_u.unsqueeze(0) - u_tile.unsqueeze(1)  # (N, P)
  3. dv = px_v.unsqueeze(0) - v_tile.unsqueeze(1)  # (N, P)