训练

命令行参数:ArgumentParser

ArgumentParser 是 Python 标准库 argparse 模块中的核心类,用于轻松编写用户友好的命令行接口。它的主要功能是解析命令行参数,使你的 Python 脚本能够接收外部输入,从而实现灵活配置。

核心工作流程

使用 ArgumentParser 通常遵循以下三个步骤:

  1. 创建解析器 (Create the parser)
    首先,你需要创建一个 ArgumentParser 对象。这个对象将保存所有解析命令行参数所需的信息。

    1
    2
    3
    4
    import argparse

    # 创建一个解析器实例,并添加描述信息
    parser = argparse.ArgumentParser(description='这是一个用于处理文件的工具')
  2. 添加参数 (Add arguments)
    接下来,通过调用 add_argument() 方法来告诉解析器,你的程序需要接收哪些命令行参数。你可以定义参数的名称、类型、默认值、帮助信息等。

    • 位置参数 (Positional Arguments): 必须提供的参数,根据它们在命令行中的位置来确定。
    • 可选参数 (Optional Arguments):--- 开头的参数,通常有默认值,是可选的。
    1
    2
    3
    4
    5
    6
    7
    8
    # 添加一个位置参数 'filename'
    parser.add_argument('filename', help='要处理的文件名')

    # 添加一个可选参数 '--verbose',它是一个开关,如果提供则为 True
    parser.add_argument('-v', '--verbose', action='store_true', help='启用详细输出模式')

    # 添加一个带类型的可选参数 '--count',默认为 1
    parser.add_argument('--count', type=int, default=1, help='处理次数')

    type

    指定该参数应该被转换成的数据类型。

    • 默认:字符串 (str)。
    • 常用int, float, bool(注意:bool类型通常配合 action 使用,直接设 type=bool 可能会有坑)。
    • 作用:如果用户输入了错误类型(比如输入了字母但要求是数字),程序会自动报错并提示。

    default

    如果命令行中没有提供该参数,变量将取这个值。

    • 场景:比如 --ip 参数,默认是 127.0.0.1,如果你不特意去改它,程序就用本地地址。

    help

    一个描述性的字符串,解释这个参数是做什么的。

    • 作用:当用户在命令行输入 python script.py -h--help 时,这段文字会显示出来,帮助用户了解如何使用程序。

    action

    指定当参数在命令行中出现时,应该采取什么动作。最常用的有两个:

    • store(默认):保存参数的值。例如 --count 5,就把 5 存起来。
    • store_true / store_false:这是一个开关。
    • 如果你定义 parser.add_argument('--debug', action='store_true')
    • 默认情况下,debugFalse
    • 如果你在命令行里写了 --debug,它的值就变成了 True。这非常适合用来控制是否开启某种模式(如调试模式、静默模式)。

    nargs

    指定应该读取多少个命令行参数。

    • ?:0个或1个。
    • *:所有剩余的参数放入一个列表。
    • +:至少一个,剩余的放入一个列表。
    • 数字:例如 2,表示需要紧跟两个数值(常用于坐标 x, y)。
    1
    2
    3
    4
    5
    6
    7
    # 示例代码
    parser.add_argument(
    "--detect_anomaly", # 1. 参数名:在命令行输入 --detect_anomaly
    action="store_true", # 2. 动作:这是一个开关。写了就是True,不写就是False
    default=False, # 3. 默认值:默认不开启
    help="Set debug mode" # 4. 帮助信息:解释它是干嘛的
    )

    运行效果:

    • 情况 A:你直接运行 python train.py
    • args.detect_anomaly 的值是 False
    • 情况 B:你运行 python train.py --detect_anomaly
    • args.detect_anomaly 的值是 True

    这就是 add_argument 的核心逻辑:定义规则,自动解析,赋值给变量。

  3. 解析参数 (Parse the arguments)
    最后,调用 parse_args() 方法。它会读取命令行输入(例如 python script.py myfile.txt -v --count 3),进行解析,并返回一个包含所有参数值的对象。

    1
    2
    3
    4
    5
    6
    7
    # 解析命令行参数
    args = parser.parse_args()

    # 通过属性访问解析出的参数
    print(f"文件名: {args.filename}")
    print(f"详细模式: {args.verbose}")
    print(f"处理次数: {args.count}")

