gsplat 移植 G-SHARP v0.2:可变形 4D 高斯场景重建的架构设计与实现解析
【免费下载链接】gsplatCUDA accelerated rasterization of gaussian splatting项目地址: https://gitcode.com/GitHub_Trending/gs/gsplat
本文基于 gsplat 仓库中的集成提案 G-SHARP v0.2 Integration Proposal,拆解这套将手术场景重建(G-SHARP v0.2)训练期算法原生移植进 gsplat 的完整方案:核心库新增的深度/遮挡/初始化组件、实验性的gsplat.contrib.dynamic4D 模块(HexPlane 场 + 形变网络 + 动态策略),以及配套的 EndoNeRF 示例训练器、测试矩阵与风险对策。读完本文,你可以理解 gsplat 如何通过"核心稳定层 + contrib 实验层"的分层设计,在gsplat.rasterization没有时间轴的前提下支持可变形动态高斯场景重建。
一、提案背景:目标与非目标
G-SHARP v0.2 是 NVIDIA holohub 中applications/surgical_scene_recon/应用(手术场景重建)的训练算法栈。该提案将其训练期算法贡献移植进 gsplat,但明确区分了"要做的"和"不做的",这是理解整个移植边界的关键。
目标(Goals):
- 将 G-SHARP v0.2 全部训练期算法贡献原生移植进 gsplat;
- 可变形 / 4D 高斯机制全部收敛到明确标注为实验性的
gsplat/contrib/命名空间下; - Test-first:每个移植单元都有正例 + 反例 pytest 用例;
- 提供 EndoNeRF 样本上的可运行端到端教程。
非目标(Non-goals):
- gsplat 包内不引入任何 Holoscan / HoloHub 依赖;
- Depth Anything V2 / MedSAM3 / VGGT 这类外部预训练网络不进入
gsplat/包,教程仅把它们作为可选的上游预处理引用; - 不做压缩 / 率失真工作(G-SHARP v0.2 本身也没有);
- 不做相机位姿优化(G-SHARP v0.2 同样不做)。
提案文档头部同时标注了状态:实现位于分支vnath_gsharp,落在gsplat.contrib.dynamic下,实验性,API 可能变化。
二、组件映射总览:核心稳定层 vs contrib 实验层
提案给出的组件映射是整个移植方案的骨架。它把 G-SHARP 的每个训练组件映射到一个 gsplat 目标文件,并注明了 G-SHARP 源出处。以下按"核心(稳定)新增"与"实验性 contrib 新增"两类完整列出。
核心(稳定)新增
| 组件 | gsplat 中的落地文件 | G-SHARP v0.2 源 |
|---|---|---|
| 双目视差 L1 损失 | gsplat/losses.py(与既有深度损失并列) | EndoRunner.compute_depth_loss,training/gsplat_train.py906–960 行 |
| 单目 Pearson 深度损失 | gsplat/losses.py | 同上 |
| 带掩码的 L1 / SSIM 包装器 | gsplat/losses.py | training/gsplat_train.py1059–1101 行 |
| 遮挡 TV(总变分)正则 | gsplat/regularizers.py | compute_tv_loss_targeted,1925–2028 行 |
| Torch 掩码膨胀(max-pool 替代 OpenCV dilation) | gsplat/regularizers.py | OpenCV dilation 替换 |
| 不可见区域掩码构造器 | gsplat/regularizers.py | create_invisible_mask_from_paths |
| 多帧深度反投影 | gsplat/init_utils.py | accumulate_multiframe_pointcloud,2036–2159 行 |
| KNN 尺度初始化 | gsplat/init_utils.py | create_splats_with_optimizers,1796–1906 行 |
| 两阶段调度器 | gsplat/training/schedulers.py | EndoRunner._train_stage,1007–1050 行 |
实验性 contrib 新增(gsplat.contrib.dynamic)
| 组件 | gsplat 中的落地文件 | G-SHARP v0.2 源 |
|---|---|---|
| HexPlane 时空特征场 | hexplane.py | training/scene/hexplane.py |
| 形变 MLP + 每高斯形变表 | deformation.py | training/scene/deformation.py,gsplat_train.py1376–1664 行 |
| 平面 / 时间平滑正则 | regulation.py | training/scene/regulation.py,训练循环 1233–1272 行 |
| DynamicStrategy | strategy.py | rasterize_splats,864–896 行 |
示例与文档交付物
| 交付物 | 路径 |
|---|---|
| EndoNeRF / SCARED 数据集加载器 | endonerf.py |
| 动态场景训练器 recipe | dynamic_surgical_trainer.py |
| 教程 | dynamic_surgical.rst |
| Contrib API 文档 | contrib.rst |
对照当前仓库源码,上述映射已全部落地:gsplat/losses.py中确有binocular_disparity_l1、pearson_depth_loss、masked_l1、masked_ssim;gsplat/regularizers.py提供compute_tv_loss_targeted、dilate_mask、create_invisible_mask(即提案中"不可见掩码构造器"在 gsplat 中的最终命名,对应 G-SHARP 的create_invisible_mask_from_paths);gsplat/training/schedulers.py提供TwoStageScheduler。唯一的细微差异是数据集加载器:教程 dynamic_surgical.rst 说明目前只实现了 EndoNeRF 解析器,SCARED 尚为占位(stub)。
三、为什么这样分层:设计 rationale
提案的 "Why this layout" 一节给出了四条分层原则,这些决策直接决定了 gsplat 的模块归属:
- 核心库保持纯 splatting 数学。深度损失、遮挡 TV、多帧初始化、两阶段调度都是通用、低风险的组件,因此放在既有的
losses.py、rendering.py、strategy/旁边,而不新建顶层模块。 - 4D / 可变形部分引入新的公共 API 面(
HexPlane、DeformNetwork、DynamicStrategy),仍属研究级,放进gsplat/contrib/dynamic/。这个命名沿用torchvision.prototype的约定,语义是"随 wheel 一起发布,但 API 可能变"——这一点在 contrib 包文档 顶部有正式的 warning 标注。 - 领域特化的数据集胶水代码不进库。EndoNeRF/SCARED 加载器与
colmap.py/ncore.py一起放在examples/datasets/下。 - 外部网络(DA2 / MedSAM3 / VGGT)不进树。教程只解释"想办法拿到
poses_bounds.npy、masks/、depth/",并指向未移植的 holohub 目录作为参考实现。从源码结构看,这保证了gsplat包的依赖面不随医学图像网络膨胀。
四、核心层组件的源码级剖析
4.1 深度监督:视差 L1、Pearson 与带掩码包装器
提案将双目 / 单目深度损失一并收入 losses.py,与既有的depth_l1_loss并列。
binocular_disparity_l1(losses.py#L227):在逆深度(视差)空间做 L1。实现上有一个值得注意的细节——只有当预测和真值深度都大于eps(默认1e-7)时该像素才参与损失(pair_valid = valid_pred & valid_gt)。源码注释解释了原因:若只做单边1/x ↔ 0替换,单边无效的深度会把|0 - 1/other|泄漏进损失。可选mask参数用于剔除器械/动态物体区域。pearson_depth_loss(losses.py#L279):单目分支的1 - Pearson r,对(掩码后的)展平深度对计算。两个数值稳定性处理:有效样本少于 2 时返回可微的 0;任一侧方差接近 0 时分母钳位到1e-12,避免 NaN。masked_l1(losses.py#L328):只在mask != 0区域求均值,掩码全零时返回可微的 0 而非 NaN;mask 可广播(如[B, 1, H, W]作用于[B, C, H, W])。masked_ssim(losses.py#L360):先把 pred 和 gt都乘上 mask 再算 SSIM,这与 G-SHARP 训练循环的约定一致。文档也诚实指出:均值是对全图取的,因此损失量级会随掩码覆盖率变化——稀疏掩码会按被剔除比例稀释报告的损失,这一偏差是有意保留的上游约定。
4.2 遮挡感知正则:遮挡 TV、掩码膨胀、不可见区域
regularizers.py 顶部 docstring 明确写着"Regularizers ported from G-SHARP v0.2",三个函数各司其职:
compute_tv_loss_targeted(regularizers.py#L53):各向异性总变分损失。无 mask 时对image.numel()取平均;带 mask 时分别对水平/垂直差分乘以裁剪到差分形状的 mask 求和,再除以mask.sum() * C + 1e-8,与 G-SHARPtraining/gsplat_train.py1992–2028 行行为一致。实现上还有一个性能细节:mask 的二值性契约检查(ENFORCE_CONTRACTS)被环境变量GSPLAT_ENFORCE_CONTRACTS=1门控,因为.all()会触发 CPU-GPU 同步,而该损失每个训练步都会被调用。dilate_mask(regularizers.py#L107):用F.max_pool2d(stride=1、padding=kernel//2)实现纯 torch 的二值掩码膨胀,替代 G-SHARP 的cv2.dilate——这是"核心库不依赖 OpenCV"这一决策的直接体现。支持 2D/3D/4D 掩码,kernel 必须是正奇数。create_invisible_mask(regularizers.py#L155):跨所有帧求并集——任何一帧中曾出现1(器械/遮挡物)的像素在结果中标1,即"在整个数据集里曾被器械遮挡的区域"。它支持两种输入模式:Tensor按已是 tool=1 约定处理;str路径则用 PIL 读入、归一化到[0,1]并取反(1 - mask),对应 G-SHARP 磁盘上 tissue=1 的 PNG 存储约定。
4.3 免 SfM 的初始化:多帧深度反投影 + KNN 尺度
G-SHARP 的accumulate_multiframe_pointcloud(2036–2159 行)在 gsplat 中落地为 init_utils.py 的两个纯 torch 函数:
multi_frame_depth_unprojection(init_utils.py#L40):对每帧取mask != 0 且 depth > 0的有效像素,用逐帧针孔内参反投影到相机系,再经 camera-to-world 位姿变换到世界系,所有帧点云拼接后按max_points随机下采样。uint8图像会自动归一化到[0,1];若无任何像素通过门控,返回空张量而非报错。knn_scale_init(init_utils.py#L145):基于每点到 k 近邻(默认k=3,不含自身)距离的均方根,经clamp_min(eps).log()得到每点初始 log-scale。它回应了提案风险清单里"sklearn 可选依赖"一条——纯 torch 实现(cdist+topk)意味着 KNN 初始化不需要 sklearn,依赖列表保持不变。实现上采用分块成对距离扫描(默认chunk_size=1024),把峰值内存从O(N^2)(N=50000 时约 10 GB)降到O(chunk * N),让 12 GB 级消费卡也能跑。
4.4 两阶段调度器:粗静态 → 细动态
G-SHARP 的EndoRunner._train_stage在 gsplat 中落地为 schedulers.py 的TwoStageScheduler(schedulers.py#L53):
- Coarse 阶段(
global_step < coarse_steps):锁定单个固定帧(默认coarse_frame_index=0,对应 G-SHARP 的self.trainset[0]),shuffle=False。目的是在时间形变介入前,用单一视角把静态高斯预热好; - Fine 阶段(
global_step >= coarse_steps):shuffle=True,帧索引按(global_step - coarse_steps) % num_frames确定性轮转,保证每帧都被访问。
step()返回TwoStageScheduleStep数据类(stage/frame_index/shuffle三字段)。值得注意的设计约束:fine_steps只是给调用方的训练预算提示,调度器本身不强制执行,从而保持为纯无状态映射。参数校验(global_step >= 0、num_frames > 0、coarse_frame_index在界内)都惰性地在step()调用时进行,因为构造时还拿不到num_frames。
五、gsplat.contrib.dynamic:4D 可变形高斯三件套
contrib/dynamic/init.py 的 docstring 声明该包"Ported from G-SHARP v0.2'straining/scenepackage",公共 API 为HexPlaneField、DeformNetwork、DynamicStrategy与一组正则函数。
5.1 HexPlaneField:4D 特征场的六平面分解
HexPlaneField(hexplane.py#L159)是对(x, y, z, t)四维特征场的多分辨率 6 平面分解,源自 K-Planes / 4DGaussians 公式(文件头注明参考了 Nerfstudio 的 HexPlane 实现)。机制如下:
- 四个输入轴两两组合出 6 个 2D 特征平面:
xy, xz, xt, yz, yt, zt; - 每个平面用
F.grid_sample双线性采样(padding_mode="border"),各平面逐元素相乘得到该尺度的特征向量; - 多个 multi-res 尺度的特征向量拼接(
concat_features=True是当前唯一支持的模式),最终特征维度为output_coordinate_dim * len(multires)。
关键实现细节(hexplane.py#L85-L111):
- 默认配置
_DEFAULT_PLANE_CONFIG:grid_dimensions=2、input_coordinate_dim=4、output_coordinate_dim=32、resolution=[64, 64, 64, 25](时间轴分辨率明显低于空间轴,默认multires=(1, 2)只放大空间轴分辨率,时间分辨率跨尺度保持不变); - 时间平面初始化为一(
nn.init.ones_),其余平面U(0.1, 0.5)均匀初始化。全 1 的时间平面意味着初始时变形近似恒等映射,与 G-SHARP 约定一致; - 空间坐标经 AABB 归一化到
[-1, 1],半宽bounds默认1.6(对齐 G-SHARP),时间坐标原样透传,越界值靠 border padding 钳位; - 类上提供了
spatial_planes()(组合索引(0, 1, 3)即 xy/xz/yz)与temporal_planes()((2, 4, 5)即 xt/yt/zt)两个访问器,把空间/时间平面的划分逻辑内聚在场对象上——这样正则器不必硬编码索引。
5.2 DeformNetwork 与 DeformationTable
deformation.py 提供两个类:
DeformNetwork(deformation.py#L49):一个num_layers深(默认 3)的 ReLU 主干,输入 HexPlane 特征(feature_dim必须与产生它的HexPlaneField.feat_dim匹配,默认隐层宽 64),后接三个线性头——位置增量 3 维、四元数增量 4 维、不透明度增量 1 维。三个头零初始化,因此构造时的前向传播是恒等映射(该性质由测试test_deform_net_zero_init_is_identity锁定);梯度仍能流过头部,所以主干可学。t参数为未来的时间感知扩展保留,当前实现假定时间信息已经由 HexPlane 编码进plane_features。前向还有三重契约检查:batch 维一致、特征维匹配、四个张量 dtype 一致,任一违反即抛ValueError。DeformationTable(deformation.py#L165):每高斯一个 bool 标志,标记该高斯是否由形变网驱动。提供set_indices/prune/duplicate/split与DefaultStrategy的稠密化算子保持同步缩放,split默认factor=2(对齐 gsplat 约定),子代继承父代的动态标志。它本身是纯torch.bool张量,不进优化器状态、零额外开销。
5.3 HexPlane 正则:平面平滑、时间平滑与 time-L1
regulation.py 移植了 G-SHARPtraining/scene/regulation.py的三个正则器,均作用于(B, C, H, W)平面张量列表:
plane_smoothness/time_smoothness:共享同一套数学——沿 H 轴的二阶差分平方和(每张平面内部取均值,跨平面求和,H < 3的平面自动跳过)。区别只在于调用者传入哪些平面:空间平面[xy, xz, yz](组合索引[0, 1, 3])做空间平滑;时空平面[xt, yt, zt]([2, 4, 5])中 H 轴即时间轴(源自_init_grid_param的反序布局),二阶差分平方即为时间平滑;time_l1:时空平面与 1.0 的 L1 偏差均值。因为时间平面初始化为 1(恒等形变),该正则惩罚偏离,鼓励静态区域保持"无时间形变",对应 G-SHARP 的L1TimePlanes;hexplane_regularization:官方推荐的便捷入口,直接接收HexPlaneField实例,用其自身的空间/时间访问器完成平面划分,再做三个正则的加权和(lambda_plane_smooth/lambda_time_smooth/lambda_time_l1均默认 1.0)。包级__init__.py的注释明确建议优先使用它而非手写平面列表调用三个底层函数——"手工划分的列表会在 HexPlaneField 重构时漂移"。
5.4 DynamicStrategy:与稠密化同频的动态掩码
DynamicStrategy(strategy.py#L50)继承自gsplat.strategy.DefaultStrategy,是"可变形感知的稠密化/剪枝策略"。它的核心设计直接回应了提案风险清单里的第一条——gsplat.rasterization没有时间轴:
- 形变传递本身(HexPlane → DeformNetwork →
(means, quats, opacities))不在这个策略类里,而是放在训练器中、在调用rasterization(...)之前执行,与 G-SHARP 的rasterize_splats同构。策略类只管稠密化策略 + 形变表的簿记; initialize_state()扩展父类,往state里塞入一个普通的(num_gaussians,)bool 张量dynamic_mask(init_dynamic=True时全 True,即所有高斯都过形变网;静态优先的工作流可设 False 再手动翻转)。之所以用普通张量而非DeformationTable包装器,是因为 gsplat 的稠密化算子会遍历 state 里的每个张量并施加与 params 相同的 per-Gaussian 置换(见 strategy/ops.py),从而让 mask 与params["means"]自动同步缩放、split 时保持身份继承;step_post_backward只负责在dynamic_mask缺失时抛RuntimeError(提示先调用initialize_state),其余全部委托父类;- HexPlane 和 DeformNet 的可训练参数不进
params:gsplat 的稠密化算子会盲目地按 per-Gaussian 索引对params中每个条目做 split/duplicate/prune,而非 per-Gaussian 张量(平面网格、MLP 权重)会被越界索引。示例训练器因此为它们单独建优化器(examples/dynamic_surgical_trainer.py中的build_deform_modules,trainer#L296)。类 docstring 还留了一段历史注记:早期版本把DeformationTable存在state["deformation_table"]里靠自定义 hook 缩放,该 hook 不能在 split 中保持身份、且训练器从未真正消费过它;现在规范 mask 就是 state 里的普通张量,包装类仅为向后兼容保留、且不进__all__。
六、端到端示例:EndoNeRF 动态手术场景训练
6.1 数据布局与前置产物
教程 dynamic_surgical.rst 以 EndoNeRF "pulling" 序列为目标场景(SCARED 支持尚未实现)。磁盘布局约定为:
<scene>/ poses_bounds.npy images/ 000000.png 000001.png ... depth/ 000000.png 000001.png ... masks/ 000000.png 000001.png ...教程明确说明:gsplat 本身不提供深度估计、器械掩码或相机位姿的代码。G-SHARP v0.2 的 holohub 应用用 Depth Anything V2、MedSAM3 和 VGGT-1B 作为可选上游预处理栈,可参考holohub/applications/surgical_scene_recon;这与提案非目标第 2 条完全对应。数据集加载器在 endonerf.py,与colmap.py/ncore.py并排放在examples/datasets/。
6.2 快速上手命令
教程给出的 quickstart(可直接复制运行):
python examples/dynamic_surgical_trainer.py \ --data_dir path/to/endonerf_pulling \ --output_dir output/dynamic_surgical \ --coarse_steps 200 --fine_steps 2500 \ --init_max_points 50000 \ --depth_mode binocular \ --render_gif_after_train参数要点:
--coarse_steps/--fine_steps直接喂给第四节讲的TwoStageScheduler(粗 200 步单帧预热 + 细 2500 步轮转);--init_max_points对应multi_frame_depth_unprojection的max_points下采样上限;--depth_mode binocular选择视差 L1 深度监督(另一分支为单目 Pearson);- 训练器会从初始化点云自动推导HexPlane 的 AABB,
--hex_bounds仅作为下界;完整 flag 列表由 trainer#L770 的tyro.cli(Config)生成,可用--help查看(Config是 trainer#L90 的 tyro dataclass)。
API 文档由 contrib.rst 通过 automodule 自动从gsplat.contrib.dynamic的四个模块生成,并在顶部以 warning 形式重申"gsplat.contrib 全为实验性"。
七、测试计划与风险对策
测试矩阵:test-first 的落地
提案要求"每个新模块都有正例 + 反例 pytest 用例,遵循仓库既有的平铺测试布局与 pytest 约定(conftest.py、seed 42、CUDA 感知)"。当前仓库tests/下的对应文件(相对仓库根目录):
| 移植组件 | 测试文件 |
|---|---|
| HexPlane 场 | test_contrib_hexplane.py |
| 形变网络 | test_contrib_deformnet.py |
| DynamicStrategy | test_contrib_dynamic_strategy.py |
| HexPlane 正则 | test_contrib_regulation.py |
| 双目/Pearson 深度损失 | test_losses_depth.py |
| 遮挡 TV / 膨胀 / 不可见掩码 | test_regularizers_occlusion.py |
| 多帧初始化 | test_init_multiframe.py |
| 两阶段调度器 | test_two_stage_scheduler.py |
| EndoNeRF 数据加载 | test_dataset_endonerf.py |
| 端到端训练器 | test_dynamic_surgical_trainer.py |
提案中提到的组件 × 测试用例矩阵位于planning/gsharp_gsplat_plan.md,文档特意说明该文件"有意不入库"(可本地渲染 HTML 预览),因此上表按仓库实际测试文件整理。
风险清单与源码级应对
提案的 Risks 一节列了四条风险,逐条都能在源码中找到对应解法:
- 形变与
rasterization()的耦合——gsplat 光栅化没有时间轴。解法即 5.4 所述:DynamicStrategy模式下的形变在rasterization(...)调用之前对(means, quats, opacities)生效,与 G-SHARP 的rasterize_splats同构; - CUDA-only 测试——HexPlane/DeformNet 的单元测试在 CPU 上跑以保证确定性,GPU 冒烟测试标记为 CUDA-only(与提案约定一致);
- EndoNeRF 样本的许可证——在教程中说明,不在树内再分发;
- sklearn 可选依赖——
knn_scale_init用纯 torch 实现(cdist+topk分块扫描),保持必需依赖列表不变。
八、小结
这份移植方案的精髓在于边界清晰:通用、低风险的训练期组件(深度损失、遮挡 TV、多帧初始化、两阶段调度)沉入核心库并复用losses.py/regularizers.py/init_utils.py/training/的既有位置;研究级的 4D 机制(HexPlane、DeformNetwork、DynamicStrategy)隔离在gsplat.contrib.dynamic实验命名空间,并靠"形变在光栅化前应用 + state 中普通 bool 掩码随稠密化算子自动缩放 + 非 per-Gaussian 参数独立优化器"三个工程决策化解了rasterization无时间轴的结构性约束;外部依赖(医学深度/分割/位姿网络、sklearn、OpenCV、HoloHub)全部被纯 torch 或"不进树"策略挡在包边界之外。对想在 gsplat 上搭动态/可变形高斯管线的开发者,examples/dynamic_surgical_trainer.py与docs/source/examples/dynamic_surgical.rst是可直接运行的起点,同时需要留意 contrib 层的实验性定位——API 可能随版本变化。
【免费下载链接】gsplatCUDA accelerated rasterization of gaussian splatting项目地址: https://gitcode.com/GitHub_Trending/gs/gsplat
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考