gaussian-splatting-lightning 训练笔记:安装、OOM 规避与多 GPU DDP
项目信息
- 仓库:yzslab/gaussian-splatting-lightning
- Changelog:releases
- 光栅化后端:nerfstudio-project/gsplat
- 3DGS-MCMC:ubc-vision/3dgs-mcmc
- Mip-Splatting:autonomousvision/mip-splatting
- 2DGS:hbb1/2d-gaussian-splatting
- SAGA / Segment Any 3D Gaussians:Jumpat/SegAnyGAussians
- Deformable Gaussians:ingra14m/Deformable-3D-Gaussians
- Feature 3DGS:ShijieZhou-UCLA/feature-3dgs
News
- 2025-06-29:论文 “Robust and Efficient 3D Gaussian Splatting for Urban Scene Reconstruction” 被 ICCV 2025 接收。
已知问题
- 多 GPU 训练原先只能在密化完成后启用;2.16 新多 GPU 策略可带密化。
能力范围
- 多 GPU/Node;diff-gaussian-rasterization 与 gsplat 可切换。
- 数据集:Blender (nerf_synthetic)、Colmap、PolyCam、Nerfies、NSVF(仅 Synthetic)、MatrixCity、PhotoTourism。
- Web viewer 支持多模型、transform、场景编辑、视频相机路径编辑;另有视频渲染、大图防 OOM、动态物体 mask。
- 派生算法通过
--config切换。Mip-Splatting、LightGaussian、2DGS、SAGA、Appearance/In the wild、3DGS-MCMC、Feature distillation、新多 GPU 策略等按对应依赖或 config 启用;其余按configs/选择。
1. 安装
1.1 克隆
git clone https://github.com/yzslab/gaussian-splatting-lightning.git
cd gaussian-splatting-lightning
1.2 虚拟环境
conda create -yn gspl python=3.9 pip
conda activate gspl
1.3 PyTorch
测试过 PyTorch==2.0.1,必须匹配 nvcc --version。CUDA 11.8:
pip install -r requirements/pyt201_cu118.txt
1.4 依赖
pip install -r requirements.txt
1.5 可选包
- ffmpeg:
sudo apt install -y ffmpeg - gsplat:只支持作者改过的 v1。
pip uninstall -y gsplat
pip install -r requirements/gsplat.txt
- SegAnyGaussian:需要 gsplat、SAM (
requirements/sam.txt)、facebookresearch/pytorch3d,并把 ViT-H SAMsam_vit_h_4b8939.pth下载到仓库根目录。
2. 训练
2.1 基本命令
python main.py fit --data.path DATASET_PATH -n EXPERIMENT_NAME
可自动检测部分数据集类型,也可指定 --data.parser。取值:Colmap、Blender、NSVF、Nerfies、MatrixCity、PhotoTourism、SegAnyColmap、Feature3DGSColmap。
默认只产 checkpoint。需要 vanilla 3DGS 格式 ply:
python utils/ckpt2ply.py TRAINING_OUTPUT_PATH- 或训练时加
--model.save_ply true
2.2 常用选项
- 带 web viewer:
--viewer。 - Blender 数据集推荐
configs/blender.yaml。 - Colmap mask:单通道,0/黑表示 masked pixel,文件名必须是图片名 +
.png;需 undistort 时用utils/colmap_undistort_mask.py。训练加--data.parser Colmap --data.parser.mask_dir MASK_DIR_PATH。 - 下采样图像(Colmap):
python utils/image_downsample.py PATH_TO_DIRECTORY_THAT_STORE_IMAGES --factor 4,训练加--data.parser.down_sample_factor 4;取整模式--data.parser.down_sample_rounding_mode可取 floor、round、round_half_up、ceil,默认 round。 - 大图防 OOM:
--data.image_uint8 true用 uint8 缓存图像。--data.train_max_num_images_to_cache 512 --data.async_caching true,训练时缓存下一 batch;或--data.train_max_num_images_to_cache 1024,当前 batch 结束时缓存下一 batch。
- 加速:
- 全部图像放 GPU:
--data.image_on_cpu false,配合--data.image_uint8 true降显存。 - 避开每 epoch 验证:
--trainer.check_val_every_n_epoch 99999。注意val/指标训练中不会更新。 - 进一步加速见 Taming 3DGS。
- 全部图像放 GPU:
2.3 gsplat
python main.py fit --config configs/gsplat.yaml ...
2.4 多 GPU DDP
DDP 只能在密化完成后启用:先单卡训练到密化结束并存 checkpoint,再 resume 开多 GPU。更多 GPU 会改善 PSNR/SSIM。
# 先单卡
python main.py fit --config ... --data.path DATASET_PATH --model.density.densify_until_iter 15000 --max_steps 15000
# 再 resume + DDP
python main.py fit --config ... --trainer configs/ddp.yaml --data.path DATASET_PATH --max_steps 30000 --ckpt_path last
2.5 Deformable 3D Gaussians
python main.py fit --config configs/deformable_blender.yaml --data.path ...
2.6 Mip-Splatting
训练:python main.py fit --config configs/mip_splatting_gsplat_v2.yaml --data.path ...
融合 3D smoothing filter:python utils/fuse_mip_filter.py TRAINED_MODEL_DIR
2.7 LightGaussian
目前只有 Prune & finetune。训练+密化+剪枝:... fit --config configs/light_gaussian/train_densify_prune-gsplat.yaml --data.path ...;剪枝+微调:... fit --config configs/light_gaussian/prune_finetune-gsplat.yaml --data.path ... --ckpt_path YOUR_CHECKPOINT_PATH,需保证 hparams 与输入模型一致。
2.8 AbsGS / EfficientGS
... fit --config configs/gsplat-absgrad.yaml --data.path ...
2.9 2D Gaussian Splatting
先安装 diff-surfel-rasterization:pip install -r requirements/2DGS.txt。
训练:... fit --config configs/vanilla_2dgs.yaml --data.path ...
Mesh extraction:Bounded 用 python utils/gs2d_mesh_extraction.py MODEL_OUTPUT_PATH;Unbounded 加 --unbounded true。
2.10 Segment Any 3D Gaussians
先训练 3DGS 场景:python main.py fit --config configs/gsplat.yaml --data.path data/Truck -n Truck -v gsplat。
生成 SAM masks 和 scales:
python utils/get_sam_masks.py data/Truck/images
python utils/get_sam_mask_scales.py outputs/Truck/gsplat
结果保存在 data/Truck/semantics。训练 SegAnyGS:python seganygs.py fit --config configs/segany_splatting.yaml --data.path data/Truck --model.initialize_from outputs/Truck/gsplat -n Truck -v seganygs。分割/聚类:python viewer.py outputs/Truck/seganygs。
2.12 Appearance Model
图像外观差异大时使用,例如不同曝光、白平衡、对比度甚至昼夜。实现上给每个 3D Gaussian 额外 feature vector,给每个 appearance group 一个 embedding,两者输入轻量 MLP 计算颜色。细节在 internal/renderers/gsplat_appearance_embedding_renderer.py。
- 生成 appearance groups(Colmap 或 PhotoTourism):
python utils/generate_image_apperance_groups.py PATH_TO_DATASET_DIR --image --name appearance_image_dedicated - 训练:
python main.py fit --config configs/appearance_embedding_renderer/view_dependent.yaml --data.path PATH_TO_DATASET_DIR --data.parser Colmap --data.parser.appearance_groups appearance_image_dedicated - 其他 configs:
view_independent.yaml(关 view dependent)、sh_view_dependent.yaml(用 SH 表示 view dependent)、*-distributed.yaml(多 GPU)、*-estimated_depth_reg.yaml。 - 渲染时去掉 MLP 依赖:
python utils/fuse_appearance_embeddings_into_shs_dc.py TRAINED_MODEL_DIR
2.13 3DGS-MCMC
... fit --config configs/gsplat-mcmc.yaml --model.density.cap_max MAX_NUM_GAUSSIANS ...
MAX_NUM_GAUSSIANS 是使用的最大 Gaussian 数量。参考 ubc-vision/3dgs-mcmc。
2.14 Feature distillation
来自 Feature 3DGS,这里用两阶段优化而非联合。先 gsplat 训练,再抽特征图(SAM:python utils/get_sam_embeddings.py data/Truck/images;LSeg 用 ShijieZhou-UCLA/feature-3dgs),最后蒸馏:python main.py fit --config configs/feature_3dgs/sam-speedup.yaml --data.path data/Truck --data.parser.down_sample_factor 2 --model.initialize_from outputs/Truck/gsplat -n Truck -v feature_3dgs-sam
高维特征光栅化慢,用 --data.parser.down_sample_factor 缩小渲染特征图加速。完成后 viewer 可视化:python viewer.py outputs/Truck/feature_3dgs。
2.15 In the wild
基于 Appearance Model,为每个训练视角生成 visibility map,判断像素是否属于 transient objects。思路类似 Ha-NeRF,但用 2D dense grid encoding 加速训练。注意:能区分 transient 像素,但可能无法去除 transient 的 artifacts/floaters,也可能把欠重建区域当 transient。
- 需要 tiny-cuda-nn:
pip install -r requirements/tcnn.txt - 下载 PhotoTourism 数据集和 split 文件;split 文件与 dense 目录同路径。
- 训练:
python main.py fit --config configs/appearance_embedding_visibility_map_renderer/view_independent-2x_ds.yaml --data.path data/brandenburg_gate -n brandenburg_gate - 训练集验证:
python main.py validate --config outputs/brandenburg_gate/lightning_logs/version_0/config.yaml --save_val --val_train
2.16 新多 GPU 训练策略
类似简化版 Scaling Up 3DGS。Gaussians 分布式存储、投影、算色,每张 GPU 为不同相机光栅化整图;目前没有 pixel-wise distribution。可带密化。
注意:尚未充分验证,仍在开发;目前仅多 GPU;与含神经网络的 derived algorithms 结合时需手动 DDP 包裹网络。
- 训练:
python main.py fit --config configs/distributed.yaml ...
默认每个进程内存中持有冗余数据集副本,可能 CPU OOM。加--data.distributed true让各进程加载不同子集。 - 合并 checkpoint:
python utils/merge_distributed_ckpts.py outputs/TRAINED_MODEL_DIR - viewer:
python viewer.py outputs/TRAINED_MODEL_DIR/checkpoints/MERGED_CHECKPOINT_FILE
2.17 SpotLessSplats
注意:没有 utilization-based pruning 和 appearance modeling。
4. Web Viewer
Web viewer 可加载多个模型、启用 transform、编辑场景、编辑视频相机路径,也可加载其他实现训练的模型(如 2D Gaussian Splatting、4D Gaussian)。训练时用 --viewer 可直接带 viewer 跑;已有模型用 python viewer.py MODEL_OUTPUT_PATH。