核心优势

  • 自动生成帮助信息: 当你运行脚本并加上 -h--help 参数时,argparse 会自动生成一个格式美观的帮助文档,列出所有可用参数及其说明。
  • 类型检查与转换: 它可以根据你设置的 type 参数自动将命令行输入的字符串转换为整数、浮点数等,并在类型不匹配时给出清晰的错误提示。
  • 处理默认值: 为可选参数设置默认值,简化了用户的使用。
  • 错误处理: 当用户提供了无效的参数或缺少必需的参数时,它会自动退出程序并显示错误信息。

总而言之,ArgumentParser 是构建专业、易用的命令行工具的标准方式,它极大地简化了参数处理逻辑,并提升了用户体验。

forward

预处理:preprocess

视锥剔除:in_frustum

1
2
3
4
float4 p_hom = transformPoint4x4(p_orig, projmatrix);
float p_w = 1.0f / (p_hom.w + 0.0000001f);
float3 p_proj = { p_hom.x * p_w, p_hom.y * p_w, p_hom.z * p_w };
p_view = transformPoint4x3(p_orig, viewmatrix);

把点进行视图变换,转换到屏幕空间。

1
2
3
4
5
6
7
8
9
10
if (p_view.z <= 0.2f)
{
if (prefiltered)
{
printf("Point is filtered although prefiltered is set. This shouldn't happen!");
__trap();
}
return false;
}
return true;

如果点的深度,说明在相机后面,retrurn false表示不在视锥内。如果 prefilteredtrue,意味着上层代码已经宣称“这个点是可见的”,现在的 GPU 代码算出来这个点其实是在相机后面的,说明出错了,__trap()强制中断。

1
2
if (!in_frustum(idx, orig_points, viewmatrix, projmatrix, prefiltered, p_view))
return;

如果不在视锥内,则不处理这个点。

投影

1
2
3
4
float3 p_orig = { orig_points[3 * idx], orig_points[3 * idx + 1], orig_points[3 * idx + 2] };
float4 p_hom = transformPoint4x4(p_orig, projmatrix);
float p_w = 1.0f / (p_hom.w + 0.0000001f);
float3 p_proj = { p_hom.x * p_w, p_hom.y * p_w, p_hom.z * p_w };

orig_points里取出这个点,再用投影矩阵进行变换,最后归一化。

计算3D协方差矩阵:computeCov3D

1
2
3
4
glm::mat3 S = glm::mat3(1.0f);
S[0][0] = mod * scale.x;
S[1][1] = mod * scale.y;
S[2][2] = mod * scale.z;


创建一个对角矩阵 。其中 mod(全局缩放系数), 是各轴向的缩放。

1
2
3
4
5
glm::mat3 R = glm::mat3(
1.f - 2.f * (y * y + z * z), 2.f * (x * y - r * z), 2.f * (x * z + r * y),
2.f * (x * y + r * z), 1.f - 2.f * (x * x + z * z), 2.f * (y * z - r * x),
2.f * (x * z - r * y), 2.f * (y * z + r * x), 1.f - 2.f * (x * x + y * y)
);


