3DGS代码解读
训练
命令行参数:ArgumentParser
ArgumentParser 是 Python 标准库 argparse 模块中的核心类,用于轻松编写用户友好的命令行接口。它的主要功能是解析命令行参数,使你的 Python 脚本能够接收外部输入,从而实现灵活配置。
核心工作流程
使用 ArgumentParser 通常遵循以下三个步骤:
创建解析器 (Create the parser)
首先,你需要创建一个ArgumentParser对象。这个对象将保存所有解析命令行参数所需的信息。1
2
3
4import argparse
# 创建一个解析器实例,并添加描述信息
parser = argparse.ArgumentParser(description='这是一个用于处理文件的工具')添加参数 (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')。 - 默认情况下,
debug是False。 - 如果你在命令行里写了
--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的核心逻辑:定义规则,自动解析,赋值给变量。解析参数 (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 | float4 p_hom = transformPoint4x4(p_orig, projmatrix); |
把点进行视图变换,转换到屏幕空间。
1 | if (p_view.z <= 0.2f) |
如果点的深度retrurn false表示不在视锥内。如果 prefiltered 为 true,意味着上层代码已经宣称“这个点是可见的”,现在的 GPU 代码算出来这个点其实是在相机后面的,说明出错了,__trap()强制中断。
1 | if (!in_frustum(idx, orig_points, viewmatrix, projmatrix, prefiltered, p_view)) |
如果不在视锥内,则不处理这个点。
投影
1 | float3 p_orig = { orig_points[3 * idx], orig_points[3 * idx + 1], orig_points[3 * idx + 2] }; |
在orig_points里取出这个点,再用投影矩阵进行变换,最后归一化。
计算3D协方差矩阵:computeCov3D
1 | glm::mat3 S = glm::mat3(1.0f); |
创建一个对角矩阵 mod(全局缩放系数),
1 | glm::mat3 R = glm::mat3( |
将四元数转换为 3x3 旋转矩阵。(代码中 r 对应
1 | glm::mat3 M = S * R; |
1 | cov3D[0] = Sigma[0][0]; // xx |
协方差矩阵
计算2D协方差矩阵:computeCov2D
1 | float3 t = transformPoint4x3(mean, viewmatrix); |
计算雅可比矩阵
1 | glm::mat3 W = glm::mat3( |
计算视图变换矩阵
1 | glm::mat3 T = W * J; |
用Vrk重构 3D 协方差矩阵,运用协方差变换公式,最后由于协方差矩阵的对称性,只返回3个值。
1 | float3 cov = computeCov2D(p_orig, focal_x, focal_y, tan_fovx, tan_fovy, cov3D, viewmatrix); |
计算渲染边界
1 | float mid = 0.5f * (cov.x + cov.z); |
利用 2x2 矩阵特征值的解析解公式,对于一个 2x2 对称矩阵,其特征值
较大的特征值对应椭圆长轴长度的平方,较小的特征值对应椭圆短轴长度的平方。
1 | float my_radius = ceil(3.f * sqrt(max(lambda1, lambda2))); |
sqrt(lambda1)取出长轴的标准差,根据3-Sigma原则,高斯分布的 99.7% 的能量集中在均值周围 3 倍标准差的范围内。ceil(...)向上取整,确保覆盖所有可能的像素。把高斯椭圆近似看成圆,my_radius 是这个圆屏幕上占据的最大像素半径。
1 | float2 point_image = { ndc2Pix(p_proj.x, W), ndc2Pix(p_proj.y, H) }; |
其中
1 | rect_min = { |
ndc2Pix将高斯中心点从NDC转换为屏幕像素坐标。根据中心点 point_image 和半径 my_radius,/ BLOCK_X将像素坐标转换为线程块坐标,计算覆盖范围涉及到的Tile,用rect_min和rect_max(左上角和右下角)定义了一个受影响的线程块矩形区域,后续计算只在这些Tile上进行。
1 | if ((rect_max.x - rect_min.x) * (rect_max.y - rect_min.y) == 0) |
如果计算出的矩形面积为 0(例如高斯球太小,或者完全在屏幕外),则直接终止函数,节省 GPU 算力。
计算颜色:computeColorFromSH
1 | glm::vec3 pos = means[idx]; |
计算从相机位置 campos 到高斯球中心 pos 的向量,并将其归一化。这代表了当前的观察角度。
1 | glm::vec3* sh = ((glm::vec3*)shs) + idx * max_coeffs; |
获取系数指针
1 | result = SH_C0 * sh[0]; |
SH_C0, SH_C1, SH_C2 等是球谐函数的常数系数(如x,y,xx,xy等是与观察角度有关的函数(如sh存储了三维向量rgb的系数(即
排序
生成键值:duplicateWithKeys
1 | // 1. 获取当前线程处理的高斯球ID |
为什么这样排序有效?
因为在计算机排序中,数字是按位从高到低比较的:
- 先比高位:如果两个 Key 的高 32 位(Tile ID)不同,那么数值大的肯定属于不同的 Tile。这保证了同一个 Tile 的数据排在一起。
- 再比低位:如果高 32 位相同(同一个 Tile),计算机就会去比低 32 位(Depth)。这保证了深度顺序(从前到后或从后到前)。
排序:SortPairs
1 | cub::DeviceRadixSort::SortPairs( |
这里调用了 NVIDIA CUB 库
对于 64 位的键值
排序过程(从低位到高位):
- Pass 1:看
,分两桶(0和1)。 - Pass 2:看
,分两桶(保持上一轮顺序)。 - …
- 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 | // 1. 获取当前线程处理的位置索引 |
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 | float T = 1.0f; |
float T = 1.0f;
- 含义:透射率。表示光线穿过当前已累积的所有高斯球后,还剩下多少能量。
- 初始值:
(即 100%)。在开始绘制前,假设没有任何物体遮挡,光线完全透过。
uint32_t contributor = 0;
- 含义:当前贡献者计数器。它记录当前正在处理列表中的第几个高斯球。
uint32_t last_contributor = 0;
- 含义:最后有效贡献者索引。记录最后一个显著改变了该像素颜色的高斯球在列表中的位置。
- 用途:这个变量主要用于反向传播(训练阶段)。在计算梯度时,我们需要知道哪些高斯球参与了该像素的成像。如果
衰减到接近 0,后面的高斯球就不再处理了, last_contributor就标记了这个截止点。
float C[CHANNELS] = { 0 };
- 含义:累积颜色。
- 作用:存储当前像素计算出的最终 RGB 颜色值。
1 | int num_done = __syncthreads_count(done); |
done 变量在之前的循环中被置为 true,意味着该像素的透射率 num_done == BLOCK_SIZE,说明当前线程块内的所有线程(即该 Tile 内的所有像素)都已经完成了渲染(光线被完全遮挡)。
1 | int progress = i * BLOCK_SIZE + block.thread_rank(); |
progress:计算当前线程负责加载的高斯球在全局列表中的相对进度。
1 | if (range.x + progress < range.y) |
确保我们要访问的高斯球索引没有超出该线程块分配的总范围(range)。为了防止处理屏幕边缘或列表末尾时的越界访问。
1 | int coll_id = point_list[range.x + progress]; |
从排序好的 point_list 中取出高斯球的全局 ID (coll_id)。根据 ID,从全局显存中读取该高斯球的属性(xy 坐标、conic_opacity 参数)。
1 | float2 xy = collected_xy[j]; |
contributor:计数器加 1,记录当前处理到了第几个高斯球。- 数学原理:这是高斯函数的指数部分,
是像素点到高斯点的距离
其中con_o的.x, .y, .z存储了协方差逆矩阵的参数。 - 如果
power > 0,说明距离太远或者概率密度极低,该高斯球对此像素没有贡献,直接跳过。
1 | float alpha = min(0.99f, con_o.w * exp(power)); |
计算当前像素位置的高斯值,并乘以高斯球本身的不透明度(
con_o.w)。
min(0.99f, ...):强制限制最大为 0.99。这是为了防止数值溢出或完全遮挡导致的梯度消失问题(如果 ,则 ,后续梯度无法回传)。 如果计算出的
太小(小于 1/255),说明不透明度太低了,肉眼几乎不可见,直接跳过以节省算力。 计算如果加上当前这个高斯球,剩余的光线透射率
会变成多少。
如果小于 0.0001,说明光线已经被阻挡了 99.99% 以上,后面的高斯球无论是什么颜色都看不见了。done = true:标记该像素渲染完成,后续循环将不再处理该像素。
1 | for (int ch = 0; ch < CHANNELS; ch++) |
- 用Alpha Blending 公式累计颜色。当前高斯球贡献的颜色 = 球的颜色
球的不透明度 之前所有球的透射率。 features存储了高斯球的 RGB 颜色值。
1 | T = test_T; |
- 更新透射率,供下一个高斯球使用。记录最后一个对该像素有贡献的高斯球索引。
backward
我们有每个像素渲染出的颜色和真实的颜色,根据这些计算出loss,从而计算出
我们需要更新的是:高斯中心位置
预处理:preprocessCUDA
已知损失函数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 | auto idx = cg::this_grid().thread_rank(); |
为每个3D高斯分配一个线程,仅处理radii[idx] > 0,即屏幕上渲染半径大于0的高斯。不可见高斯对损失无贡献。
回顾正向投影过程:
展开为分量形式:
再通过透视除法得到2D归一化坐标:
具体计算过程:
1 | float3 m = means[idx]; |
m = means[idx]:取出当前3D高斯的均值。 m_hom = transformPoint4x4(m, proj):正向计算齐次裁剪空间坐标(对应上述数学公式)。 m_w = 1/(m_hom.w + eps):计算透视除法分母的倒数,加小常数 eps防止除零。mul1和mul2:
预计算和 ,用于简化后续商的导数计算。
1 | 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; |
已知
再根据商的导数法则(
代入投影矩阵的偏导数(
1 | dL_dmeans[idx] += dL_dmean; |
将投影路径的梯度累加到3D均值的总梯度中。3D均值还会通过“颜色路径”和“形状路径”影响损失,总梯度需累加所有路径的贡献。
颜色路径的反向传播:computeColorFromSH
已知损失对RGB颜色的梯度 dL_dcolor[idx]),也即
- 损失对SH系数的梯度
(代码中 dL_dshs), - 损失对3D高斯位置的梯度
(代码中 dL_dmeans)。
计算对SH系数的梯度:
1 | glm::vec3 pos = means[idx]; |
计算视角方向 dir_orig),并归一化得到单位方向 dir)。
由于RGB是SH系数的线性函数,因此
代码中按阶数
阶数
1 | float dRGBdsh0 = SH_C0; |
阶数
其中
1 | float dRGBdsh1 = -SH_C1 * y; |
阶数
代码逻辑与 dL_dRGB 得到SH系数梯度。
计算对均值的梯度:
其中
1 | glm::vec3 dRGBdx(0, 0, 0); // ∂RGB / ∂dir.x |
在 1 阶、2 阶、3 阶 SH 里不断 累加:
1 | // 1阶 SH |
d是从相机指向高斯的单位向量:
对 pos 求导:
其中
代码的巧妙之处在于,它将这两步合并了。
在第一步中计算的那些包含 x, y, z 的导数表达式(如 -SH_C1 * sh[3]),实际上已经是
例如,对于一阶球谐函数,其基函数与方向分量成正比(如
而
形状路径的反向传播(缩放+旋转→协方差)
1 | if (scales) |
数学背景:3D高斯的协方差矩阵由缩放和旋转合成。正向过程为:
- 将四元数
rotations转换为旋转矩阵。 - 构建缩放对角矩阵
(对角元素为 scales * scale_modifier)。 - 计算3D协方差:
。
- 将四元数
反向传播任务:
已知( dL_dcov3D),通过链式法则计算:
1.( dL_dscale):缩放参数的梯度。
2.( dL_drot):旋转四元数的梯度。
我们已知损失函数 dL_dSigma)。
我们需要求
根据矩阵求导法则,若
代码对应:
1 | // dSigma_dM = 2 * M |
反向传播需要用到前向传播时的中间结果,所以代码首先重新计算了旋转矩阵
我们需要求
因为
这本质上是一个矩阵乘法:
1 | glm::mat3 Rt = glm::transpose(R); |
接下来我们需要求
路径是:
首先,利用
代码中通过 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 }; |