编程 Gaussian Splatting Lightning 训练笔记:多 GPU 要在密化之后开,默认不出 ply

2026-09-24 00:04:58

gaussian-splatting-lightning 训练笔记:安装、OOM 规避与多 GPU DDP

项目信息

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 SAM sam_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。

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

推荐文章

程序员茄子在线接单