将四元数转换为 3x3 旋转矩阵。(代码中 r 对应

1
2
glm::mat3 M = S * R;
glm::mat3 Sigma = glm::transpose(M) * M;


1
2
3
4
5
6
cov3D[0] = Sigma[0][0]; // xx
cov3D[1] = Sigma[0][1]; // xy
cov3D[2] = Sigma[0][2]; // xz
cov3D[3] = Sigma[1][1]; // yy
cov3D[4] = Sigma[1][2]; // yz
cov3D[5] = Sigma[2][2]; // zz

协方差矩阵 是一个对称矩阵(即 )。为了节省显存和带宽,代码只存储了矩阵的上三角部分

计算2D协方差矩阵:computeCov2D

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
float3 t = transformPoint4x3(mean, viewmatrix);

const float limx = 1.3f * tan_fovx;
const float limy = 1.3f * tan_fovy;
const float txtz = t.x / t.z;
const float tytz = t.y / t.z;
t.x = min(limx, max(-limx, txtz)) * t.z;
t.y = min(limy, max(-limy, tytz)) * t.z;
```
将高斯球的中心点 `mean` 从世界坐标系变换到相机坐标系 `t` 。计算视野边界`limx`,`limy`,乘以 `1.3f` 是为了留出额外的缓冲区域,计算点在图像平面上的投影位置`txtz`,`tytz`,确保数值不小于左边界(`-limx`),不大于右边界(`limx`),保证点在视锥体内。

```cpp
glm::mat3 J = glm::mat3(
focal_x / t.z, 0.0f, -(focal_x * t.x) / (t.z * t.z),
0.0f, focal_y / t.z, -(focal_y * t.y) / (t.z * t.z),
0, 0, 0);


计算雅可比矩阵,透视除法后2D协方差矩阵不需要深度信息,第三行置0

1
2
3
4
glm::mat3 W = glm::mat3(
viewmatrix[0], viewmatrix[4], viewmatrix[8],
viewmatrix[1], viewmatrix[5], viewmatrix[9],
viewmatrix[2], viewmatrix[6], viewmatrix[10]);

计算视图变换矩阵,这里只提取了视图矩阵的 3x3 旋转部分(去掉了平移,因为协方差只与形状和方向有关,与位置无关)

1
2
3
4
5
6
7
8
9
10
glm::mat3 T = W * J;

glm::mat3 Vrk = glm::mat3(
cov3D[0], cov3D[1], cov3D[2],
cov3D[1], cov3D[3], cov3D[4],
cov3D[2], cov3D[4], cov3D[5]);

glm::mat3 cov = glm::transpose(T) * glm::transpose(Vrk) * T;

return { float(cov[0][0]), float(cov[0][1]), float(cov[1][1]) };


Vrk重构 3D 协方差矩阵,运用协方差变换公式,最后由于协方差矩阵的对称性,只返回3个值。

1
float3 cov = computeCov2D(p_orig, focal_x, focal_y, tan_fovx, tan_fovy, cov3D, viewmatrix);

计算渲染边界

1
2
3
float mid = 0.5f * (cov.x + cov.z);
float lambda1 = mid + sqrt(max(0.1f, mid * mid - det));
float lambda2 = mid - sqrt(max(0.1f, mid * mid - det));

利用 2x2 矩阵特征值的解析解公式,对于一个 2x2 对称矩阵,其特征值 满足:

较大的特征值对应椭圆长轴长度的平方,较小的特征值对应椭圆短轴长度的平方。

1
float my_radius = ceil(3.f * sqrt(max(lambda1, lambda2)));

sqrt(lambda1)取出长轴的标准差,根据3-Sigma原则,高斯分布的 99.7% 的能量集中在均值周围 3 倍标准差的范围内。ceil(...)向上取整,确保覆盖所有可能的像素。把高斯椭圆近似看成圆,my_radius 是这个圆屏幕上占据的最大像素半径。

1
2
float2 point_image = { ndc2Pix(p_proj.x, W), ndc2Pix(p_proj.y, H) };
getRect(point_image, my_radius, rect_min, rect_max, grid);

其中

1
2
3
4
5
6
7
8
rect_min = {
min(grid.x, max((int)0, (int)((p.x - max_radius) / BLOCK_X))),
min(grid.y, max((int)0, (int)((p.y - max_radius) / BLOCK_Y)))
};
rect_max = {
min(grid.x, max((int)0, (int)((p.x + max_radius + BLOCK_X - 1) / BLOCK_X))),
min(grid.y, max((int)0, (int)((p.y + max_radius + BLOCK_Y - 1) / BLOCK_Y)))
};

ndc2Pix将高斯中心点从NDC转换为屏幕像素坐标。根据中心点 point_image 和半径 my_radius/ BLOCK_X将像素坐标转换为线程块坐标,计算覆盖范围涉及到的Tile,用rect_minrect_max(左上角和右下角)定义了一个受影响的线程块矩形区域,后续计算只在这些Tile上进行。

1
2
if ((rect_max.x - rect_min.x) * (rect_max.y - rect_min.y) == 0)
return;

如果计算出的矩形面积为 0(例如高斯球太小,或者完全在屏幕外),则直接终止函数,节省 GPU 算力。

计算颜色:computeColorFromSH

1
2
3
glm::vec3 pos = means[idx];
glm::vec3 dir = pos - campos;
dir = dir / glm::length(dir);

计算从相机位置 campos 到高斯球中心 pos 的向量,并将其归一化。这代表了当前的观察角度。

1
glm::vec3* sh = ((glm::vec3*)shs) + idx * max_coeffs;

获取系数指针

1
2
3
result = SH_C0 * sh[0];
result = result - SH_C1 * y * sh[1] + SH_C1 * z * sh[2] - SH_C1 * x * sh[3];
...






SH_C0, SH_C1, SH_C2 等是球谐函数的常数系数(如)。xyxxxy等是与观察角度有关的函数(如),它们共同构成了基函数,sh存储了三维向量rgb的系数(即)。

排序

生成键值:duplicateWithKeys

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
// 1. 获取当前线程处理的高斯球ID
auto idx = cg::this_grid().thread_rank();
if (idx >= P) return; // 如果ID超过总数,退出

// 2. 可见性剔除:如果半径为0,说明这个球看不见,不参与排序
if (radii[idx] > 0)
{
// 3. 计算写入位置的偏移量
// offsets 数组记录了前面所有高斯球生成了多少个“分身”
uint32_t off = (idx == 0) ? 0 : offsets[idx - 1];

// 4. 计算这个高斯球覆盖了屏幕上的哪些 Tiles,存在rect_min和rect_max里
uint2 rect_min, rect_max;
getRect(points_xy[idx], radii[idx], rect_min, rect_max, grid);

// 5. 遍历这个矩形内的每一个 Tile
for (int y = rect_min.y; y < rect_max.y; y++)
{
for (int x = rect_min.x; x < rect_max.x; x++)
{
// 计算 Tile ID ,第 y 行第 x 列的 Tile ID = y * 总列数 + x
uint64_t key = y * grid.x + x;

// 将 Tile ID 左移 32 位,这样 Tile ID 占据 64位整数的高32位
key <<= 32;

// 将深度值强制转换为整数,按位或,填入低32位
// 这里直接操作内存,把 float 变成了 uint32
key |= *((uint32_t*)&depths[idx]);

// 把拼好的大数字(key)和高斯球ID(value)存起来
gaussian_keys_unsorted[off] = key;
gaussian_values_unsorted[off] = idx;
off++;
}
}
}

为什么这样排序有效?
因为在计算机排序中,数字是按位从高到低比较的:

  1. 先比高位:如果两个 Key 的高 32 位(Tile ID)不同,那么数值大的肯定属于不同的 Tile。这保证了同一个 Tile 的数据排在一起。
  2. 再比低位:如果高 32 位相同(同一个 Tile),计算机就会去比低 32 位(Depth)。这保证了深度顺序(从前到后或从后到前)。

排序:SortPairs

1
2
3
4
5
6
7
8
9
10
11
cub::DeviceRadixSort::SortPairs(
binningState.list_sorting_space, // 临时工作区(刚才算出来的那块地)
binningState.sorting_size, // 工作区大小
binningState.point_list_keys_unsorted, // 输入:乱序的键值 (TileID + Depth)
binningState.point_list_keys, // 输出:排好序的键值
binningState.point_list_unsorted, // 输入:乱序的高斯球ID
binningState.point_list, // 输出:排好序的高斯球ID
num_rendered, // 要排序的元素总数
0, // 排序的起始位(从第0位开始)
32 + bit // 排序的结束位(排多少位)
);

这里调用了 NVIDIA CUB 库

对于 64 位的键值 ,基数排序算法会把它看作二进制串:

排序过程(从低位到高位):

  1. Pass 1:看 ,分两桶(0和1)。
  2. Pass 2:看 ,分两桶(保持上一轮顺序)。
  3. Pass 64:看 (最高位,即 TileID 的最高位)。

0, 32 + bit 含义:处理从第 0 位到第 32+bit 位的所有二进制位

如果 bit 是 32,那就是处理全部 64 位。这就保证了先比深度(低32位),再比 TileID(高32位),最终 TileID 占主导地位。

划分边界:identifyTileRanges

这段代码是 3D Gaussian Splatting 渲染管线中 Tile-Based Radix Sort 的第三阶段(范围识别)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
// 1. 获取当前线程处理的位置索引
auto idx = cg::this_grid().thread_rank();
if (idx >= L) return; // 如果索引超过总数,退出

// 2. 读取当前元素的 Key,并提取 Tile ID
uint64_t key = point_list_keys[idx];
uint32_t currtile = key >> 32; // 取高32位,得到当前属于哪个 Tile

// 3. 处理列表开头
if (idx == 0)
ranges[currtile].x = 0;

// 4. 处理中间部分
else
{
uint32_t prevtile = point_list_keys[idx - 1] >> 32; // 读前一个元素的 Tile ID
if (currtile != prevtile) // 如果当前 Tile 和前一个不一样,说明到了“交界处”
{
ranges[prevtile].y = idx; // 前一个 Tile 在这里结束
ranges[currtile].x = idx; // 当前 Tile 从这里开始
}
}

// 5. 处理列表结尾
if (idx == L - 1)
ranges[currtile].y = L;

ranges[tile_id].x代表该 Tile 的数据在排序列表中的起始索引ranges[tile_id].y代表该 Tile 的数据在排序列表中的结束索引

如果是列表的第一个元素(idx == 0),那么当前这个 Tile 的起始位置(.x)就是 0。如果不是第一个元素,通过比较当前元素前一个元素TileID,来判断是否到了边界。如果是列表的最后一个元素(idx == L - 1),那么当前这个 Tile 的结束位置(.y)就是列表的总长度 L

颜色混合:renderCUDA

block = tile

thread = pixel

1
2
3
4
float T = 1.0f;
uint32_t contributor = 0;
uint32_t last_contributor = 0;
float C[CHANNELS] = { 0 };

float T = 1.0f;

  • 含义透射率。表示光线穿过当前已累积的所有高斯球后,还剩下多少能量。
  • 初始值(即 100%)。在开始绘制前,假设没有任何物体遮挡,光线完全透过。

uint32_t contributor = 0;

  • 含义当前贡献者计数器。它记录当前正在处理列表中的第几个高斯球。

uint32_t last_contributor = 0;

  • 含义最后有效贡献者索引。记录最后一个显著改变了该像素颜色的高斯球在列表中的位置。
  • 用途:这个变量主要用于反向传播(训练阶段)。在计算梯度时,我们需要知道哪些高斯球参与了该像素的成像。如果 衰减到接近 0,后面的高斯球就不再处理了,last_contributor 就标记了这个截止点。

float C[CHANNELS] = { 0 };

  • 含义累积颜色
  • 作用:存储当前像素计算出的最终 RGB 颜色值。
1
2
3
int num_done = __syncthreads_count(done);
if (num_done == BLOCK_SIZE)
break;

done 变量在之前的循环中被置为 true,意味着该像素的透射率 已经趋近于 0,后续的高斯球对它不再有贡献。如果 num_done == BLOCK_SIZE,说明当前线程块内的所有线程(即该 Tile 内的所有像素)都已经完成了渲染(光线被完全遮挡)。

1
int progress = i * BLOCK_SIZE + block.thread_rank();

progress:计算当前线程负责加载的高斯球在全局列表中的相对进度。

1
if (range.x + progress < range.y)

确保我们要访问的高斯球索引没有超出该线程块分配的总范围(range)。为了防止处理屏幕边缘或列表末尾时的越界访问。

1
2
3
4
int coll_id = point_list[range.x + progress];
collected_id[block.thread_rank()] = coll_id;
collected_xy[block.thread_rank()] = points_xy_image[coll_id];
collected_conic_opacity[block.thread_rank()] = conic_opacity[coll_id];

从排序好的 point_list 中取出高斯球的全局 ID (coll_id)。根据 ID,从全局显存中读取该高斯球的属性(xy 坐标、conic_opacity 参数)。

1
2
3
4
5
6
float2 xy = collected_xy[j];
float2 d = { xy.x - pixf.x, xy.y - pixf.y };
float4 con_o = collected_conic_opacity[j];
float power = -0.5f * (con_o.x * d.x * d.x + con_o.z * d.y * d.y) - con_o.y * d.x * d.y;
if (power > 0.0f)
continue;
  • contributor:计数器加 1,记录当前处理到了第几个高斯球。
  • 数学原理:这是高斯函数的指数部分,是像素点到高斯点的距离


    其中 con_o.x, .y, .z 存储了协方差逆矩阵的参数。
  • 如果 power > 0,说明距离太远或者概率密度极低,该高斯球对此像素没有贡献,直接跳过。
1
2
3
4
5
6
7
8
9
float alpha = min(0.99f, con_o.w * exp(power));
if (alpha < 1.0f / 255.0f)
continue;
float test_T = T * (1 - alpha);
if (test_T < 0.0001f)
{
done = true;
continue;
}
  • 计算当前像素位置的高斯值,并乘以高斯球本身的不透明度(con_o.w)。

  • min(0.99f, ...):强制限制 最大为 0.99。这是为了防止数值溢出或完全遮挡导致的梯度消失问题(如果 ,则 ,后续梯度无法回传)。

  • 如果计算出的 太小(小于 1/255),说明不透明度太低了,肉眼几乎不可见,直接跳过以节省算力。

  • 计算如果加上当前这个高斯球,剩余的光线透射率 会变成多少。

    如果 小于 0.0001,说明光线已经被阻挡了 99.99% 以上,后面的高斯球无论是什么颜色都看不见了。done = true:标记该像素渲染完成,后续循环将不再处理该像素。

1
2
for (int ch = 0; ch < CHANNELS; ch++)
C[ch] += features[collected_id[j] * CHANNELS + ch] * alpha * T;
  • 用Alpha Blending 公式累计颜色。当前高斯球贡献的颜色 = 球的颜色 球的不透明度 之前所有球的透射率。features存储了高斯球的 RGB 颜色值。
1
2
T = test_T;
last_contributor = contributor;
  • 更新透射率,供下一个高斯球使用。记录最后一个对该像素有贡献的高斯球索引。

backward

我们有每个像素渲染出的颜色和真实的颜色,根据这些计算出loss,从而计算出

我们需要更新的是:高斯中心位置,高斯形状,高斯的颜色,高斯的不透明度。因此需要计算出loss对每个参数的梯度

预处理:preprocessCUDA

已知损失函数对2D屏幕空间高斯均值的梯度(代码中 dL_dmean2D),通过链式法则反向传播,计算损失对3D高斯以下参数的梯度:

  • 3D均值dL_dmeans
  • 球谐系数dL_dsh
  • 缩放dL_dscale
  • 旋转dL_drot

投影路径的反向传播

3D高斯的参数通过多条路径影响最终损失,总梯度是所有路径梯度的累加
$$\frac{\partial L}{\partial \text{means}} = \underbrace{\frac{\partial L}{\partial \text{means}}}{\text{投影路径}} + \underbrace{\frac{\partial L}{\partial \text{means}}}{\text{颜色路径}} + \underbrace{\frac{\partial L}{\partial \text{means}}}_{\text{形状路径}}$$

1
2
3
auto idx = cg::this_grid().thread_rank();
if (idx >= P || !(radii[idx] > 0))
return;

为每个3D高斯分配一个线程,仅处理radii[idx] > 0,即屏幕上渲染半径大于0的高斯。不可见高斯对损失无贡献。

回顾正向投影过程



展开为分量形式:

再通过透视除法得到2D归一化坐标:

具体计算过程:

1
2
3
4
5
6
7
float3 m = means[idx];
float4 m_hom = transformPoint4x4(m, proj);
float m_w = 1.0f / (m_hom.w + 0.0000001f);

glm::vec3 dL_dmean;
float mul1 = (proj[0] * m.x + proj[4] * m.y + proj[8] * m.z + proj[12]) * m_w * m_w;
float mul2 = (proj[1] * m.x + proj[5] * m.y + proj[9] * m.z + proj[13]) * m_w * m_w;
  • m = means[idx]:取出当前3D高斯的均值
  • m_hom = transformPoint4x4(m, proj):正向计算齐次裁剪空间坐标(对应上述数学公式)。
  • m_w = 1/(m_hom.w + eps):计算透视除法分母的倒数,加小常数 eps 防止除零。
  • mul1mul2
    预计算,用于简化后续商的导数计算。
1
2
3
dL_dmean.x = (proj[0] * m_w - proj[3] * mul1) * dL_dmean2D[idx].x + (proj[1] * m_w - proj[3] * mul2) * dL_dmean2D[idx].y;
dL_dmean.y = (proj[4] * m_w - proj[7] * mul1) * dL_dmean2D[idx].x + (proj[5] * m_w - proj[7] * mul2) * dL_dmean2D[idx].y;
dL_dmean.z = (proj[8] * m_w - proj[11] * mul1) * dL_dmean2D[idx].x + (proj[9] * m_w - proj[11] * mul2) * dL_dmean2D[idx].y;

已知,需计算(y、z分量同理)。根据链式法则

再根据商的导数法则):

代入投影矩阵的偏导数(),最终得到代码中的表达式。
'_' allowed only in math mode \frac{\partial \text{mean2D}.x}{\partial m.x} = \frac{\text{proj}[0] \cdot P_{hom}.w - P_{hom}.x \cdot \text{proj}[3]}{(P_{hom}.w)^2} = \frac{\text{proj}[0]}{P_{hom}.w} - \frac{\text{proj}[3] \cdot P_{hom}.x}{(P_{hom}.w)^2} = \texttt{proj[0] * m_w - proj[3] * mul1}

1
dL_dmeans[idx] += dL_dmean;

投影路径的梯度累加到3D均值的总梯度中。3D均值还会通过“颜色路径”和“形状路径”影响损失,总梯度需累加所有路径的贡献。

颜色路径的反向传播:computeColorFromSH

已知损失对RGB颜色的梯度 (代码中 dL_dcolor[idx]),也即,需通过链式法则计算两个目标梯度:

  1. 损失对SH系数的梯度 (代码中 dL_dshs),
  2. 损失对3D高斯位置的梯度 (代码中 dL_dmeans)。

计算对SH系数的梯度:

1
2
3
glm::vec3 pos = means[idx];
glm::vec3 dir_orig = pos - campos;
glm::vec3 dir = dir_orig / glm::length(dir_orig);

计算视角方向 dir_orig),并归一化得到单位方向 dir)。

由于RGB是SH系数的线性函数,因此 是第 个SH基函数在 处的值)。根据链式法则:

代码中按阶数 从0到3依次计算,对应各阶基函数的导数:

阶数 时实球谐基函数:'_' allowed only in math modeY_0^0 = \text{SH_C0},其中 '_' allowed only in math mode\text{SH_C0} = \frac{1}{\sqrt{4\pi}}(归一化常数)。

1
2
float dRGBdsh0 = SH_C0;
dL_dsh[0] = dRGBdsh0 * dL_dRGB;

阶数 时实球谐基函数:
'_' allowed only in math mode \begin{cases} Y_1^{-1} = -\text{SH_C1} \cdot y \ Y_1^{0} = \text{SH_C1} \cdot z \ Y_1^{1} = -\text{SH_C1} \cdot x \end{cases}
其中 '_' allowed only in math mode\text{SH_C1} = \sqrt{\frac{3}{4\pi}}(归一化常数), 是单位方向 的分量。

1
2
3
4
5
6
float dRGBdsh1 = -SH_C1 * y;
float dRGBdsh2 = SH_C1 * z;
float dRGBdsh3 = -SH_C1 * x;
dL_dsh[1] = dRGBdsh1 * dL_dRGB;
dL_dsh[2] = dRGBdsh2 * dL_dRGB;
dL_dsh[3] = dRGBdsh3 * dL_dRGB;

阶数 时部分实球谐基函数
'_' allowed only in math mode \begin{cases} Y_2^{-2} = \text{SH_C2}[0] \cdot xy \ Y_2^{-1} = \text{SH_C2}[1] \cdot yz \ Y_2^{0} = \text{SH_C2}[2] \cdot (2z^2 - x^2 - y^2) \ \cdots \end{cases}
代码逻辑与 类似,计算基函数值并乘以梯度 dL_dRGB 得到SH系数梯度。

时基函数形式更复杂(如 ),代码逻辑同上。

计算对均值的梯度:


其中 方向向量对位置的雅可比矩阵,代码中通过球谐系数与方向向量的乘积实现。

1
2
3
glm::vec3 dRGBdx(0, 0, 0);   // ∂RGB / ∂dir.x
glm::vec3 dRGBdy(0, 0, 0); // ∂RGB / ∂dir.y
glm::vec3 dRGBdz(0, 0, 0); // ∂RGB / ∂dir.z

在 1 阶、2 阶、3 阶 SH 里不断 累加

1
2
3
4
5
6
7
8
9
10
11
12
13
14
// 1阶 SH
dRGBdx = -SH_C1 * sh[3];
dRGBdy = -SH_C1 * sh[1];
dRGBdz = SH_C1 * sh[2];

// 2阶 SH
dRGBdx += ...;
dRGBdy += ...;
dRGBdz += ...;

// 3阶 SH
dRGBdx += ...;
dRGBdy += ...;
dRGBdz += ...;

d是从相机指向高斯的单位向量

对 pos 求导:

其中

代码的巧妙之处在于,它将这两步合并了。

在第一步中计算的那些包含 x, y, z 的导数表达式(如 -SH_C1 * sh[3]),实际上已经是 相乘后的结果

例如,对于一阶球谐函数,其基函数与方向分量成正比(如 )。当我们对它求关于位置 的导数时,链式法则告诉我们:

是归一化方向向量的分量 。对它求关于 的导数,就会得到一个包含 的项。代码中那些看似复杂的表达式,正是这些导数经过化简后的形式。

形状路径的反向传播(缩放+旋转→协方差)

1
2
if (scales)
computeCov3D(idx, scales[idx], scale_modifier, rotations[idx], dL_dcov3D, dL_dscale, dL_drot);
  • 数学背景:3D高斯的协方差矩阵由缩放旋转合成。正向过程为:

    1. 将四元数 rotations 转换为旋转矩阵
    2. 构建缩放对角矩阵(对角元素为 scales * scale_modifier)。
    3. 计算3D协方差:
  • 反向传播任务
    已知dL_dcov3D),通过链式法则计算:
    1.dL_dscale):缩放参数的梯度。
    2.dL_drot):旋转四元数的梯度。

我们已知损失函数 对协方差矩阵 的梯度 (代码中为 dL_dSigma)。
我们需要求 对中间矩阵 的梯度。

根据矩阵求导法则,若 ,则:

代码对应:

1
2
// dSigma_dM = 2 * M
glm::mat3 dL_dM = 2.0f * M * dL_dSigma;

反向传播需要用到前向传播时的中间结果,所以代码首先重新计算了旋转矩阵 和缩放矩阵

我们需要求
因为 ,根据链式法则:

这本质上是一个矩阵乘法:

1
2
3
4
5
6
7
glm::mat3 Rt = glm::transpose(R);
glm::mat3 dL_dMt = glm::transpose(dL_dM);

// 点乘操作实际上是在做矩阵乘法的行提取
dL_dscale->x = glm::dot(Rt[0], dL_dMt[0]); // 对应 S_xx 的梯度
dL_dscale->y = glm::dot(Rt[1], dL_dMt[1]);
dL_dscale->z = glm::dot(Rt[2], dL_dMt[2]);

接下来我们需要求
路径是:

首先,利用 ,我们可以得到 的梯度:

代码中通过 dL_dMt 乘以 s 来实现这一步(利用了转置性质)。

接下来是 旋转矩阵对四元数的导数 ()
因为 的每个元素都是 的二次函数(如 等),其导数是线性的。
代码中那几行极其复杂的算式(如 2 * z * (dL_dMt[0][1] - dL_dMt[1][0])...)正是 的展开形式。

最后构建出

1
*dL_drot = float4{ dL_dq.x, dL_dq.y, dL_dq.z, dL_dq.w };