mirror of
https://github.com/harry7557558/spirula-studio.git
synced 2026-10-02 02:44:54 +08:00
rename and i18n plan, phase 0
This commit is contained in:
+1
-3
@@ -1,8 +1,6 @@
|
||||
# Third-party dependencies (miniz, stb, npy, ...)
|
||||
src/external/** linguist-vendored
|
||||
|
||||
# Auto-generated code (generate_headers.py, generate_kernel_instantiation.py,
|
||||
# generate_cli_config.py)
|
||||
# Auto-generated code (generate_headers.py, generate_kernel_instantiation.py)
|
||||
src/generated/** linguist-vendored linguist-generated
|
||||
src/app/generated/** linguist-vendored linguist-generated
|
||||
src/instantiations/** linguist-vendored linguist-generated
|
||||
|
||||
@@ -123,7 +123,9 @@ src/
|
||||
│ ├── bind_data.cpp native dataset parsers
|
||||
│ ├── bind_viewer.cpp native web-viewer server + post-split bake
|
||||
│ └── bind_trainer.cpp SsplatConfig + TrainerSession
|
||||
├── generated/ app/generated/ instantiations/ GENERATED — do not hand-edit
|
||||
├── config/ TrainConfig.h — the training config's single source
|
||||
│ of truth: one X-macro row per flag, hand-written
|
||||
├── generated/ instantiations/ GENERATED — do not hand-edit
|
||||
└── external/ vendored (miniz, stb, npy)
|
||||
```
|
||||
|
||||
@@ -168,14 +170,15 @@ the comment on the option in `cmake/SsplatOptions.cmake` before changing it.
|
||||
## Codegen — the invariants that bite
|
||||
|
||||
Generated trees are marked in `.gitattributes` (`src/generated/`,
|
||||
`src/app/generated/`, `src/instantiations/`) and are **committed**, so a fresh checkout
|
||||
builds with no Python at all. Five generators, all run from the repo root:
|
||||
`src/instantiations/`) and are **committed**, so a fresh checkout
|
||||
builds with no Python at all. Four generators, all run from the repo root.
|
||||
Every one of them reads C++/CUDA sources — **no generator reads Python**, and
|
||||
none should be added:
|
||||
|
||||
| generator | reads | writes |
|
||||
|---|---|---|
|
||||
| `tools/codegen/generate_headers.py` | `/*[AutoHeaderGeneratorExport]*/` markers in `src/kernels/**/*.cu` | the declaration section of the matching `<Name>.cuh` |
|
||||
| `tools/codegen/generate_kernel_instantiation.py` | kernel decls in `src/kernels/**/*_kernel.cuh` | `src/instantiations/*.cu` |
|
||||
| `tools/codegen/generate_cli_config.py` | the Python config dataclasses (`ast`-parsed, no torch import) | `src/app/generated/cli_config.h` |
|
||||
| `tools/codegen/generate_backend_api.py` | per-kernel `.cuh` headers | `src/backend/api/*.h` forwarders |
|
||||
| `tools/codegen/generate_vulkan_stubs.py` | link-probes the Vulkan build | throwing stubs for unported kernels |
|
||||
|
||||
@@ -194,10 +197,13 @@ Rules:
|
||||
and add it to the list. Both build systems glob via `cmake/sources.txt`, so
|
||||
no build file changes are needed. A listed file that doesn't exist is a hard error,
|
||||
so a rename can't silently drop declarations.
|
||||
4. The Python config dataclasses are the **single source of truth** for the
|
||||
training config. Adding a field there makes it appear in the native CLI,
|
||||
the GUI's "All Options" editor, and `--help` automatically after codegen.
|
||||
A new field that collides across groups must be listed in `RENAMES`.
|
||||
4. `src/config/TrainConfig.h` is the **single source of truth** for the
|
||||
training config, and it is hand-written, not generated. Adding a row to
|
||||
`SSPLAT_CONFIG_FIELDS` makes the field appear in the native CLI, `--help`,
|
||||
the GUI's "All Options" editor, the run's `config.json` and `TrainerCore` —
|
||||
the struct is expanded from the same table, so the two cannot drift. The
|
||||
Python dataclasses are now downstream copies on their way out; a field
|
||||
added here must be mirrored there until they go.
|
||||
5. `.cuh` declaration sections must stay CUDA-include-free — they have to
|
||||
parse under `-DSSPLAT_BACKEND_VULKAN` without the CUDA toolkit.
|
||||
|
||||
|
||||
+2
-3
@@ -4,12 +4,11 @@
|
||||
# ./build_develop.bash -DSSPLAT_NO_TORCH=ON
|
||||
# builds only the standalone `ssplat` executable (no Torch/Python needed).
|
||||
|
||||
# Regenerate headers/config. Skipped when python3 is unavailable -- the
|
||||
# generated files are committed, so the build still works without it.
|
||||
# Regenerate headers. Skipped when python3 is unavailable -- the generated
|
||||
# files are committed, so the build still works without it.
|
||||
if command -v python3 >/dev/null 2>&1; then
|
||||
python3 tools/codegen/generate_headers.py
|
||||
python3 tools/codegen/generate_kernel_instantiation.py
|
||||
python3 tools/codegen/generate_cli_config.py
|
||||
else
|
||||
echo "python3 not found -- skipping codegen (using committed generated files)"
|
||||
fi
|
||||
|
||||
@@ -18,7 +18,6 @@ rem ---------------------------------------------------------------------------
|
||||
where python >nul 2>&1 || goto :skip_codegen
|
||||
python tools\codegen\generate_headers.py || goto :codegen_warn
|
||||
python tools\codegen\generate_kernel_instantiation.py || goto :codegen_warn
|
||||
python tools\codegen\generate_cli_config.py || goto :codegen_warn
|
||||
goto :msvc_env
|
||||
:codegen_warn
|
||||
echo codegen failed -- continuing with committed generated files
|
||||
|
||||
@@ -12,6 +12,7 @@ the detail.
|
||||
| [datasets.md](datasets.md) | supported dataset layouts and how they are parsed |
|
||||
| [testing.md](testing.md) | native parity tests, the three Python/C++ parity gates, the reference-dump workflow |
|
||||
| [restructure-proposal.md](restructure-proposal.md) | the in-progress plan for reorganizing the tree |
|
||||
| [notes/rename-and-i18n-plan.md](notes/rename-and-i18n-plan.md) | the Spirula Studio rename, 13-locale localization, and retiring the Python client |
|
||||
| [notes/pose-normalization.md](notes/pose-normalization.md) | orientation/centering: what the native parser implements, and the kept Python reference for what it doesn't |
|
||||
| [notes/](notes/) | design notes for individual subsystems |
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ repo: read `src/backend/README.md`, then `src/backend/vulkan/README.md`.
|
||||
|
||||
| responsibility | code |
|
||||
|---|---|
|
||||
| training config (source of truth) | Python dataclasses in `spirulae_splat/modules/` → codegen → `src/app/generated/cli_config.h` |
|
||||
| training config (source of truth) | `src/config/TrainConfig.h` — hand-written, one X-macro row per flag. The Python dataclasses in `spirulae_splat/modules/` are downstream copies until they are deleted. |
|
||||
| dataset parsing (native) | `src/data/parsers/{Colmap,Nerfstudio,Metashape}Parser.cpp`, `DatasetCommon.cpp`, `DatasetParser.h` |
|
||||
| dataset parsing (Python client) | `spirulae_splat/modules/native_dataparser.py` — an adapter, not a parser. The Python implementation is gone; `dataparser.py` is now just the config dataclass, `scripts/{colmap,metashape}_utils.py` keep a Python reader for preprocessing, and `camera_utils.py` is retained on no code path as the reference for the unported `orientation_method` / `center_method` (docs/notes/pose-normalization.md) |
|
||||
| image cache / prefetch / warp | `src/data/DataManager.cpp` |
|
||||
|
||||
+24
-23
@@ -5,11 +5,14 @@ Generated trees are marked `linguist-generated` / `linguist-vendored` in
|
||||
no Python at all — the dev build scripts run codegen when `python3` is
|
||||
available and fall back to the committed files when it isn't.
|
||||
|
||||
Every generator reads C++/CUDA sources. **None reads Python**, and none should
|
||||
be added: the training config used to be generated from the Python dataclasses
|
||||
and is now hand-written in `src/config/TrainConfig.h`.
|
||||
|
||||
Generated locations:
|
||||
|
||||
```
|
||||
src/generated/ device-math headers from Slang
|
||||
src/app/generated/ cli_config.h, viewer_html.h
|
||||
src/backend/api/ backend module forwarders
|
||||
src/instantiations/ kernel instantiation TUs
|
||||
```
|
||||
@@ -19,7 +22,6 @@ All generators run from the **repo root**:
|
||||
```bash
|
||||
python3 tools/codegen/generate_headers.py
|
||||
python3 tools/codegen/generate_kernel_instantiation.py
|
||||
python3 tools/codegen/generate_cli_config.py
|
||||
python3 tools/codegen/generate_backend_api.py
|
||||
# generate_vulkan_stubs.py takes arguments; see below
|
||||
```
|
||||
@@ -99,32 +101,31 @@ rather than instantiating everything in one TU.
|
||||
|
||||
---
|
||||
|
||||
## `generate_cli_config.py` — Python dataclasses → native config
|
||||
## The training config — *not* generated
|
||||
|
||||
**The training config's single source of truth is the Python dataclasses** —
|
||||
`TrainerConfig` and the nested `SpirulaeSplatDataParserConfig` /
|
||||
`SpirulaeSplatDataManagerConfig` / `SpirulaeSplatModelConfig` /
|
||||
`OptimizerConfig`, plus the tyro preset subclasses.
|
||||
`src/config/TrainConfig.h` is hand-written and is the training config's single
|
||||
source of truth. It holds:
|
||||
|
||||
The script parses them with `ast` (no torch import — it runs on a fresh
|
||||
checkout before `csrc.so` exists) and emits `src/app/generated/cli_config.h`:
|
||||
- `SSPLAT_CONFIG_FIELDS(X)` — one X-macro row per flag,
|
||||
`(type, member, default, group, choices, help)`;
|
||||
- `struct SsplatConfig` — **expanded from that same table**, so a field cannot
|
||||
exist in one and not the other;
|
||||
- `kSsplatPresets` + `ssplat_apply_preset()` — one branch per preset.
|
||||
|
||||
- `struct SsplatConfig` — every field, flattened, with the Python defaults
|
||||
baked in;
|
||||
- `SSPLAT_CONFIG_FIELDS(X)` — an X-macro over
|
||||
`(member, cli_key, group, choices, help)`, expanded by the CLI's generic
|
||||
parser and `--help` printer **and** by the GUI's "All Options" editor;
|
||||
- `ssplat_apply_preset()` — one branch per preset, assigning exactly the
|
||||
fields that preset overrides.
|
||||
The table is expanded by the CLI's generic parser and `--help` printer, the
|
||||
GUI's "All Options" editor, `TrainerCore`'s `config.json` dump and the pybind
|
||||
module. Add a row and the flag appears in all of them.
|
||||
|
||||
Consequences: add a field to the Python dataclass and it appears in
|
||||
`ssplat train --help`, in the GUI, and in the preset machinery after codegen.
|
||||
Names are flattened (`--model.sh-degree` → `--sh-degree`); a collision across
|
||||
groups must be resolved in the script's `RENAMES` table, and the script
|
||||
**errors on any unlisted collision** so a new Python field cannot silently
|
||||
shadow an existing flag.
|
||||
The CLI flag is `member` stringified (`--sh-degree` sets `sh_degree`; `-` and
|
||||
`_` are interchangeable), so a flag name cannot drift from its member. The
|
||||
`config.json` key is `ssplat_json_key(flag)`, which is the identity for
|
||||
everything except `dm_split_batch` → `split_batch`; that shim exists because
|
||||
`config.json` is read back by `ssplat mesh` and `--resume`, and the
|
||||
datamanager field would otherwise collide with `model.split_batch`.
|
||||
|
||||
Docstrings on the dataclass fields become CLI help text and GUI tooltips.
|
||||
This used to be generated from the Python dataclasses by
|
||||
`generate_cli_config.py`. It isn't any more — the Python dataclasses are
|
||||
downstream copies until they are deleted.
|
||||
|
||||
---
|
||||
|
||||
|
||||
@@ -0,0 +1,633 @@
|
||||
# Spirula Studio: rename, localization, and the end of the Python backend
|
||||
|
||||
Plan for three changes that look independent but share one keystone file:
|
||||
|
||||
1. **Localize** the application into 13 locales
|
||||
(en, ja, zh-Hans, zh-Hant, ko, de, fr, es, pt, it, nl, ru, tr).
|
||||
2. **Rename** `spirulae-splat` → **Spirula Studio**, `SSPLAT_*` → `SS_*`, and
|
||||
give each locale an official product name.
|
||||
3. **Retire the Python/PyTorch client.** Anything generated *from* Python
|
||||
becomes hand-written C++; a named subset of Python survives as reference
|
||||
for features that are not ported yet.
|
||||
|
||||
Status: **not started.** This document is the plan; delete phase checklists as
|
||||
they land and fold the surviving content into `docs/i18n.md` and
|
||||
`src/i18n/README.md`.
|
||||
|
||||
---
|
||||
|
||||
## 1. The keystone: `src/app/generated/cli_config.h`
|
||||
|
||||
All three changes converge on one generated file. It is produced by
|
||||
`tools/codegen/generate_cli_config.py` from the Python config dataclasses, it
|
||||
holds **192 English help strings**, and it defines `SsplatConfig` plus the
|
||||
`SSPLAT_CONFIG_FIELDS` X-macro that the CLI parser and the GUI options editor
|
||||
both expand.
|
||||
|
||||
So: it is the Python dependency (change 3), it is a third of the translation
|
||||
surface (change 1), and it carries two of the renamed identifiers (change 2).
|
||||
Migrating it first means doing that work once. Migrating it last means doing it
|
||||
three times.
|
||||
|
||||
Note that `generate_cli_config.py` is the **only** one of the five generators
|
||||
that reads Python *source*. The other four read `.cu`/`.cuh` and are ordinary
|
||||
build tooling whose outputs are committed — they are not part of "the Python
|
||||
backend" and stay. See `docs/codegen.md`.
|
||||
|
||||
**Phase order: 0 → 1 → 2 → 3 → 4.** Rename before localizing (renaming 1,400
|
||||
catalog entries afterwards is pure waste); localize after the Python client is
|
||||
gone (otherwise the CLI has two config surfaces to keep in sync).
|
||||
|
||||
---
|
||||
|
||||
## 2. Measured surface
|
||||
|
||||
Numbers from the tree as of this writing, to size the work honestly.
|
||||
|
||||
| surface | candidate literals | translate? |
|
||||
|---|---|---|
|
||||
| `src/app/gui/` (10.3k LOC) | ~800 | yes — the bulk of it |
|
||||
| `src/app/cli/` | ~460 | yes — usage/help |
|
||||
| `cli_config.h` field help | 192 | yes, last tier |
|
||||
| `src/app/webviewer/` + `viewer.html` | ~40 | later |
|
||||
| `viewer/js/` (standalone WebGL viewer) | ~161 | separate mechanism, out of scope here |
|
||||
| `src/engine/ src/sfm/ src/sam/ src/nn/ src/mesh/ src/data/` | ~1,300 | **no** — diagnostics, stay English |
|
||||
|
||||
The grep counts literals with two or more words; perhaps half of the GUI's 800
|
||||
are real UI copy and the rest are ImGui IDs, format strings and paths. Working
|
||||
estimate: **~1,400 translatable messages**, i.e. ~16,800 strings across the
|
||||
twelve non-English locales.
|
||||
|
||||
`SSPLAT_` appears on **601 lines in 95 files**: CMake options, 23 `getenv`
|
||||
names, internal CMake variables, include guards and compile-time macros.
|
||||
|
||||
---
|
||||
|
||||
## 3. Phase 0 — config source of truth moves to C++ — **LANDED 2026-08-04**
|
||||
|
||||
`tools/codegen/generate_cli_config.py` and `src/app/generated/cli_config.h`
|
||||
are deleted. `src/config/TrainConfig.h` is hand-written and is now the source
|
||||
of truth: 190 rows of
|
||||
|
||||
```cpp
|
||||
// X(type, member, default, group, choices, help)
|
||||
#define SSPLAT_CONFIG_FIELDS(X) \
|
||||
X(int, num_iterations, 30000, "trainer", "", \
|
||||
"Number of training iterations") \
|
||||
...
|
||||
```
|
||||
|
||||
with `struct SsplatConfig` **expanded from the same table**, so a field cannot
|
||||
exist in one and not the other — the drift the generator existed to prevent.
|
||||
|
||||
Four deviations from the plan as written, each for the same reason (§11: don't
|
||||
edit code you are about to delete):
|
||||
|
||||
- **`cli_key` dropped**, derived as `#member` instead. It was identical to the
|
||||
member in all 190 rows; a derivable column is exactly the duplication this
|
||||
phase is removing. Every expansion site stringifies instead.
|
||||
- **`pyname` dropped**, but it was *not* only a Python artifact: it is the key
|
||||
schema for the run's `config.json`, which `ssplat mesh` and `--resume` read
|
||||
back. It differed from the flag in exactly one row, so it became the
|
||||
`ssplat_json_key()` shim — `dm_split_batch` → `split_batch`, documented as
|
||||
compatibility rather than a mechanism.
|
||||
- **`help` stays a literal**, not a `Msg` reference. `Msg` does not exist until
|
||||
Phase 3; the table gets rewritten then.
|
||||
- **`SsplatConfig` keeps its name.** Renaming it to `TrainConfig` now would
|
||||
touch `bind_trainer.cpp` and `native_trainer.py`, both deleted in Phase 1,
|
||||
and would muddy the byte-identical gate below. Moved to the Phase 2
|
||||
checklist as an explicit item (the `Ssplat` → `Ss` sed does not produce
|
||||
`TrainConfig` on its own).
|
||||
|
||||
Verified: `train --help` byte-identical for **all 7 presets** against the last
|
||||
generated build (covers every default, preset override, group ordering,
|
||||
choices list and help string), a run's `config.json` byte-identical, CUDA and
|
||||
Vulkan builds green, and `bind_trainer.cpp` compiles.
|
||||
|
||||
Not touched: `scripts/`, which Phase 0 turned out not to reach.
|
||||
|
||||
Follow-ups this surfaced, deliberately left alone to keep the gate clean:
|
||||
|
||||
- The 33 `optimizer`-group fields have **empty help strings** — `OptimizerConfig`
|
||||
never had docstrings. They are the worst entries in `--help` and should be
|
||||
written before Phase 4 translates anything.
|
||||
- The `meshing` preset's help says "Use `spirulae-meshing`", a binary that no
|
||||
longer exists; it is `ssplat mesh`.
|
||||
|
||||
---
|
||||
|
||||
## 4. Phase 1 — retiring the Python client
|
||||
|
||||
The 9.9k lines of `spirulae_splat/` are mostly the torch client path, which
|
||||
`TrainerCore` already replaced. What genuinely has no C++ counterpart:
|
||||
|
||||
| Python | LOC | disposition |
|
||||
|---|---|---|
|
||||
| `resume.py` + `resume_codecs.py` + `resume_adapt.py` | 728 | **port** → `src/checkpoint/`. Serialization and slot adaptation; needs no torch. |
|
||||
| `lpips.py` | 530 | **port** → `src/lpips/`, on top of `src/nn/`. Model-specific weights and constants stay out of `nn/` per the layering rule in AGENTS.md. |
|
||||
| `camera_utils.py` | 680 | **keep as reference** — already on no code path; the `orientation_method` / `center_method` reference (`docs/notes/pose-normalization.md`). |
|
||||
| `enhancer.py`, `resample.py`, `edge_detector.py`, `debug_image.py` | ~700 | audit individually; most are thin wrappers over `_C` and die with the binding. |
|
||||
| `model.py`, `trainer.py`, `core.py`, `dataset.py`, `datamanager.py`, `dataparser.py`, `splat/` | ~4,700 | **delete** once Phase 0 + the two ports land. |
|
||||
|
||||
Steps:
|
||||
|
||||
1. Port resume. This is the single blocker — until it lands, Python is the only
|
||||
way to continue an interrupted run.
|
||||
2. Port LPIPS onto `nn/`. Cross-check against the Python implementation on a
|
||||
fixed image pair before deleting it (a parity gate in the style of
|
||||
`docs/testing.md` §4-6).
|
||||
3. Move survivors to **`reference/python/`** at the repo root — outside the
|
||||
importable package, absent from `pyproject.toml`, with a README stating they
|
||||
are on no code path. The current `spirulae_splat/modules/` layout makes
|
||||
accidental imports possible; this makes them impossible and `grep`
|
||||
unambiguous.
|
||||
4. Delete `spirulae_splat/`, `setup.py`, the `ss_*` console scripts, and
|
||||
`src/bindings/` (144 `m.def`s). Drop `SSPLAT_NO_TORCH` / `SSPLAT_WITH_TORCH`
|
||||
and the torch half of `cmake/SsplatBackendCuda.cmake`.
|
||||
5. Update AGENTS.md §"What this project is" — the "the Python path must keep
|
||||
working" rule is retired, and `tests/python/` goes with it.
|
||||
|
||||
**Do not rename the Python package.** It is being deleted; renaming it first is
|
||||
work spent on a corpse.
|
||||
|
||||
---
|
||||
|
||||
## 5. Phase 2 — the rename
|
||||
|
||||
### 5.1 Names
|
||||
|
||||
| thing | from | to |
|
||||
|---|---|---|
|
||||
| product | spirulae-splat | **Spirula Studio** |
|
||||
| repository | `spirulae-splat` | `spirula-studio` |
|
||||
| executable | `ssplat` | `spirula` (`spirula train\|sfm\|sam\|mesh`; symlink `spirula-sfm`) |
|
||||
| macro / option / env prefix | `SSPLAT_` | `SS_` |
|
||||
| CMake modules | `cmake/Ssplat*.cmake` | `cmake/Ss*.cmake` |
|
||||
| config type | `SsplatConfig` | `TrainConfig` (Phase 0) |
|
||||
| C++ namespace | `ssplat` (6 files) | `spirula` |
|
||||
|
||||
Leave `viewer/` alone — the GitHub Pages URL depends on the directory name.
|
||||
Renaming the repository **will** change that URL; GitHub redirects git remotes
|
||||
and web UI paths on rename but does not reliably redirect Pages project paths.
|
||||
Verify before flipping, and consider keeping a redirect stub at the old repo.
|
||||
|
||||
### 5.2 `SS_` and the avoid list
|
||||
|
||||
`SS_` is short enough to collide with system headers, and the two that matter
|
||||
are both on our platforms:
|
||||
|
||||
- **`<signal.h>`** (glibc, macOS): `SS_ONSTACK`, `SS_DISABLE`.
|
||||
- **`winuser.h`** (Windows SDK) — static-control styles:
|
||||
`SS_LEFT`, `SS_CENTER`, `SS_RIGHT`, `SS_ICON`, `SS_BLACKRECT`,
|
||||
`SS_GRAYRECT`, `SS_WHITERECT`, `SS_BLACKFRAME`, `SS_GRAYFRAME`,
|
||||
`SS_WHITEFRAME`, `SS_USERITEM`, `SS_SIMPLE`, `SS_LEFTNOWORDWRAP`,
|
||||
`SS_OWNERDRAW`, `SS_BITMAP`, `SS_ENHMETAFILE`, `SS_ETCHEDHORZ`,
|
||||
`SS_ETCHEDVERT`, `SS_ETCHEDFRAME`, `SS_TYPEMASK`, `SS_REALSIZECONTROL`,
|
||||
`SS_NOPREFIX`, `SS_NOTIFY`, `SS_CENTERIMAGE`, `SS_RIGHTJUST`,
|
||||
`SS_REALSIZEIMAGE`, `SS_SUNKEN`, `SS_EDITCONTROL`, `SS_ENDELLIPSIS`,
|
||||
`SS_PATHELLIPSIS`, `SS_WORDELLIPSIS`, `SS_ELLIPSISMASK`.
|
||||
|
||||
The GUI includes `winuser.h` transitively through GLFW, and the webviewer's
|
||||
`HttpServer.cpp` includes `winsock2.h`, so both are live. Regenerate the list
|
||||
against the SDK you ship with rather than trusting this one:
|
||||
|
||||
```bash
|
||||
grep -hoE '^\s*#\s*define\s+SS_[A-Z0-9_]+' <sdk>/winuser.h | awk '{print $NF}' | sort -u
|
||||
```
|
||||
|
||||
Enforcement: `tools/check_ss_prefix.sh`, run from `build_develop.bash`, greps
|
||||
our own `#define SS_*`, `option(SS_*` and `set(SS_*` declarations against
|
||||
`tools/ss_reserved_names.txt` and fails on a hit. None of the current 95
|
||||
`SSPLAT_` names maps onto a reserved one, so the initial rename is clean — the
|
||||
lint exists to stop the next one.
|
||||
|
||||
### 5.3 Compatibility shim
|
||||
|
||||
Two shims, both designed to be deletable in a single commit one release later.
|
||||
|
||||
**CMake options and cache variables** — a table-driven loop in
|
||||
`cmake/SsOptions.cmake`, ahead of the `option()` declarations:
|
||||
|
||||
```cmake
|
||||
set(SS_LEGACY_OPTIONS
|
||||
BUILD_CLI BUILD_GUI BUILD_SFM BUILD_SAM BUILD_BACKEND_TESTS
|
||||
BACKEND SEPARATE_TOOLS DEBUG_SYMBOLS ENABLE_PATENTED)
|
||||
foreach(opt ${SS_LEGACY_OPTIONS})
|
||||
if(DEFINED SSPLAT_${opt} AND NOT DEFINED SS_${opt})
|
||||
message(DEPRECATION
|
||||
"SSPLAT_${opt} is deprecated; use SS_${opt}. "
|
||||
"The alias will be removed in the release after next.")
|
||||
set(SS_${opt} "${SSPLAT_${opt}}" CACHE STRING "" FORCE)
|
||||
endif()
|
||||
endforeach()
|
||||
```
|
||||
|
||||
`SSPLAT_NO_TORCH` and `SSPLAT_WITH_TORCH` get no alias — they are deleted in
|
||||
Phase 1, and a deprecation warning is friendlier than silently ignoring them.
|
||||
|
||||
**Environment variables** — replace all 23 raw `std::getenv("SSPLAT_...")`
|
||||
sites with one helper, so the deprecation lives in exactly one function:
|
||||
|
||||
```cpp
|
||||
// src/core/Env.h
|
||||
namespace spirula {
|
||||
// Reads SS_<suffix>, falling back to the deprecated SSPLAT_<suffix> with a
|
||||
// one-shot warning. Delete the fallback with the alias table in cmake/.
|
||||
const char* env(const char* suffix);
|
||||
}
|
||||
```
|
||||
|
||||
Call sites become `spirula::env("VK_DEVICE")`. That is a strict improvement
|
||||
over the status quo regardless of the rename: today the 23 names are scattered
|
||||
and undocumented.
|
||||
|
||||
### 5.4 Execution
|
||||
|
||||
The identifier rename is mechanical (`sed` over `SSPLAT_` → `SS_`, `Ssplat` →
|
||||
`Ss`, `ssplat` → `spirula`) but touches the committed generated trees
|
||||
(`src/generated/`, `src/instantiations/`, `src/app/generated/`, 100k+ lines).
|
||||
Rename the sources, re-run codegen, commit the regenerated output in the *same*
|
||||
commit, and confirm the regenerated diff contains nothing but the rename.
|
||||
|
||||
Do this on a quiet tree, in one commit, with no other change riding along.
|
||||
|
||||
---
|
||||
|
||||
## 6. Phase 3 — the localization mechanism
|
||||
|
||||
Design goals, in priority order: **a missing translation must fail the build**;
|
||||
no new dependency; catalogs modular, one per module.
|
||||
|
||||
The trick is to make a translation a *type* rather than a lookup key. Then
|
||||
"missing translation" is a `static_assert`, not a runtime fallback.
|
||||
|
||||
### 6.1 The language list
|
||||
|
||||
One file declares the set; everything else derives from it.
|
||||
|
||||
```cpp
|
||||
// src/i18n/Languages.h -- the ONE place the language set is written
|
||||
#define SS_LANGUAGES(X) \
|
||||
X(en, "English") \
|
||||
X(ja, "日本語") \
|
||||
X(zh_hans, "简体中文") \
|
||||
X(zh_hant, "繁體中文") \
|
||||
X(ko, "한국어") \
|
||||
X(de, "Deutsch") \
|
||||
X(fr, "Français") \
|
||||
X(es, "Español") \
|
||||
X(pt, "Português") \
|
||||
X(it, "Italiano") \
|
||||
X(nl, "Nederlands") \
|
||||
X(ru, "Русский") \
|
||||
X(tr, "Türkçe")
|
||||
```
|
||||
|
||||
Adding a locale here breaks the build on every incomplete message — which is
|
||||
the point.
|
||||
|
||||
### 6.2 `Msg`
|
||||
|
||||
```cpp
|
||||
// src/i18n/Message.h
|
||||
namespace spirula::i18n {
|
||||
|
||||
enum class Lang : unsigned {
|
||||
#define X(id, native) id,
|
||||
SS_LANGUAGES(X)
|
||||
#undef X
|
||||
};
|
||||
inline constexpr unsigned kLangCount = 0
|
||||
#define X(id, native) + 1
|
||||
SS_LANGUAGES(X)
|
||||
#undef X
|
||||
;
|
||||
|
||||
template <Lang L> struct Tr { const char* s; };
|
||||
|
||||
struct Msg {
|
||||
const char* v[kLangCount] = {};
|
||||
|
||||
template <Lang... Ls>
|
||||
constexpr explicit Msg(Tr<Ls>... t) : v{} {
|
||||
static_assert(sizeof...(Ls) == kLangCount,
|
||||
"i18n: wrong number of translations");
|
||||
((v[unsigned(Ls)] = t.s), ...); // C++17 fold
|
||||
}
|
||||
constexpr bool complete() const {
|
||||
for (unsigned i = 0; i < kLangCount; i++)
|
||||
if (!v[i] || !*v[i]) return false; // also catches a duplicate tag
|
||||
return true;
|
||||
}
|
||||
const char* get() const { return v[unsigned(current())]; }
|
||||
};
|
||||
|
||||
Lang current();
|
||||
void set_current(Lang);
|
||||
|
||||
} // namespace spirula::i18n
|
||||
|
||||
#define SS_MSG(name, ...) \
|
||||
inline constexpr ::spirula::i18n::Msg name{__VA_ARGS__}; \
|
||||
static_assert(name.complete(), \
|
||||
"i18n: '" #name "' is missing a translation")
|
||||
```
|
||||
|
||||
Properties worth noting: the tags carry their own slot, so **order does not
|
||||
matter**; a duplicated tag necessarily leaves a hole and is caught by
|
||||
`complete()`; a wrong count is caught by its own `static_assert`. C++17
|
||||
throughout — no need to raise `CMAKE_CXX_STANDARD` from 17.
|
||||
|
||||
Cost: all 13 languages are always linked. At ~1,400 messages that is roughly
|
||||
1 MB of `.rodata`. Acceptable, and far simpler than a resource-file scheme.
|
||||
|
||||
### 6.3 A catalog
|
||||
|
||||
One header per module, included only by that module's `.cpp` files. The short
|
||||
tag macros are defined by `Message.h` and `#undef`'d by `EndCatalog.h`, so
|
||||
`EN`/`JA`/… never leak into ordinary code:
|
||||
|
||||
```cpp
|
||||
// src/i18n/catalog/Gui.h
|
||||
#include "i18n/BeginCatalog.h"
|
||||
namespace spirula::i18n::gui {
|
||||
|
||||
SS_MSG(open_dataset,
|
||||
EN("Open a Dataset..."), JA("データセットを開く…"),
|
||||
ZH_HANS("打开数据集…"), ZH_HANT("開啟資料集…"),
|
||||
KO("데이터셋 열기…"), DE("Datensatz öffnen …"),
|
||||
FR("Ouvrir un jeu de données…"), ES("Abrir un conjunto de datos…"),
|
||||
PT("Abrir um conjunto de dados…"), IT("Apri un set di dati…"),
|
||||
NL("Dataset openen…"), RU("Открыть набор данных…"),
|
||||
TR("Veri kümesi aç…"));
|
||||
|
||||
}
|
||||
#include "i18n/EndCatalog.h"
|
||||
```
|
||||
|
||||
Planned catalogs: `Gui.h`, `Cli.h`, `Config.h` (the 192 field help strings),
|
||||
`Errors.h`, `Brand.h` (§7).
|
||||
|
||||
### 6.4 The second half: unmarked strings
|
||||
|
||||
`Msg` cannot catch `ImGui::Button("Start")` — that compiles fine. Close it by
|
||||
routing every text-rendering call through a thin wrapper whose parameters take
|
||||
`const Msg&`:
|
||||
|
||||
```cpp
|
||||
// src/app/gui/Ui.h -- the only ImGui text entry points the GUI may call
|
||||
namespace ui {
|
||||
inline bool Button(const i18n::Msg& m, ImVec2 sz = {}) {
|
||||
return ImGui::Button(m.get(), sz);
|
||||
}
|
||||
inline void Text(const i18n::Msg& m) { ImGui::TextUnformatted(m.get()); }
|
||||
// Explicitly NOT translated: paths, numbers, engine log lines.
|
||||
inline void TextRaw(const char* s) { ImGui::TextUnformatted(s); }
|
||||
}
|
||||
```
|
||||
|
||||
A bare literal then does not compile. Back it with `tools/check_i18n.sh`,
|
||||
which fails if `ImGui::{Text,TextWrapped,TextColored,TextUnformatted,Button,
|
||||
SmallButton,Checkbox,Combo,MenuItem,BeginMenu,CollapsingHeader,SeparatorText,
|
||||
RadioButton,SetTooltip,LabelText,SliderFloat,SliderInt,InputText,InputInt,
|
||||
InputFloat}` appears anywhere in `src/app/gui/` outside `Ui.h`. Type system
|
||||
covers "incomplete translation"; the lint covers "unmarked string".
|
||||
|
||||
### 6.5 Two rules to fix now
|
||||
|
||||
Retrofitting either of these is expensive.
|
||||
|
||||
- **Never concatenate sentences.** Use positional substitution —
|
||||
`format(msg, {a, b})` over `{0}` / `{1}` placeholders, ~30 lines, no
|
||||
dependency. Every one of these 13 languages reorders clauses relative to
|
||||
English; Japanese, Korean and Turkish are verb-final.
|
||||
- **No plural-sensitive sentences in the catalog.** Write `Images: 5`, not
|
||||
`5 images`. Otherwise Russian needs a three-form CLDR plural rule
|
||||
(one/few/many) and every message that counts anything triples.
|
||||
|
||||
### 6.6 Staged rollout
|
||||
|
||||
1,400 messages × 12 locales will not land atomically, and an unenforced
|
||||
fallback rots silently. One explicit escape hatch:
|
||||
|
||||
```cpp
|
||||
#define SS_MSG_EN(name, s) \
|
||||
inline constexpr ::spirula::i18n::Msg name = ::spirula::i18n::en_only(s)
|
||||
```
|
||||
|
||||
It is greppable and countable — `grep -rc SS_MSG_EN src/i18n/catalog/` **is**
|
||||
the TODO list — and catalogs flip to full enforcement one at a time. Tiers:
|
||||
|
||||
| tier | scope | messages |
|
||||
|---|---|---|
|
||||
| 1 | GUI chrome, menus, errors, the ~40 basic options | ~350 |
|
||||
| 2 | remaining GUI + CLI usage | ~850 |
|
||||
| 3 | the 192 config help strings | ~190 |
|
||||
|
||||
Machine translation is a reasonable first pass for tiers 2-3. It is **not**
|
||||
acceptable for the ~40 messages attached to irreversible actions (delete,
|
||||
overwrite, "this will erase") — those get human review in every locale before
|
||||
shipping. ja/zh/ko/de deserve review at tier 1 regardless; they are the four
|
||||
where a bad string is most visible.
|
||||
|
||||
### 6.7 Locale resolution
|
||||
|
||||
In order, first hit wins:
|
||||
|
||||
1. `--lang <code>` on the command line
|
||||
2. `SS_LANG` environment variable
|
||||
3. the settings file (`%APPDATA%\Spirula Studio\settings.json`,
|
||||
`~/.config/spirula-studio/settings.json`)
|
||||
4. the OS locale — `GetUserDefaultLocaleName` (Windows), `NSLocale` (macOS),
|
||||
`LC_ALL` / `LC_MESSAGES` / `LANG` (Linux)
|
||||
5. **`SS_DEFAULT_LANG`, the compile-time default** (§8.3)
|
||||
|
||||
Chinese mapping needs care: `zh_CN`, `zh_SG`, `zh-Hans-*` → zh-Hans;
|
||||
`zh_TW`, `zh_HK`, `zh_MO`, `zh-Hant-*` → zh-Hant. A bare `zh` follows
|
||||
`SS_DEFAULT_LANG` if that is a Chinese locale, else zh-Hans. Unknown locales
|
||||
fall to `SS_DEFAULT_LANG`, never to a hard-coded `en`.
|
||||
|
||||
---
|
||||
|
||||
## 7. Product names per locale
|
||||
|
||||
`Brand.h` holds the product name as an ordinary `SS_MSG`, so the policy is data
|
||||
rather than code.
|
||||
|
||||
| locale | name | note |
|
||||
|---|---|---|
|
||||
| en, de, fr, es, pt, it, nl, tr | Spirula Studio | Latin-script markets keep the wordmark |
|
||||
| zh-Hans / zh-Hant | 旋影工坊 | **all four characters are identical in Simplified and Traditional**, so one name serves both scripts — and 旋 preserves the spiral sense of *Spirula* |
|
||||
| ja | スピルラ・スタジオ | pin the katakana, or users independently invent スピルーラ / スパイルラ |
|
||||
| ko | 스피룰라 스튜디오 | same reasoning |
|
||||
| ru | Spirula Studio (Спирула Студио) | Russian technical users keep Latin brand names; Cyrillic as a first-mention gloss only |
|
||||
|
||||
Recommended policy: the **Latin wordmark stays the logo in every locale**, and
|
||||
the localized name is in-text copy — window titles, About, prose. This is what
|
||||
Blender, Krita and Godot do, and it avoids maintaining five logo lockups. The
|
||||
exception is zh, where 旋影工坊 can stand alone; CJK markets genuinely adopt
|
||||
local names, which is the failure mode this whole section exists to prevent.
|
||||
|
||||
---
|
||||
|
||||
## 8. Fonts
|
||||
|
||||
**This is the unbudgeted part.** `GuiMain.cpp` contains no font code at all, so
|
||||
the GUI runs on ImGui's built-in ProggyClean — a 13px ASCII-only bitmap font.
|
||||
German `ö`, French `é` and Turkish `ğ` are already broken today, before any CJK
|
||||
is involved.
|
||||
|
||||
ImGui is pinned at **v1.92.8**, which has the dynamic font atlas: glyphs
|
||||
rasterize on demand and ranges need not be enumerated up front. `GuiMain.cpp`
|
||||
already uses `style.FontScaleMain`, so the codebase is on the new API. What the
|
||||
new system does *not* do is find a font for you — the glyphs still have to ship.
|
||||
|
||||
### 8.1 Font choice
|
||||
|
||||
Nothing with CJK coverage looks like ProggyClean; it is a pixel font and they
|
||||
are outline fonts. The UI **will** change appearance, and pretending otherwise
|
||||
sets up a bad surprise. The least jarring choice:
|
||||
|
||||
**Source Sans 3** (Latin / Greek / Cyrillic, OFL-1.1) + **Source Han Sans**
|
||||
(CJK, OFL-1.1).
|
||||
|
||||
The reason for that specific pair over Noto Sans: Source Han Sans embeds Source
|
||||
Sans as its own Latin, so the two are designed together — matching x-height,
|
||||
weights and vertical metrics, which is exactly what keeps a mixed
|
||||
Latin/CJK line from visibly stepping. (Noto Sans CJK and Source Han Sans are
|
||||
the same typeface under two names; Noto Sans, the Latin family, is a different
|
||||
design.) Turkish `ı ğ ş İ` and full Cyrillic are covered.
|
||||
|
||||
Mitigations for the ProggyClean → outline transition: raise the base size from
|
||||
13 to 15px, and consider adding `imgui_freetype` with light hinting — the
|
||||
appeal of ProggyClean is pixel-crispness, and stb_truetype at 13px is
|
||||
noticeably softer. FreeType is the one dependency worth the argument here; it
|
||||
is optional and gated.
|
||||
|
||||
An alternative was considered and rejected: keep ProggyClean for English and
|
||||
switch fonts only for other locales. It preserves the current look exactly for
|
||||
English users, but it means two UI appearances to maintain and screenshot, and
|
||||
mixed pixel/outline glyphs on the same line look worse in French and German
|
||||
than a clean switch does.
|
||||
|
||||
### 8.2 `SS_FONT_CJK` — embed or fetch
|
||||
|
||||
```cmake
|
||||
set(SS_FONT_CJK "fetch" CACHE STRING
|
||||
"CJK font: fetch | none | sc | tc | jp | kr | all")
|
||||
```
|
||||
|
||||
| value | behaviour | binary cost |
|
||||
|---|---|---|
|
||||
| `fetch` *(default)* | Latin/Cyrillic embedded; the matching Source Han Sans regional face is downloaded on first use of a CJK locale, through the consent-gated path already built for SAM weights in `src/app/gui/ModelCache.cpp` | ~0.4 MB |
|
||||
| `none` | Latin/Cyrillic only; CJK locales render tofu | ~0.4 MB |
|
||||
| `sc` / `tc` / `jp` / `kr` | embed exactly that regional face — **this is the per-region binary path** | ~+16 MB |
|
||||
| `all` | embed all four | ~+65 MB |
|
||||
|
||||
Sizes are for the language-specific OTFs and are approximate; check against the
|
||||
release you vendor, and note that the *Subset* OTFs are smaller but drop
|
||||
coverage.
|
||||
|
||||
**Han unification matters here.** The regional faces of Source Han Sans differ
|
||||
in default glyph forms for shared codepoints (直, 骨, 雪, 戸 and many more). A
|
||||
`sc` binary shown to a Japanese reader renders kanji in Chinese forms — legible,
|
||||
and visibly wrong. So a per-region build should pair `SS_FONT_CJK` with a
|
||||
matching `SS_DEFAULT_LANG`, and the runtime should still offer to fetch the
|
||||
correct regional face if the user switches away from the embedded region. Keep
|
||||
the fetch path compiled in for every value except `none`.
|
||||
|
||||
Embedding reuses `ssplat_embed_file()` from `cmake/SsplatEmbed.cmake` — the
|
||||
same mechanism as `viewer.html` and `mask.py`. Fonts are OFL-1.1, which is
|
||||
GPLv3-compatible for bundling; ship the licence text alongside, and do not
|
||||
rename the font files (OFL reserved font name clause).
|
||||
|
||||
### 8.3 `SS_DEFAULT_LANG`
|
||||
|
||||
```cmake
|
||||
set(SS_DEFAULT_LANG "en" CACHE STRING
|
||||
"Locale used when none is detected: en|ja|zh_hans|zh_hant|ko|de|fr|es|pt|it|nl|ru|tr")
|
||||
```
|
||||
|
||||
The last resort in the §6.7 chain — headless runs, containers with no `LANG`,
|
||||
and stripped Windows environments all land here. A regional build sets it:
|
||||
|
||||
```bash
|
||||
# Simplified-Chinese build
|
||||
cmake -DSS_DEFAULT_LANG=zh_hans -DSS_FONT_CJK=sc ...
|
||||
# Japanese build
|
||||
cmake -DSS_DEFAULT_LANG=ja -DSS_FONT_CJK=jp ...
|
||||
```
|
||||
|
||||
Add a configure-time consistency check, in the spirit of the rest of the
|
||||
localization design — an inconsistent combination should fail early rather than
|
||||
ship tofu:
|
||||
|
||||
```cmake
|
||||
if(SS_DEFAULT_LANG MATCHES "^(ja|ko|zh_hans|zh_hant)$" AND SS_FONT_CJK STREQUAL "none")
|
||||
message(FATAL_ERROR
|
||||
"SS_DEFAULT_LANG=${SS_DEFAULT_LANG} needs a CJK font; "
|
||||
"set SS_FONT_CJK to fetch, sc, tc, jp, kr or all.")
|
||||
endif()
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 9. Comments
|
||||
|
||||
Trim comments **in files you are already editing**, never as a standalone pass —
|
||||
a repo-wide comment diff has no test coverage and will conflict with all three
|
||||
phases above.
|
||||
|
||||
The rule: **keep a comment that records a non-obvious invariant or a
|
||||
why-not; delete one that restates the code.** AGENTS.md's "Gotchas" section is
|
||||
the right register — every entry there is something that cost someone a day.
|
||||
The header comment on `ConfigUI.cpp` ("keep this file free of per-field special
|
||||
cases") is a keeper. Most of the inline narration in `GuiMain.cpp` is not.
|
||||
|
||||
Internal comments stay in English in every file, permanently. Only catalog
|
||||
strings are localized.
|
||||
|
||||
---
|
||||
|
||||
## 10. Phase checklist
|
||||
|
||||
- [x] **0.** `TrainConfig.h` hand-written; `generate_cli_config.py` deleted;
|
||||
`--help` diffed byte-for-byte across all 7 presets. *(2026-08-04)*
|
||||
- [ ] **1.** Resume ported → `src/checkpoint/`. LPIPS ported → `src/lpips/`
|
||||
with a parity gate. Survivors moved to `reference/python/`.
|
||||
`spirulae_splat/`, `setup.py`, `src/bindings/`, `tests/python/` deleted.
|
||||
AGENTS.md updated.
|
||||
- [ ] **2.** `tools/ss_reserved_names.txt` + `check_ss_prefix.sh`.
|
||||
`SSPLAT_` → `SS_` in sources, generated trees regenerated in the same
|
||||
commit. `spirula::env()` replaces 23 `getenv` sites. CMake alias loop.
|
||||
`SsplatConfig` → `TrainConfig` (a separate rename, not the prefix sed),
|
||||
and the `SsplatVec3*` / `ssplat_v3*` / `ssplat_json_key` helpers with it.
|
||||
Executable and repository renamed; Pages URL verified.
|
||||
- [ ] **3.** `src/i18n/` (`Languages.h`, `Message.h`, `Begin/EndCatalog.h`),
|
||||
`gui/Ui.h`, `check_i18n.sh`, font loading + `SS_FONT_CJK` +
|
||||
`SS_DEFAULT_LANG`, locale detection. All catalogs `SS_MSG_EN` at first.
|
||||
- [ ] **4.** Translate tier 1 → 2 → 3; each catalog flips from `SS_MSG_EN` to
|
||||
`SS_MSG` as it completes. Human review for ja/zh/ko/de and for every
|
||||
irreversible-action message.
|
||||
- [ ] **5.** Fold this document into `docs/i18n.md` + `src/i18n/README.md`;
|
||||
drop the `SSPLAT_` aliases one release later.
|
||||
|
||||
## 11. Recommendation
|
||||
|
||||
Phase 0 first and alone — it is a few days, it unblocks everything, and it is
|
||||
the one piece that is pure gain even if the rest slips.
|
||||
|
||||
Then Phase 1, which is gated almost entirely on the resume port. Do not start
|
||||
the rename before it: renaming code you are about to delete is the largest
|
||||
avoidable cost in this plan.
|
||||
|
||||
Phase 2 is a day of mechanical work plus a careful pass over the public option
|
||||
and environment-variable names. Do it on a quiet tree, in one commit.
|
||||
|
||||
Phase 3's plumbing is small — `Msg` is 60 lines. Budget the time for fonts
|
||||
instead; §8 is where the surprises are, and `SS_FONT_CJK=fetch` is what keeps
|
||||
the default binary from growing 16 MB for a feature most users of a given build
|
||||
will not use.
|
||||
|
||||
Phase 4 is the long tail and is the only part that can ship incrementally,
|
||||
which is exactly why `SS_MSG_EN` exists.
|
||||
@@ -416,6 +416,12 @@ The config dataclasses stay the single source of truth
|
||||
that generator to also emit the pybind `TrainerConfig` struct so the three
|
||||
representations can't diverge.
|
||||
|
||||
> **Superseded twice.** The binding was done through the existing
|
||||
> `SSPLAT_CONFIG_FIELDS` X-macro instead (see §7.2 below), and on 2026-08-04
|
||||
> the direction reversed entirely: `generate_cli_config.py` is gone and the
|
||||
> hand-written `src/config/TrainConfig.h` is the source of truth, with the
|
||||
> dataclasses as the downstream copy. See `docs/notes/rename-and-i18n-plan.md`.
|
||||
|
||||
**Verification gate:** train a short run on a public scene through both paths
|
||||
before/after and compare per-step loss to within float noise; the existing
|
||||
`backend/tests/engine/engine_train_parity.cpp` is the model for this.
|
||||
@@ -435,7 +441,7 @@ One screen of orientation plus links. Target contents:
|
||||
install`*, for development.
|
||||
- **Codegen invariants** — what `generate_headers.py` (the
|
||||
`/*[AutoHeaderGeneratorExport]*/` marker and the `Name.*.cu` collection
|
||||
rule), `generate_kernel_instantiation.py`, `generate_cli_config.py`,
|
||||
rule), `generate_kernel_instantiation.py`,
|
||||
`generate_backend_api.py`, `generate_vulkan_stubs.py` each own; never
|
||||
hand-edit below the `AUTO HEADER GENERATOR` splitter line; `.gitattributes`
|
||||
marks generated/vendored trees.
|
||||
|
||||
+4
-4
@@ -134,9 +134,9 @@ viewer read its step counter / pause flag / progress JSON straight off a
|
||||
|
||||
1. **Config conversion** — `to_native_config(PresetClass())` must equal
|
||||
`SsplatConfig()` + `ssplat_apply_preset(name)` for all seven presets. These
|
||||
are still *two live representations*: one side reads the dataclasses, the
|
||||
other is baked into `cli_config.h` at codegen time. Still a true parity
|
||||
test; it fails the moment the generated header goes stale.
|
||||
are *two live representations*: `src/config/TrainConfig.h` is the source of
|
||||
truth and the Python dataclasses are the downstream copy, so this test is
|
||||
what catches the copy drifting until the dataclasses are deleted.
|
||||
2. **Per-step `EngineStepConfig`** — 8 config variants × 4 run states × 20
|
||||
steps chosen to straddle every warmup/decay boundary, checked against
|
||||
`step_config_golden.json`. Those 640 configs are the frozen result of the
|
||||
@@ -190,7 +190,7 @@ proof that no longer has a second implementation to re-derive it from. So:
|
||||
|---|---|
|
||||
| any kernel | CUDA build + Vulkan build + the relevant parity test on both |
|
||||
| engine logic | both builds + `engine_render_parity` + `engine_train_step`-level check |
|
||||
| config field | rerun `generate_cli_config.py`; check `ssplat train --help`; `test_trainer_parity.py` |
|
||||
| config field | add the row in `src/config/TrainConfig.h`, mirror it on the Python dataclass; check `ssplat train --help`; `test_trainer_parity.py` |
|
||||
| training-loop logic | change `TrainerCore.cpp`, not the Python mirror; `test_trainer_parity.py` |
|
||||
| build system | all four modes in [build.md](build.md) |
|
||||
| Python-facing | a short `spirulae-train` run with `--no-keep-viewer-alive` |
|
||||
|
||||
@@ -5,16 +5,14 @@ and `metashape_utils.py`) is gone: COLMAP / Nerfstudio / Metashape are parsed
|
||||
by `src/data/parsers/` and reached from Python through
|
||||
`modules/native_dataparser.py`. See docs/restructure-proposal.md §4.1.
|
||||
|
||||
The config dataclass stays, and stays HERE, because it is the source of truth
|
||||
for two generated artifacts and one CLI:
|
||||
* `tools/codegen/generate_cli_config.py` AST-parses this file to emit
|
||||
`src/app/generated/cli_config.h` (and so the native CLI's flags);
|
||||
The config dataclass stays, and stays HERE, because:
|
||||
* `ss_trainer.py` builds its tyro CLI from it, and `--resume` reads it back
|
||||
out of a run's config.json;
|
||||
* `native_dataparser.to_native_parser_config()` maps it onto the native
|
||||
`DatasetParserConfig`.
|
||||
Editing a field here therefore changes the native trainer too -- re-run the
|
||||
generator.
|
||||
It is no longer the source of truth: `src/config/TrainConfig.h` is, and this
|
||||
is a downstream copy that has to be kept in step until the Python client is
|
||||
deleted. `tests/python/test_trainer_parity.py` is what catches the drift.
|
||||
|
||||
Fields with no native counterpart (scene_scale, orientation_method,
|
||||
center_method, auto_scale_poses, train_frame != "points") are kept so old
|
||||
|
||||
@@ -15,13 +15,13 @@ step. See docs/restructure-proposal.md §4.3.
|
||||
|
||||
Config conversion
|
||||
-----------------
|
||||
`SsplatConfig` is generated from the Python dataclasses by
|
||||
`tools/codegen/generate_cli_config.py`, and that generator also emits the
|
||||
field table this module walks (`_C.ssplat_config_fields()` -> (cli_key,
|
||||
pyname, group, choices, help)). So the flattening -- which Python sub-config a
|
||||
field lives in, and what it is called after the rename pass -- exists exactly
|
||||
once, in the generator. Adding a field to a dataclass and re-running the
|
||||
generator is enough; nothing here needs editing.
|
||||
`SsplatConfig` and the field table this module walks
|
||||
(`_C.ssplat_config_fields()` -> (cli_key, pyname, group, choices, help)) both
|
||||
come from the hand-written `src/config/TrainConfig.h`. That header is the
|
||||
source of truth as of 2026-08-04; the dataclasses below it are a downstream
|
||||
copy, kept only until the Python client is deleted. So the direction of the
|
||||
check has flipped: a field present natively and missing on the dataclass is
|
||||
the error, and that is what to_native_config() raises on.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -43,8 +43,8 @@ PRESET_BY_CLASS_NAME = {
|
||||
"TrainerConfigAcademicBaseline": "academic-baseline",
|
||||
}
|
||||
|
||||
# Fields whose Python type the generator overrode (TYPE_OVERRIDES in
|
||||
# generate_cli_config.py), so a straight copy is wrong.
|
||||
# Fields whose Python type differs from the native one, so a straight copy is
|
||||
# wrong.
|
||||
# rescale_camera_to_fit: Union[bool, int] on the Python side.
|
||||
# True -> probe the image resolution -> -1 natively
|
||||
# False -> off -> 0
|
||||
@@ -82,9 +82,9 @@ def to_native_config(config):
|
||||
value = getattr(holder, pyname)
|
||||
except AttributeError as e:
|
||||
raise AttributeError(
|
||||
f"{type(config).__name__}.{group}.{pyname} is missing; "
|
||||
f"cli_config.h is stale -- re-run "
|
||||
f"tools/codegen/generate_cli_config.py") from e
|
||||
f"{type(config).__name__}.{group}.{pyname} is missing; the "
|
||||
f"dataclass has drifted from src/config/TrainConfig.h -- add "
|
||||
f"the field here to match the native table") from e
|
||||
|
||||
override = _VALUE_OVERRIDES.get(cli_key)
|
||||
if override is not None:
|
||||
|
||||
+27
-25
@@ -37,7 +37,7 @@ system libs), `SSPLAT_BUILD_CLI` is forced ON, CUDA archs come from
|
||||
`nvidia-smi --query-gpu=compute_cap` (override with `-DTORCH_CUDA_ARCH_LIST`),
|
||||
and the libpython link + static-libstdc++/nftw interposition workarounds are
|
||||
skipped (they exist only because of libtorch). Generated headers
|
||||
(`src/generated/`, `app/generated/`, `src/instantiations/`) are committed, so a fresh
|
||||
(`src/generated/`, `src/instantiations/`) are committed, so a fresh
|
||||
checkout builds with no Python at all. With Torch present the extension build
|
||||
is unchanged (shared `libcsrc`, same flags as before).
|
||||
|
||||
@@ -128,8 +128,7 @@ extraction (see below).
|
||||
| `WriterPool.h` | Bounded-queue worker threads that JPEG/PNG-encode and write frames and masks off the calling thread. Used by `FrameExtract`, `ssplat sam track` and the GUI's folder-masking loop. Encoding a 1080p mask through stb's deflate is ~75 ms — a third of a SAM 2.1 Tiny frame — and none of it needs the GPU, so a caller that writes inline sets the frame rate with zlib. The queue bound is what keeps a slow disk applying back-pressure instead of growing until memory runs out. |
|
||||
| `HttpServer.h/.cpp` | Minimal HTTP/1.0 GET server (POSIX sockets; winsock shim compiles but untested). Serial request handling — parity with Python's non-threading `HTTPServer`. |
|
||||
| `Viewer.h/.cpp` | Web-viewer server port (viewer/server.py + http_server.py + render_worker.py + annotation.py): latest-wins render worker, `get_outputs` viewer subset, `engine_blit_view` GPU annotation/colormap, stb JPEG encode. Serves the **unchanged** `viewer.html` (embedded at configure time via CMake hex; `SSPLAT_VIEWER_HTML=<path>` env overrides for dev). `/pick?px=&py=&<camera params>` returns the 3D point under a pixel as JSON for viewer.html's double-click centering (the Python server has no /pick; the client treats non-OK responses as a no-op). |
|
||||
| `generated/cli_config.h` | AUTO-GENERATED — do not edit. `SsplatConfig` struct (all 189 config fields, defaults baked), `SSPLAT_CONFIG_FIELDS(X)` X-macro flag table, `ssplat_apply_preset()`. |
|
||||
| `../../../../generate_cli_config.py` | The generator. AST-parses the Python config dataclasses (no torch import → works on fresh checkout). Run by `build_develop.bash`. |
|
||||
| `../config/TrainConfig.h` | Hand-written, the training config's single source of truth: the `SSPLAT_CONFIG_FIELDS(X)` X-macro flag table (190 rows), `struct SsplatConfig` expanded from it, `kSsplatPresets` and `ssplat_apply_preset()`. |
|
||||
| `../external/` | All vendored third-party code (marked `linguist-vendored` in `.gitattributes` along with the generated dirs): `stb_image.h`/`stb_image_write.h` (images), `npy.hpp` (checkpoints), `miniz.c/.h` (zip reading for the Metashape `.psx` camera table; compiled into `ssplat` only). |
|
||||
|
||||
Debug: `SSPLAT_DUMP_CAMERAS=<path> ssplat train ...` dumps parsed + post-split
|
||||
@@ -167,30 +166,33 @@ HTTP/viewer.
|
||||
- `splat.ply` reader expects float32 binary-little-endian properties (what
|
||||
both the Python trainer and `EngineCheckpoint.cpp` write).
|
||||
|
||||
## Config codegen (source of truth = Python dataclasses)
|
||||
## The config table (source of truth = `src/config/TrainConfig.h`)
|
||||
|
||||
`tools/codegen/generate_cli_config.py` parses `TrainerConfig` (+ preset
|
||||
subclasses) in `modules/trainer.py` and the nested
|
||||
`SpirulaeSplatDataParserConfig` / `SpirulaeSplatDataManagerConfig` /
|
||||
`SpirulaeSplatModelConfig` / `OptimizerConfig`. Field docstrings become help
|
||||
text; preset `default_factory` lambdas become `ssplat_apply_preset` branches.
|
||||
Hand-written, one `SSPLAT_CONFIG_FIELDS` row per flag:
|
||||
`X(type, member, default, group, choices, help)`. `struct SsplatConfig` is
|
||||
expanded from the same table, so the declaration and the metadata cannot
|
||||
drift. Add a row and the flag appears in the CLI parser, `--help`, the GUI's
|
||||
"All Options" editor, `config.json` and the pybind module.
|
||||
|
||||
- **Collisions**: flattening drops group prefixes; the generator hard-errors
|
||||
on un-listed collisions. Resolve in its `RENAMES` dict. Currently only
|
||||
`datamanager.split_batch` → `--dm-split-batch` (the model one keeps the
|
||||
plain name; datamanager's is the legacy Python-path OOM workaround, no-op
|
||||
on the managed path).
|
||||
- **Type mapping**: `Optional[int/float]` → `std::optional`;
|
||||
`Optional[str/Path]` → `std::string` with `""` = None; `Literal[str...]` →
|
||||
string + validated choices; `Literal[True,False,None]` → `optional<bool>`;
|
||||
`Tuple[...]` → `std::array`; `Union[bool,int]`
|
||||
(`rescale_camera_to_fit`) via `TYPE_OVERRIDES` → float (0=off, -1=auto,
|
||||
>0=factor).
|
||||
- **Verified** (2026-07-10): all 189 defaults + all 7 presets match
|
||||
instantiated Python configs; a full `in-the-wild` command line with mixed
|
||||
flag styles resolves identically to `tyro.cli` (189/189 fields).
|
||||
Verification scripts were session-scratch; the approach: flatten a tyro
|
||||
config, diff against the CLI's `config.json` dump.
|
||||
- **Flag names**: `member` stringified. `-` and `_` are interchangeable, so
|
||||
`--sh-degree` sets `sh_degree`. A flag cannot drift from its member.
|
||||
- **`config.json` keys**: `ssplat_json_key(flag)`, the identity except
|
||||
`dm_split_batch` → `split_batch`. `config.json` is read back by
|
||||
`ssplat mesh` and `--resume`, and the datamanager field would otherwise
|
||||
collide with `model.split_batch` (datamanager's is the legacy Python-path
|
||||
OOM workaround, a no-op on the managed path). New fields never need an
|
||||
entry — the shim is compatibility, not a mechanism.
|
||||
- **Commas**: macro arguments split on them, so `std::array<T, N>` fields use
|
||||
the `SsplatVec3i` / `SsplatVec3f` aliases and the `ssplat_v3i()` /
|
||||
`ssplat_v3f()` makers.
|
||||
- **Grouping**: rows must stay contiguous per group — `--help` and the GUI
|
||||
stream group headers as they walk the table rather than sorting first.
|
||||
|
||||
This was generated from the Python dataclasses by `generate_cli_config.py`
|
||||
until 2026-08-04. The migration was verified by diffing `train --help` for
|
||||
all 7 presets and a run's `config.json` against the last generated build:
|
||||
byte-identical, so all 190 defaults, preset overrides, choices, help strings
|
||||
and JSON keys carried over unchanged.
|
||||
|
||||
## Python → C++ port mapping (all in `TrainerCore.cpp` unless noted)
|
||||
|
||||
|
||||
+10
-7
@@ -442,27 +442,30 @@ template <typename T, size_t N> std::string json_str(const std::array<T, N>& v)
|
||||
|
||||
} // namespace
|
||||
|
||||
// Nested-by-group dump, using the original Python field names. Close to the
|
||||
// trainer's config.json (enough for humans / tooling; the Python --resume
|
||||
// path needs a tyro round-trip and is not supported by the CLI yet).
|
||||
// Nested-by-group dump. Keys are ssplat_json_key(flag) -- see the note on it
|
||||
// in config/TrainConfig.h; this is a read-back format (`ssplat mesh` and
|
||||
// --resume parse it), so the key set is not free to change.
|
||||
void save_config_json(const SsplatConfig& c, const fs::path& out_dir,
|
||||
const std::string& preset) {
|
||||
FILE* f = std::fopen((out_dir / "config.json").string().c_str(), "w");
|
||||
if (!f) throw std::runtime_error("cannot write config.json");
|
||||
std::fprintf(f, "{\n \"preset\": \"%s\"", preset.c_str());
|
||||
const char* open_group = "";
|
||||
#define SSPLAT_DUMP(member, cli_key, pyname, group, choices, help) \
|
||||
#define SSPLAT_DUMP(type, member, default_, group, choices, help) \
|
||||
{ \
|
||||
const char* key = ssplat_json_key(#member); \
|
||||
if (std::strcmp(group, "trainer") == 0) { \
|
||||
std::fprintf(f, ",\n \"%s\": %s", pyname, json_str(c.member).c_str()); \
|
||||
std::fprintf(f, ",\n \"%s\": %s", key, json_str(c.member).c_str()); \
|
||||
} else { \
|
||||
if (std::strcmp(open_group, group) != 0) { \
|
||||
if (*open_group) std::fprintf(f, "\n }"); \
|
||||
std::fprintf(f, ",\n \"%s\": {", group); \
|
||||
open_group = group; \
|
||||
std::fprintf(f, "\n \"%s\": %s", pyname, json_str(c.member).c_str()); \
|
||||
std::fprintf(f, "\n \"%s\": %s", key, json_str(c.member).c_str()); \
|
||||
} else { \
|
||||
std::fprintf(f, ",\n \"%s\": %s", pyname, json_str(c.member).c_str()); \
|
||||
std::fprintf(f, ",\n \"%s\": %s", key, json_str(c.member).c_str()); \
|
||||
} \
|
||||
} \
|
||||
}
|
||||
SSPLAT_CONFIG_FIELDS(SSPLAT_DUMP)
|
||||
#undef SSPLAT_DUMP
|
||||
|
||||
@@ -35,7 +35,7 @@
|
||||
#include "engine/Engine.h"
|
||||
#include "data/DatasetParser.h"
|
||||
#include "app/webviewer/RenderWorker.h"
|
||||
#include "app/generated/cli_config.h"
|
||||
#include "config/TrainConfig.h"
|
||||
|
||||
#include <array>
|
||||
#include <atomic>
|
||||
@@ -120,7 +120,7 @@ build_loss_weights(const SsplatConfig& c, int step);
|
||||
EngineStepConfig build_step_config(const SsplatConfig& c, const RunState& st,
|
||||
int step);
|
||||
|
||||
// Nested-by-group config.json dump using the original Python field names.
|
||||
// Nested-by-group config.json dump, keyed by ssplat_json_key(flag).
|
||||
void save_config_json(const SsplatConfig& c, const std::filesystem::path& out_dir,
|
||||
const std::string& preset);
|
||||
|
||||
|
||||
@@ -7,8 +7,7 @@
|
||||
// academic-baseline). Flags are the FLATTENED training config fields
|
||||
// (--sh-degree, not --model.sh-degree); '-' and '_' are interchangeable;
|
||||
// booleans take a value (--warp-to-pinhole 1). The config struct, flag
|
||||
// table, and preset appliers are code-generated from the Python config
|
||||
// dataclasses -- see generated/cli_config.h and generate_cli_config.py.
|
||||
// table, and preset appliers all come from config/TrainConfig.h.
|
||||
//
|
||||
// The engine plumbing (dataset -> seeding -> per-step configs -> train loop)
|
||||
// lives in TrainerCore.{h,cpp}, shared with the native GUI (gui/). This file
|
||||
@@ -121,8 +120,8 @@ void check_choices(const T&, const std::string&, const char*) {}
|
||||
// Returns false if the key is unknown.
|
||||
bool set_config_field(SsplatConfig& c, const std::string& key,
|
||||
int argc, char** argv, int& i) {
|
||||
#define SSPLAT_TRY_SET(member, cli_key, pyname, group, choices, help) \
|
||||
if (key == cli_key) { \
|
||||
#define SSPLAT_TRY_SET(type, member, default_, group, choices, help) \
|
||||
if (key == #member) { \
|
||||
consume(c.member, key, argc, argv, i); \
|
||||
check_choices(c.member, key, choices); \
|
||||
return true; \
|
||||
@@ -206,7 +205,7 @@ void print_help(const char* argv0, const SsplatConfig& c) {
|
||||
std::printf("\nflags ('-' and '_' interchangeable; bools take 0/1; 'none' clears "
|
||||
"optional values;\n defaults shown for the selected preset):\n");
|
||||
const char* cur_group = "";
|
||||
#define SSPLAT_PRINT_HELP(member, cli_key, pyname, group, choices, help) \
|
||||
#define SSPLAT_PRINT_HELP(type, member, default_, group, choices, help) \
|
||||
if (std::strcmp(cur_group, group) != 0) { \
|
||||
cur_group = group; \
|
||||
std::printf("\n [%s]\n", group); \
|
||||
@@ -217,7 +216,7 @@ void print_help(const char* argv0, const SsplatConfig& c) {
|
||||
if (dot != std::string::npos) h = h.substr(0, dot + 1); \
|
||||
if (h.size() > 110) h = h.substr(0, 107) + "..."; \
|
||||
std::string ch = choices; \
|
||||
std::string key_disp = cli_key; \
|
||||
std::string key_disp = #member; \
|
||||
for (auto& ck : key_disp) if (ck == '_') ck = '-'; \
|
||||
std::printf(" --%-38s [%s]%s%s\n %s\n", key_disp.c_str(), \
|
||||
value_str(c.member).c_str(), \
|
||||
|
||||
@@ -1,545 +0,0 @@
|
||||
#pragma once
|
||||
|
||||
// AUTO-GENERATED by tools/codegen/generate_cli_config.py -- DO NOT EDIT.
|
||||
// Source of truth: the Python training config dataclasses
|
||||
// (TrainerConfig + presets, SpirulaeSplatDataParserConfig,
|
||||
// SpirulaeSplatDataManagerConfig, SpirulaeSplatModelConfig,
|
||||
// OptimizerConfig). Re-run the generator after editing those.
|
||||
|
||||
#include <array>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
|
||||
struct SsplatConfig {
|
||||
// ==== trainer (TrainerConfig) ====
|
||||
std::string data = {}; // REQUIRED
|
||||
std::string resume = "";
|
||||
std::string output_dir_prefix = "outputs";
|
||||
std::string output_dir_name = "";
|
||||
int steps_per_save = 2000;
|
||||
bool save_only_latest_checkpoint = true;
|
||||
bool save_full_checkpoint = false;
|
||||
bool save_eval_images = false;
|
||||
int num_iterations = 30000;
|
||||
int viewer_port = 7007;
|
||||
bool disable_viewer = false;
|
||||
bool keep_viewer_alive = true;
|
||||
// ==== dataparser (SpirulaeSplatDataParserConfig) ====
|
||||
std::string data_format = "";
|
||||
std::string colmap_recon_dir = "";
|
||||
std::string image_dir = "images";
|
||||
std::string mask_dir = "masks";
|
||||
std::string depth_dir = "depths";
|
||||
std::string normal_dir = "normals";
|
||||
std::string metashape_xml = "";
|
||||
std::string metashape_ply = "";
|
||||
std::string metashape_psx = "";
|
||||
float rescale_camera_to_fit = 0.0f;
|
||||
std::string downscale_rounding_mode = "floor";
|
||||
float scene_scale = 1.0f;
|
||||
std::string orientation_method = "up";
|
||||
std::string center_method = "poses";
|
||||
bool auto_scale_poses = true;
|
||||
float outlier_threshold = std::numeric_limits<float>::infinity();
|
||||
std::string train_frame = "points";
|
||||
std::string eval_mode = "all";
|
||||
float train_split_fraction = 0.9f;
|
||||
int eval_interval = 8;
|
||||
float depth_unit_scale_factor = 0.001f;
|
||||
float validation_fraction = 0.0f;
|
||||
// ==== datamanager (SpirulaeSplatDataManagerConfig) ====
|
||||
int max_batch_per_epoch = 800;
|
||||
bool dm_split_batch = false;
|
||||
std::string cache_images = "disk";
|
||||
bool load_depths = true;
|
||||
bool load_normals = true;
|
||||
float mask_boundary_offset = 0.0f;
|
||||
bool warp_to_pinhole = false;
|
||||
bool warp_spherical_to_pinhole = true;
|
||||
bool deblur_training_images = false;
|
||||
// ==== model (SpirulaeSplatModelConfig) ====
|
||||
std::string primitive = "3dgs";
|
||||
int sh_degree = 3;
|
||||
int sh_degree_warmup_every = 1000;
|
||||
std::string background_mode = "black";
|
||||
int background_noise_warmup = 2000;
|
||||
float background_noise_pre_warmup = 0.25f;
|
||||
int background_sh_degree = 4;
|
||||
std::optional<float> relative_scale = std::nullopt;
|
||||
float l1_weight = 1.0f;
|
||||
float l2_weight = 0.0f;
|
||||
float ssim_lambda = 0.2f;
|
||||
float l1_weight_y = 0.0f;
|
||||
float l2_weight_y = 0.0f;
|
||||
float l2_weight_u = 0.0f;
|
||||
float l2_weight_v = 0.0f;
|
||||
int num_loss_scales = 0;
|
||||
int loss_scale_min_pixels = 1920;
|
||||
bool use_camera_optimizer = false;
|
||||
bool packed = true;
|
||||
bool use_bvh = false;
|
||||
bool use_fused_proj_bwd_optim = true;
|
||||
bool split_batch = true;
|
||||
int quantization_level = 1;
|
||||
std::string optimizer_offload = "";
|
||||
int resolution_schedule = 3000;
|
||||
int num_downscales = 0;
|
||||
bool use_mcmc = true;
|
||||
bool preallocate_splat_tensors = true;
|
||||
int cap_max = 1000000;
|
||||
float min_init_fraction = 0.0f;
|
||||
int refine_every = 100;
|
||||
int refine_start_iter = 500;
|
||||
int refine_stop_num_iter = 5000;
|
||||
int refine_stop_iter = 25000;
|
||||
float noise_lr = 80.0f;
|
||||
float noise_lr_final = 0.8f;
|
||||
float min_opacity = 0.005f;
|
||||
float growth_factor = 1.05f;
|
||||
bool use_revised_densification = true;
|
||||
std::string densify_score_mode = "mean";
|
||||
float densify_score_blend_world_grad = 0.0f;
|
||||
std::string densify_loss_map_mode = "ssim_structure";
|
||||
float densify_robust_edge_aware_quantile = 0.9f;
|
||||
bool use_long_axis_split = true;
|
||||
std::array<float, 3> long_axis_split_opacity_k = {0.6f, 0.6f, 4500.0f};
|
||||
float relocate_screen_size = std::numeric_limits<float>::infinity();
|
||||
float max_screen_size = 0.3f;
|
||||
float max_screen_size_clip_hardness = 1.5f;
|
||||
float max_world_size = std::numeric_limits<float>::infinity();
|
||||
int reset_alpha_every = 30;
|
||||
bool use_bilateral_grid = true;
|
||||
std::array<int, 3> bilagrid_shape = {16, 16, 8};
|
||||
std::string bilagrid_type = "ppisp";
|
||||
bool use_bilateral_grid_for_geometry = true;
|
||||
std::array<int, 3> bilagrid_shape_geometry = {8, 8, 4};
|
||||
bool use_adagrad_bilagrid_optim = true;
|
||||
float bilagrid_tv_loss_weight = 10.0f;
|
||||
float color_shift_reg_weight = 0.0f;
|
||||
int color_shift_reg_ema_period = 750;
|
||||
float bilagrid_tv_loss_weight_geometry = 10.0f;
|
||||
bool use_ppisp = true;
|
||||
std::string ppisp_param_type = "no_crf";
|
||||
bool use_adagrad_ppisp_optim = true;
|
||||
bool apply_ppisp_before_bilagrid = true;
|
||||
float ppisp_reg_exposure_mean = 1.0f;
|
||||
float ppisp_reg_vig_center = 0.02f;
|
||||
float ppisp_reg_vig_non_pos = 0.01f;
|
||||
float ppisp_reg_vig_channel_var = 0.1f;
|
||||
float ppisp_reg_color_mean = 1.0f;
|
||||
float ppisp_reg_crf_channel_var = 0.1f;
|
||||
bool image_color_is_linear = false;
|
||||
std::string image_color_gamut = "";
|
||||
std::optional<bool> splat_color_is_linear = std::nullopt;
|
||||
std::string splat_color_gamut = "";
|
||||
std::optional<bool> convert_initial_point_cloud_color = std::nullopt;
|
||||
std::optional<float> scale_init = std::nullopt;
|
||||
std::optional<float> opacity_init = std::nullopt;
|
||||
bool suppress_initial_scales = false;
|
||||
float scale_regularization_weight = 0.0f;
|
||||
float max_gauss_ratio = 10.0f;
|
||||
float depth_distortion_reg = 0.0f;
|
||||
float normal_distortion_reg = 0.0f;
|
||||
float rgb_distortion_reg = 0.0f;
|
||||
int distortion_reg_warmup = 6000;
|
||||
float normal_reg_weight = 0.04f;
|
||||
int normal_reg_warmup = 6000;
|
||||
float alpha_reg_weight = 0.0f;
|
||||
int alpha_reg_warmup = 12000;
|
||||
int reg_warmup_length = 0;
|
||||
bool apply_loss_for_mask = false;
|
||||
bool enable_sky_masking = true;
|
||||
float alpha_loss_weight = 0.01f;
|
||||
float alpha_loss_weight_under = 0.0f;
|
||||
float opacity_reg = 0.01f;
|
||||
float scale_reg = 0.01f;
|
||||
float opacity_decay = 0.0f;
|
||||
float scale_decay = 0.0f;
|
||||
float erank_reg = 0.0f;
|
||||
float erank_reg_s3 = 0.0f;
|
||||
float quat_norm_reg = 0.01f;
|
||||
float sh_reg = 0.001f;
|
||||
float overexposure_reg = 0.0f;
|
||||
int supervision_warmup = 0;
|
||||
float depth_supervision_weight = 0.0f;
|
||||
bool input_depth_is_ray_depth = false;
|
||||
float normal_supervision_weight = 0.01f;
|
||||
float mean_median_depth_weight = 0.0f;
|
||||
float median_depth_normal_reg_weight = 0.0f;
|
||||
float median_normal_supervision_weight = 0.0f;
|
||||
float median_render_normal_reg_weight = 0.0f;
|
||||
int median_warmup = 6000;
|
||||
std::string overfit_score_aggregation_mode = "min";
|
||||
int validation_loss_average_window = 500;
|
||||
int early_stop_patience = 1000;
|
||||
int early_stop_warmup = 12000;
|
||||
// ==== optimizer (OptimizerConfig) ====
|
||||
std::optional<int> max_steps = std::nullopt;
|
||||
bool use_scale_agnostic_mean = true;
|
||||
bool use_per_splat_bias_correction = true;
|
||||
float means_lr = 0.000128f;
|
||||
std::optional<float> means_lr_final = 1.6e-06f;
|
||||
float scales_lr = 0.02f;
|
||||
std::optional<float> scales_lr_final = 0.005f;
|
||||
float quats_lr = 0.0015f;
|
||||
float opacities_lr = 0.025f;
|
||||
float features_dc_lr = 0.005f;
|
||||
float features_sh_lr = 0.00025f;
|
||||
float background_dc_lr = 0.0025f;
|
||||
float background_sh_lr = 0.0005f;
|
||||
float bilagrid_lr = 0.002f;
|
||||
std::optional<float> bilagrid_lr_final = 0.0001f;
|
||||
int bilagrid_lr_warmup = 1000;
|
||||
float bilagrid_depth_lr = 0.002f;
|
||||
std::optional<float> bilagrid_depth_lr_final = 0.0001f;
|
||||
int bilagrid_depth_lr_warmup = 2000;
|
||||
float bilagrid_normal_lr = 0.0005f;
|
||||
std::optional<float> bilagrid_normal_lr_final = 4e-05f;
|
||||
int bilagrid_normal_lr_warmup = 2000;
|
||||
float bilagrid_adagrad_lr = 0.04f;
|
||||
float bilagrid_adagrad_depth_lr = 0.04f;
|
||||
float bilagrid_adagrad_normal_lr = 0.01f;
|
||||
float ppisp_lr = 0.002f;
|
||||
std::optional<float> ppisp_lr_final = 2e-05f;
|
||||
int ppisp_lr_warmup = 500;
|
||||
float ppisp_adagrad_lr = 0.1f;
|
||||
float camera_opt_lr = 0.0001f;
|
||||
std::optional<float> camera_opt_lr_final = 5e-07f;
|
||||
int camera_opt_lr_warmup = 1000;
|
||||
};
|
||||
|
||||
// X(member, cli_key, pyname, group, choices, help). cli_key uses '_';
|
||||
// the parser treats '-' and '_' as equivalent. pyname is the field's
|
||||
// original name inside its Python config class (differs from cli_key
|
||||
// only for RENAMES entries); group is the Python sub-config it lives
|
||||
// in. choices is a '|' list for string fields ('' = free-form);
|
||||
// 'none' selects the empty string.
|
||||
#define SSPLAT_CONFIG_FIELDS(X) \
|
||||
X(data, "data", "data", "trainer", "", "Path to dataset. Can be a Nerfstudio or a COLMAP dataset.") \
|
||||
X(resume, "resume", "resume", "trainer", "none", "Resume training from a checkpoint. Pass a run output dir (latest step-*.ckpt is used) or a specific step-*.ckpt dir. The run's config.json supplies the architecture/model/data config (run-control flags like num_iterations / save cadence / viewer are still taken from the CLI); requires the checkpoint to have been written with save_full_checkpoint.") \
|
||||
X(output_dir_prefix, "output_dir_prefix", "output_dir_prefix", "trainer", "", "Prefix to output directory") \
|
||||
X(output_dir_name, "output_dir_name", "output_dir_name", "trainer", "none", "Output directory name relative to output_dir_prefix. If not specified, will set a generic combining current timestamp and dataset name.") \
|
||||
X(steps_per_save, "steps_per_save", "steps_per_save", "trainer", "", "Save checkpoint every this number of steps. If -1, save only at the end. If zero, never save (used in benchmark).") \
|
||||
X(save_only_latest_checkpoint, "save_only_latest_checkpoint", "save_only_latest_checkpoint", "trainer", "", "Whether to save only last checkpoint") \
|
||||
X(save_full_checkpoint, "save_full_checkpoint", "save_full_checkpoint", "trainer", "", "If True, each checkpoint's `state.tar` additionally includes the Resume slots -- world raw parameters (at max_num_splats) plus all optimizer state -- making the checkpoint sufficient to resume training, not just for inference. If False, only the Always slots (appearance/inference params) are saved alongside `splat.ply`.") \
|
||||
X(save_eval_images, "save_eval_images", "save_eval_images", "trainer", "", "Whether to save eval images at end of training") \
|
||||
X(num_iterations, "num_iterations", "num_iterations", "trainer", "", "Number of training iterations") \
|
||||
X(viewer_port, "viewer_port", "viewer_port", "trainer", "", "Port used by the web viewer") \
|
||||
X(disable_viewer, "disable_viewer", "disable_viewer", "trainer", "", "If True, ss_trainer skips starting the viewer thread. Used by ss_benchmark so each scene runs without competing for the viewer port.") \
|
||||
X(keep_viewer_alive, "keep_viewer_alive", "keep_viewer_alive", "trainer", "", "If True, ss_trainer keeps the process (and thus the viewer) running after training + eval finish, so the result can still be inspected in the browser. Press Ctrl-C to exit. Ignored when disable_viewer=True.") \
|
||||
X(data_format, "data_format", "data_format", "dataparser", "colmap|nerfstudio|metashape|none", "Data format, leave None to auto detect") \
|
||||
X(colmap_recon_dir, "colmap_recon_dir", "colmap_recon_dir", "dataparser", "none", "Path to COLMAP reconstruction relative to dataset directory (e.g. sparse/0). Will auto detect if not specified (picking the model with the most registered images when several exist).") \
|
||||
X(image_dir, "image_dir", "image_dir", "dataparser", "", "Path to images relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(mask_dir, "mask_dir", "mask_dir", "dataparser", "", "Path to masks relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(depth_dir, "depth_dir", "depth_dir", "dataparser", "", "Path to depth maps relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(normal_dir, "normal_dir", "normal_dir", "dataparser", "", "Path to normal maps relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(metashape_xml, "metashape_xml", "metashape_xml", "dataparser", "none", "Path to the Metashape xml file. Will automatically detect if not specified.") \
|
||||
X(metashape_ply, "metashape_ply", "metashape_ply", "dataparser", "none", "Path to the Metashape point export ply file. Will automatically detect if not specified.") \
|
||||
X(metashape_psx, "metashape_psx", "metashape_psx", "dataparser", "none", "Path to Metashape PSX file, used to resolve file name ambiguity when there are multiple images with the same file name") \
|
||||
X(rescale_camera_to_fit, "rescale_camera_to_fit", "rescale_camera_to_fit", "dataparser", "", "Whether to check if image resolution match camera resolution and scale camera intrinsics accordingly if not. Set this to a number to divide intrinsics by that number, e.g. Mip-NeRF 360 and Zip-NeRF with images_(2|4) Set this to True to detect resolution, e.g. tankt_db [CLI: 0 = off, -1 = auto-detect from image resolution, > 0 = divide intrinsics by this]") \
|
||||
X(downscale_rounding_mode, "downscale_rounding_mode", "downscale_rounding_mode", "dataparser", "floor|ceil|round", "Rounding mode applied to camera width/height when dividing by `rescale_camera_to_fit`. Use `round` to match the convention used by most image downscalers (e.g. Mip-NeRF 360 images_(2|4|8)).") \
|
||||
X(scene_scale, "scene_scale", "scene_scale", "dataparser", "", "How much to scale the region of interest by.") \
|
||||
X(orientation_method, "orientation_method", "orientation_method", "dataparser", "pca|up|vertical|none|gsplat", "The method to use for orientation.") \
|
||||
X(center_method, "center_method", "center_method", "dataparser", "poses|focus|none|gsplat", "The method to use to center the poses.") \
|
||||
X(auto_scale_poses, "auto_scale_poses", "auto_scale_poses", "dataparser", "", "Whether to automatically scale the poses to fit in +/- 1 bounding box.") \
|
||||
X(outlier_threshold, "outlier_threshold", "outlier_threshold", "dataparser", "", "Threshold to reject outlier camera poses.") \
|
||||
X(train_frame, "train_frame", "train_frame", "dataparser", "normalized|camera|points", "Coordinate frame in which splats are trained.") \
|
||||
X(eval_mode, "eval_mode", "eval_mode", "dataparser", "fraction|filename|interval|all", "The method to use for splitting the dataset into train and eval. Fraction splits based on a percentage for train and the remaining for eval. Filename splits based on filenames containing train/eval. Interval uses every nth frame for eval. All uses all the images for any split.") \
|
||||
X(train_split_fraction, "train_split_fraction", "train_split_fraction", "dataparser", "", "The percentage of the dataset to use for training. Only used when eval_mode is train-split-fraction.") \
|
||||
X(eval_interval, "eval_interval", "eval_interval", "dataparser", "", "The interval between frames to use for eval. Only used when eval_mode is eval-interval.") \
|
||||
X(depth_unit_scale_factor, "depth_unit_scale_factor", "depth_unit_scale_factor", "dataparser", "", "Scales the depth values to meters. Default value is 0.001 for a millimeter to meter conversion.") \
|
||||
X(validation_fraction, "validation_fraction", "validation_fraction", "dataparser", "", "Use this fraction of training images for validation. Stop training when performance on validation images start to drop.") \
|
||||
X(max_batch_per_epoch, "max_batch_per_epoch", "max_batch_per_epoch", "datamanager", "", "Maximum number of batches per epoch, used for configuring batch size") \
|
||||
X(dm_split_batch, "dm_split_batch", "split_batch", "datamanager", "", "Whether to one large batch into many small batches to avoid OOM, at cost of slower training") \
|
||||
X(cache_images, "cache_images", "cache_images", "datamanager", "cpu-pageable|cpu|gpu|disk", "Image cache location. If \"cpu\", caches on cpu. If \"gpu\", caches on device. If \"cpu-pageable\", cache on cpu pageable memory (saves RAM but may cause error if spill to swap memory). If \"disk\", cache on disk (limited support).") \
|
||||
X(load_depths, "load_depths", "load_depths", "datamanager", "", "Whether to load depth maps, if exist") \
|
||||
X(load_normals, "load_normals", "load_normals", "datamanager", "", "Whether to load normal maps, if exist") \
|
||||
X(mask_boundary_offset, "mask_boundary_offset", "mask_boundary_offset", "datamanager", "", "Signed boundary offset applied to binarized masks at decode time, as a fraction of sqrt(W*H) of the decoded mask. Positive dilates (grows) foreground, negative erodes (shrinks). Runs on CPU during data loading via separable Felzenszwalb-Huttenlocher squared-Euclidean DT (exact, O(N) per row + col).") \
|
||||
X(warp_to_pinhole, "warp_to_pinhole", "warp_to_pinhole", "datamanager", "", "Whether to split a fisheye image into 5 undistorted pinhole images. Can sometimes give better quality and compatibility for dataset captured by fisheye/360 cameras.") \
|
||||
X(warp_spherical_to_pinhole, "warp_spherical_to_pinhole", "warp_spherical_to_pinhole", "datamanager", "", "Whether to split an equirectangular (spherical panorama) image into 6 pinhole cubemap faces for training. When True (default), equirectangular images are split into 6 undistorted pinhole sub-images (the historical behavior). When False, train directly on the equirectangular image using an equirectangular projection/unprojection plugged into the linear/UT projection pipeline (supports 3dgs/mip and 3dgut primitives). Direct equirectangular training does not support depth/normal supervision.") \
|
||||
X(deblur_training_images, "deblur_training_images", "deblur_training_images", "datamanager", "", "Whether to use a custom trained deep learning model to deblur images before training") \
|
||||
X(primitive, "primitive", "primitive", "model", "3dgs|mip|3dgut", "Splat primitive to use") \
|
||||
X(sh_degree, "sh_degree", "sh_degree", "model", "", "Maximum degree of spherical harmonics to use.") \
|
||||
X(sh_degree_warmup_every, "sh_degree_warmup_every", "sh_degree_warmup_every", "model", "", "Increase SH degree every this number of iterations") \
|
||||
X(background_mode, "background_mode", "background_mode", "model", "black|noise|sh", "Background mode, black per convention, noise to discourage transparency, sh for skybox.") \
|
||||
X(background_noise_warmup, "background_noise_warmup", "background_noise_warmup", "model", "", "Number of steps to warmup background noise. This applies when background_mode is noise") \
|
||||
X(background_noise_pre_warmup, "background_noise_pre_warmup", "background_noise_pre_warmup", "model", "", "Weight of background noise at start of training (0 to 1). Higher value reduce the chance of washing away splat opacities near the beginning of training.") \
|
||||
X(background_sh_degree, "background_sh_degree", "background_sh_degree", "model", "", "SH degree for background color, only used when background_mode is sh.") \
|
||||
X(relative_scale, "relative_scale", "relative_scale", "model", "", "Manually set scale when a scene is poorly scaled, i.e. increase this for large datasets. If not set, will use a scale agnostic optimizer. To prevent this, set it to 1.0.") \
|
||||
X(l1_weight, "l1_weight", "l1_weight", "model", "", "Weight of L1 loss, default 1.0") \
|
||||
X(l2_weight, "l2_weight", "l2_weight", "model", "", "Weight of L2 loss, default 0.0") \
|
||||
X(ssim_lambda, "ssim_lambda", "ssim_lambda", "model", "", "Weight of ssim loss; 0.2 for academic baseline, higher for potentially more high-frequency details, lower for less blurry background in outdoor scenes") \
|
||||
X(l1_weight_y, "l1_weight_y", "l1_weight_y", "model", "", "Weight of per-pixel BT.601 luma (Y) L1 loss.") \
|
||||
X(l2_weight_y, "l2_weight_y", "l2_weight_y", "model", "", "Weight of per-pixel BT.601 luma (Y) L2 loss.") \
|
||||
X(l2_weight_u, "l2_weight_u", "l2_weight_u", "model", "", "Weight of per-pixel BT.601 chroma U L2 loss.") \
|
||||
X(l2_weight_v, "l2_weight_v", "l2_weight_v", "model", "", "Weight of per-pixel BT.601 chroma V L2 loss.") \
|
||||
X(num_loss_scales, "num_loss_scales", "num_loss_scales", "model", "", "Number of scales for image loss. For multi-scale loss, image is downscaled by 2 this number of times, and losses are averaged across scales. Improves convergence for high-resolution images.") \
|
||||
X(loss_scale_min_pixels, "loss_scale_min_pixels", "loss_scale_min_pixels", "model", "", "If positive, overrides num_loss_scales per image based on resolution, in units of pixels. num_loss_scales is chosen so the smallest image dimension is halved down toward (but not below) this many pixels. e.g. with 2000: min dim 1999 -> num_loss_scales=0, 2000 -> 1, 4000 -> 2, 8000 -> 3, etc. Adapts per training step, so datasets with mixed image resolutions get the right count per image automatically.") \
|
||||
X(use_camera_optimizer, "use_camera_optimizer", "use_camera_optimizer", "model", "", "Whether to use camera optimizer Note: this only works well in patch batching mode") \
|
||||
X(packed, "packed", "packed", "model", "", "Pack projection outputs, reduce VRAM usage at large batch size but can be slightly slower") \
|
||||
X(use_bvh, "use_bvh", "use_bvh", "model", "", "Use BVH for splat-patch intersection test, may be faster when batching large number of small patches") \
|
||||
X(use_fused_proj_bwd_optim, "use_fused_proj_bwd_optim", "use_fused_proj_bwd_optim", "model", "", "Whether to use fused projection backward and optimizer. More memory efficient for large number of Gaussians, with slight performance hit.") \
|
||||
X(split_batch, "split_batch", "split_batch", "model", "", "Split the camera batch into one-camera sub-batches inside the C++ train step. Per-splat grads accumulate via atomicAdd across sub-batches; a single optim+densify pass at the end consumes the accumulator with grad_scale = 1/B. Drops peak VRAM for the immediate projection / rasterization buffers by roughly 1/B. Per-image grad magnitude vs regularization weight stays batch-size invariant. Not compatible with use_fused_proj_bwd_optim or with the warped train-step path.") \
|
||||
X(quantization_level, "quantization_level", "quantization_level", "model", "", "SH quantization level: a single int that selects one of two (param bits, optim bits) configurations. 0 = off : 32-bit param, fp32 optim 1 = light : 16-bit param, 8-bit packed optim (2 B / cell) Collapsing the prior independent param+optim bit controls into a single level minimizes the FPBO kernel instantiations.") \
|
||||
X(optimizer_offload, "optimizer_offload", "optimizer_offload", "model", "sh|all|none", "Whether to offload optimizer momentum to CPU to save VRAM. This is only supported for Adam optimizer.") \
|
||||
X(resolution_schedule, "resolution_schedule", "resolution_schedule", "model", "", "training starts at 1/d resolution, every n steps this is doubled") \
|
||||
X(num_downscales, "num_downscales", "num_downscales", "model", "", "At the beginning, resolution is 1/2^d, where d is this number") \
|
||||
X(use_mcmc, "use_mcmc", "use_mcmc", "model", "", "Must be True for 3DGS methods.") \
|
||||
X(preallocate_splat_tensors, "preallocate_splat_tensors", "preallocate_splat_tensors", "model", "", "Whether to pre-allocate Gaussian attribute tensors to cap_max to avoid OOM during densification") \
|
||||
X(cap_max, "cap_max", "cap_max", "model", "", "maximum number of splats, dataset-specific tuning required") \
|
||||
X(min_init_fraction, "min_init_fraction", "min_init_fraction", "model", "", "minimum fraction of splats out of cap_max at initialization") \
|
||||
X(refine_every, "refine_every", "refine_every", "model", "", "Densify every this number of steps") \
|
||||
X(refine_start_iter, "refine_start_iter", "refine_start_iter", "model", "", "Start densification at this number of steps") \
|
||||
X(refine_stop_num_iter, "refine_stop_num_iter", "refine_stop_num_iter", "model", "", "Stop densification at this number of steps before maximum number of training iterations") \
|
||||
X(refine_stop_iter, "refine_stop_iter", "refine_stop_iter", "model", "", "Densification runs until max(this, num_iterations - refine_stop_num_iter). Without this floor, runs shorter than refine_stop_num_iter would never densify at all (num_iterations - refine_stop_num_iter goes negative), which confuses users.") \
|
||||
X(noise_lr, "noise_lr", "noise_lr", "model", "", "MCMC-like noise injection magnitude at start of training") \
|
||||
X(noise_lr_final, "noise_lr_final", "noise_lr_final", "model", "", "MCMC-like noise injection magnitude at end of training") \
|
||||
X(min_opacity, "min_opacity", "min_opacity", "model", "", "Minimum Gaussian opacity before relocation") \
|
||||
X(growth_factor, "growth_factor", "growth_factor", "model", "", "Multiply number of splats by this number at each densification step") \
|
||||
X(use_revised_densification, "use_revised_densification", "use_revised_densification", "model", "", "Whether to use revised densification instead of original MCMC.") \
|
||||
X(densify_score_mode, "densify_score_mode", "densify_score_mode", "model", "mean|max|median|geom", "How to accumulate per-splat scores across iterations for densification. \"mean\": running mean of |w|. \"max\": running max of |w|. \"median\": running median of |w| (approximation). \"geom\": running geometric mean of |w|.") \
|
||||
X(densify_score_blend_world_grad, "densify_score_blend_world_grad", "densify_score_blend_world_grad", "model", "", "Blend weight `w` in [0, 1] between the image-space loss score and the world-space gradient score for densification. The per-step score is (image-space accum_weight)^(1-w) * (||dL/dmean_world|| * max post-exp world scale)^w. The world-grad term favors world-space-large splats (e.g. distant background in unbounded outdoor scenes) that the image-space score under-weights; the geometric blend is invariant to each score's global scale so no cross-normalization is needed. 0 (default) = image-space score only, identical cost and behavior to before. 1 = world-grad score only; the per-pixel densification loss map (densify_loss_map_mode) is skipped entirely. In between, both scores are computed (one extra float per splat of VRAM).") \
|
||||
X(densify_loss_map_mode, "densify_loss_map_mode", "densify_loss_map_mode", "model", "none|loss_full|ssim_full|ssim_cs|ssim_structure|edge_aware|robust_edge_aware", "What gets accumulated into the per-pixel densification loss map. The loss map is read by raster bwd to weight the per-splat accum_weight. Only active when use_revised_densification. Modes: \"none\": no loss map (uniform alpha*T accumulation). \"loss_full\": per-pixel L1/L2 + auxiliary supervisory terms + full SSIM (luminance*contrast*structure). \"ssim_full\": full SSIM only. \"ssim_cs\": contrast*structure SSIM (no luminance). \"ssim_structure\": structure-only SSIM, biases toward pattern/edge mismatches and ignores brightness/contrast errors. \"edge_aware\": canny edge magnitude of GT rgb (Plenoxels-style, https://arxiv.org/abs/2603.08661). Biases densification toward GT edges directly, regardless of how well the splats already reconstruct them. \"robust_edge_aware\": RobustNeRF-style Tukey biweight on the BT.601 luma of |render - GT|, capped at the per-image `densify_robust_edge_aware_quantile`, then canny. Near-zero where the render already matches GT, zeroed past the quantile cutoff so distractor pixels (people/cars/operator) don't pull splats toward them, and luminance-shift tolerant since a global DC residual has no spatial gradient. For num_loss_scales > 0 the map is computed per scale and the per-scale results are averaged (matches the multi-scale loss accumulation). Affects loss_map only; training gradients and scalar losses unchanged.") \
|
||||
X(densify_robust_edge_aware_quantile, "densify_robust_edge_aware_quantile", "densify_robust_edge_aware_quantile", "model", "", "Per-image quantile of the luma residual used as the Tukey biweight cutoff in `robust_edge_aware` mode. Pixels whose residual exceeds this quantile get zero densification weight. Lower values are more aggressive outlier rejection (good for distractor-heavy datasets); higher are more permissive (good for clean datasets where real edges may produce large residuals). Ignored unless `densify_loss_map_mode == \"robust_edge_aware\"`.") \
|
||||
X(use_long_axis_split, "use_long_axis_split", "use_long_axis_split", "model", "", "whether to use long-axis split described in https://arxiv.org/abs/2508.12313 for relocation and sample add. When combined with use_revised_densification, this can give less blurry background details for unbounded outdoor scenes.") \
|
||||
X(long_axis_split_opacity_k, "long_axis_split_opacity_k", "long_axis_split_opacity_k", "model", "", "opacity split factor `k` for long-axis split, as (initial, final, warmup_steps). Each split child keeps opacity `logit^-1(k / (1 + exp(-logit_opacity) - k))`; `k` is linearly scheduled from `initial` to `final` over the first `warmup_steps` training steps, then held at `final`. Larger `k` preserves more opacity per child (denser, sharper); smaller `k` fades children faster.") \
|
||||
X(relocate_screen_size, "relocate_screen_size", "relocate_screen_size", "model", "", "if a gaussian is more than this fraction of screen space, relocate it Useful for fisheye with 3DGUT, may drop PSNR for conventional cameras For likely better quality, use max_screen_size instead") \
|
||||
X(max_screen_size, "max_screen_size", "max_screen_size", "model", "", "if a gaussian is more than this fraction of screen space, clip scale and increase opacity Intended to be an MCMC-friendly alternative of relocate_screen_size") \
|
||||
X(max_screen_size_clip_hardness, "max_screen_size_clip_hardness", "max_screen_size_clip_hardness", "model", "", "clip hardness for Gaussians with large screen space size, between 1 and infinity, larger is harder") \
|
||||
X(max_world_size, "max_world_size", "max_world_size", "model", "", "if a gaussian is more than this of world space, clip scale Useful if you see huge floaters at a distance in large indoor space") \
|
||||
X(reset_alpha_every, "reset_alpha_every", "reset_alpha_every", "model", "", "Every this many refinement steps, reset the alpha. Only applies for opaque triangle splatting.") \
|
||||
X(use_bilateral_grid, "use_bilateral_grid", "use_bilateral_grid", "model", "", "If True, use bilateral grid to handle the ISP changes in the image space. This technique was introduced in the paper 'Bilateral Guided Radiance Field Processing' (https://bilarfpro.github.io/).") \
|
||||
X(bilagrid_shape, "bilagrid_shape", "bilagrid_shape", "model", "", "Shape of the bilateral grid, typically `16 16 8`, or `8 8 4` for scenes with low-texture surfaces.") \
|
||||
X(bilagrid_type, "bilagrid_type", "bilagrid_type", "model", "affine|ppisp|loglinear", "What the bilateral grid predicts. affine: 4x3 matrix per original bilateral grid. ppisp: PPISP exposure and color parameters, generally gives less color shift but can be less numerically stable. loglinear: 3x3 linear transformation matrix with log-encoded diagonals, balances color shift and numerical stability.") \
|
||||
X(use_bilateral_grid_for_geometry, "use_bilateral_grid_for_geometry", "use_bilateral_grid_for_geometry", "model", "", "If True, use bilateral grid for depth and normal (e.g. AI generated biased ones)") \
|
||||
X(bilagrid_shape_geometry, "bilagrid_shape_geometry", "bilagrid_shape_geometry", "model", "", "Shape of the bilateral grid for depth and normal (X, Y, W)") \
|
||||
X(use_adagrad_bilagrid_optim, "use_adagrad_bilagrid_optim", "use_adagrad_bilagrid_optim", "model", "", "Use AdaGrad (lr_decay=0, weight_decay=0, initial_accumulator_value=0, eps=1e-15) instead of Adam for all bilateral-grid parameters (RGB + depth + normal). When True, the bilagrid LR fields read from OptimizerConfig switch to ``bilagrid_adagrad_*_lr``. Bilagrid bit depths are coupled to `quantization_level`: level 0 = fp32; level 1 = 16-bit value + 8x2-bit optimizer state across all three bilagrid types.") \
|
||||
X(bilagrid_tv_loss_weight, "bilagrid_tv_loss_weight", "bilagrid_tv_loss_weight", "model", "", "Total variation loss weight for bilateral grid used for radiance") \
|
||||
X(color_shift_reg_weight, "color_shift_reg_weight", "color_shift_reg_weight", "model", "", "Color-shift regularizer for the combined bilagrid + PPISP color transform. Penalizes the dataset-wide mean of sign(post - pre) per channel, where `pre` is the splat-side rendered RGB (input to whichever of bilagrid / PPISP runs first) and `post` is the final image fed to the photometric loss: R = w * ||EMA[mean_p sign(post - pre)]||^2. The gradient is injected on the POST-transforms v_render_rgb buffer and flows through each transform's vjp into its parameters (and as a small leak into the splats). Active when at least one of bilagrid_rgb / PPISP is enabled; 0 disables. Typical values: 0.01--1.0.") \
|
||||
X(color_shift_reg_ema_period, "color_shift_reg_ema_period", "color_shift_reg_ema_period", "model", "", "EMA decay length for the color-shift regularizer, in BATCHES. beta = max(0, 1 - 1/period). Should be roughly the number of batches per epoch so the EMA estimates the dataset-wide mean. Ignored when color_shift_reg_weight = 0.") \
|
||||
X(bilagrid_tv_loss_weight_geometry, "bilagrid_tv_loss_weight_geometry", "bilagrid_tv_loss_weight_geometry", "model", "", "Total variation loss weight for bilateral grid used for geometry") \
|
||||
X(use_ppisp, "use_ppisp", "use_ppisp", "model", "", "If True, use the PPISP model (https://research.nvidia.com/labs/sil/projects/ppisp/) to handle per-pixel color distortions.") \
|
||||
X(ppisp_param_type, "ppisp_param_type", "ppisp_param_type", "model", "original|rqs|no_crf", "Parameterization for PPISP. \"original\" implements the original paper, \"rqs\" uses a parameterization that is more friendly to optimization and can produce better results in darker areas, \"no_crf\" omits the tone-curve (CRF) stage entirely and just clamps colors to [0,1] after exposure/vignetting/color correction.") \
|
||||
X(use_adagrad_ppisp_optim, "use_adagrad_ppisp_optim", "use_adagrad_ppisp_optim", "model", "", "Use unscheduled AdaGrad (lr_decay=0, weight_decay=0, initial_accumulator_value=0, eps=1e-15) instead of Adam for the PPISP parameter table. When True, the PPISP LR reads from ``OptimizerConfig.ppisp_adagrad_lr`` (constant) instead of the scheduled ``ppisp_lr``. No quantization path either way.") \
|
||||
X(apply_ppisp_before_bilagrid, "apply_ppisp_before_bilagrid", "apply_ppisp_before_bilagrid", "model", "", "When True, the PPISP forward runs BEFORE the RGB bilagrid (and PPISP backward runs AFTER bilagrid backward), i.e. render -> PPISP -> bilagrid -> loss. Otherwise: render -> bilagrid -> PPISP -> loss. Only meaningful when both ``use_bilateral_grid`` and ``use_ppisp`` are enabled.") \
|
||||
X(ppisp_reg_exposure_mean, "ppisp_reg_exposure_mean", "ppisp_reg_exposure_mean", "model", "", "Encourage exposure mean ~ 0 to resolve SH <-> exposure ambiguity in PPISP.") \
|
||||
X(ppisp_reg_vig_center, "ppisp_reg_vig_center", "ppisp_reg_vig_center", "model", "", "Encourage vignetting optical center near image center in PPISP.") \
|
||||
X(ppisp_reg_vig_non_pos, "ppisp_reg_vig_non_pos", "ppisp_reg_vig_non_pos", "model", "", "Penalize positive vignetting alpha coefficients in PPISP (should be <= 0).") \
|
||||
X(ppisp_reg_vig_channel_var, "ppisp_reg_vig_channel_var", "ppisp_reg_vig_channel_var", "model", "", "Encourage similar vignetting across RGB channels in PPISP.") \
|
||||
X(ppisp_reg_color_mean, "ppisp_reg_color_mean", "ppisp_reg_color_mean", "model", "", "Encourage color correction mean ~ 0 across frames in PPISP.") \
|
||||
X(ppisp_reg_crf_channel_var, "ppisp_reg_crf_channel_var", "ppisp_reg_crf_channel_var", "model", "", "Encourage similar CRF parameters across RGB channels in PPISP.") \
|
||||
X(image_color_is_linear, "image_color_is_linear", "image_color_is_linear", "model", "", "Whether to assume training images are in linear color space.") \
|
||||
X(image_color_gamut, "image_color_gamut", "image_color_gamut", "model", "ACES2065-1|ACEScg|Rec.2020|AdobeRGB|DCI-P3|none", "Color gamut of input images. If None, Rec.709 will be used. Note that tonemap is not applied.") \
|
||||
X(splat_color_is_linear, "splat_color_is_linear", "splat_color_is_linear", "model", "", "Whether to train splats in linear color space. If None, will use same as images.") \
|
||||
X(splat_color_gamut, "splat_color_gamut", "splat_color_gamut", "model", "Rec.709|ACES2065-1|ACEScg|Rec.2020|AdobeRGB|DCI-P3|none", "Color gamut of trained splats. If None, will use same as images. Note that tonemap is not applied.") \
|
||||
X(convert_initial_point_cloud_color, "convert_initial_point_cloud_color", "convert_initial_point_cloud_color", "model", "", "If True, this will assume color in initial point cloud is sRGB, and convert if images are in a linear or wide-gamut color space.") \
|
||||
X(scale_init, "scale_init", "scale_init", "model", "", "Initial scale. If not set, auto decide") \
|
||||
X(opacity_init, "opacity_init", "opacity_init", "model", "", "Initial opacity. If not set, auto decide") \
|
||||
X(suppress_initial_scales, "suppress_initial_scales", "suppress_initial_scales", "model", "", "Whether to suppress scales during initialization to discourage large floaters in vacant areas") \
|
||||
X(scale_regularization_weight, "scale_regularization_weight", "scale_regularization_weight", "model", "", "If enabled, a scale regularization introduced in PhysGauss (https://xpandora.github.io/PhysGaussian/) is used for reducing huge spikey gaussians.") \
|
||||
X(max_gauss_ratio, "max_gauss_ratio", "max_gauss_ratio", "model", "", "Threshold of ratio of gaussian max to min scale before applying regularization loss from the PhysGaussian paper") \
|
||||
X(depth_distortion_reg, "depth_distortion_reg", "depth_distortion_reg", "model", "", "Weight for depth distortion regularizer") \
|
||||
X(normal_distortion_reg, "normal_distortion_reg", "normal_distortion_reg", "model", "", "Weight for normal distortion regularizer") \
|
||||
X(rgb_distortion_reg, "rgb_distortion_reg", "rgb_distortion_reg", "model", "", "Weight for rgb distortion regularizer") \
|
||||
X(distortion_reg_warmup, "distortion_reg_warmup", "distortion_reg_warmup", "model", "", "warmup steps for depth regularizer, regularization weight ramps up") \
|
||||
X(normal_reg_weight, "normal_reg_weight", "normal_reg_weight", "model", "", "Weight for normal regularizer") \
|
||||
X(normal_reg_warmup, "normal_reg_warmup", "normal_reg_warmup", "model", "", "warmup steps for normal regularizer, regularization weight ramps up") \
|
||||
X(alpha_reg_weight, "alpha_reg_weight", "alpha_reg_weight", "model", "", "Weight for alpha regularizer (encourage alpha to go to either 0 or 1)") \
|
||||
X(alpha_reg_warmup, "alpha_reg_warmup", "alpha_reg_warmup", "model", "", "warmup steps for alpha regularizer, regularization weight ramps up") \
|
||||
X(reg_warmup_length, "reg_warmup_length", "reg_warmup_length", "model", "", "Warmup steps for depth, normal, and alpha regularizers. only apply regularizers after this many steps.") \
|
||||
X(apply_loss_for_mask, "apply_loss_for_mask", "apply_loss_for_mask", "model", "", "Set this to False to use masks to ignore distractors (e.g. people and cars, area outside fisheye circle, over exposure) Set this to True to remove background (e.g. sky, background outside centered object)") \
|
||||
X(enable_sky_masking, "enable_sky_masking", "enable_sky_masking", "model", "", "If enabled, sky from depth map will be used for masking. Alpha loss will always be applied for sky.") \
|
||||
X(alpha_loss_weight, "alpha_loss_weight", "alpha_loss_weight", "model", "", "Loss weight for alpha, applies when rendered alpha is above reference alpha") \
|
||||
X(alpha_loss_weight_under, "alpha_loss_weight_under", "alpha_loss_weight_under", "model", "", "Loss weight for alpha, applies when rendered alpha is below reference alpha") \
|
||||
X(opacity_reg, "opacity_reg", "opacity_reg", "model", "", "Encourage low opacity to aid densification, per MCMC.") \
|
||||
X(scale_reg, "scale_reg", "scale_reg", "model", "", "Encourage low scale, per MCMC.") \
|
||||
X(opacity_decay, "opacity_decay", "opacity_decay", "model", "", "Decay opacity to aid densification, per MRNF.") \
|
||||
X(scale_decay, "scale_decay", "scale_decay", "model", "", "Decay scale to aid densification, per MRNF.") \
|
||||
X(erank_reg, "erank_reg", "erank_reg", "model", "", "erank regularization weight, for 3DGS only - see *Effective Rank Analysis and Regularization for Enhanced 3D Gaussian Splatting, Hyung et al.*") \
|
||||
X(erank_reg_s3, "erank_reg_s3", "erank_reg_s3", "model", "", "erank regularization weight for smallest dimension, for 3DGS only") \
|
||||
X(quat_norm_reg, "quat_norm_reg", "quat_norm_reg", "model", "", "Weight to regularize quaternion norm to identity") \
|
||||
X(sh_reg, "sh_reg", "sh_reg", "model", "", "Regularize SH magnitude to find a balance between bilagrid/PPISP and improve generalizability.") \
|
||||
X(overexposure_reg, "overexposure_reg", "overexposure_reg", "model", "", "Image-space L2 penalty on rendered RGB outside [0, 1]: L = w * mean(max(-x, x-1, 0)^2) over all pixels and channels of the raw rendered image (pre-bilagrid / pre-PPISP / pre-color-space). When non-zero a dedicated CUDA kernel adds dL/dx directly into the rendered-RGB gradient; the scalar loss value is never materialized.") \
|
||||
X(supervision_warmup, "supervision_warmup", "supervision_warmup", "model", "", "Start using foundation model depth at this number of steps") \
|
||||
X(depth_supervision_weight, "depth_supervision_weight", "depth_supervision_weight", "model", "", "Weight for depth supervision by comparing rendered depth with depth predicted by a foundation model Warn that this can reduce quality if AI generated depth is heavily biased") \
|
||||
X(input_depth_is_ray_depth, "input_depth_is_ray_depth", "input_depth_is_ray_depth", "model", "", "Whether the input/supervision depth maps store ray depth (Euclidean distance along the camera ray) rather than linear (z) depth. The rasterizer renders ray depth, so when this is False (the common case, e.g. most foundation-model depths) the GT depth is converted from linear to ray depth in place on the GPU before the depth bilateral grid / loss. Set True for depth maps already in ray depth, e.g. >180deg fisheye captures where linear depth is ill-defined.") \
|
||||
X(normal_supervision_weight, "normal_supervision_weight", "normal_supervision_weight", "model", "", "Weight for normal supervision by comparing normal from rendered depth with normal from depth predicted by a foundation model") \
|
||||
X(mean_median_depth_weight, "mean_median_depth_weight", "mean_median_depth_weight", "model", "", "L1 between the mean (expected) depth and the median depth, where both are nonzero.") \
|
||||
X(median_depth_normal_reg_weight, "median_depth_normal_reg_weight", "median_depth_normal_reg_weight", "model", "", "normal_loss between the normal from the median depth and the normal from the mean depth.") \
|
||||
X(median_normal_supervision_weight, "median_normal_supervision_weight", "median_normal_supervision_weight", "model", "", "normal_loss between the normal from the median depth and the reference (foundation-model) normal.") \
|
||||
X(median_render_normal_reg_weight, "median_render_normal_reg_weight", "median_render_normal_reg_weight", "model", "", "normal_loss between the normal from the median depth and the rendered normal (placeholder until render_normal exists).") \
|
||||
X(median_warmup, "median_warmup", "median_warmup", "model", "", "Linear warmup length (steps) shared by all four median-depth loss weights: each ramps 0->full over this many steps.") \
|
||||
X(overfit_score_aggregation_mode, "overfit_score_aggregation_mode", "overfit_score_aggregation_mode", "model", "max|min|mean", "Mode to aggregate multiple overfitting objectives. Use max for more aggressive early stopping, min for more conservative early stopping, and mean for something in between.") \
|
||||
X(validation_loss_average_window, "validation_loss_average_window", "validation_loss_average_window", "model", "", "Window to calculate moving average validation loss for early stop") \
|
||||
X(early_stop_patience, "early_stop_patience", "early_stop_patience", "model", "", "Stop training if overfitting score remains positive for this number of iterations") \
|
||||
X(early_stop_warmup, "early_stop_warmup", "early_stop_warmup", "model", "", "Warmup steps for early stop, will not early stop before this number of steps Recommend setting this number no less than regularization warmups") \
|
||||
X(max_steps, "max_steps", "max_steps", "optimizer", "", "") \
|
||||
X(use_scale_agnostic_mean, "use_scale_agnostic_mean", "use_scale_agnostic_mean", "optimizer", "", "") \
|
||||
X(use_per_splat_bias_correction, "use_per_splat_bias_correction", "use_per_splat_bias_correction", "optimizer", "", "") \
|
||||
X(means_lr, "means_lr", "means_lr", "optimizer", "", "") \
|
||||
X(means_lr_final, "means_lr_final", "means_lr_final", "optimizer", "", "") \
|
||||
X(scales_lr, "scales_lr", "scales_lr", "optimizer", "", "") \
|
||||
X(scales_lr_final, "scales_lr_final", "scales_lr_final", "optimizer", "", "") \
|
||||
X(quats_lr, "quats_lr", "quats_lr", "optimizer", "", "") \
|
||||
X(opacities_lr, "opacities_lr", "opacities_lr", "optimizer", "", "") \
|
||||
X(features_dc_lr, "features_dc_lr", "features_dc_lr", "optimizer", "", "") \
|
||||
X(features_sh_lr, "features_sh_lr", "features_sh_lr", "optimizer", "", "") \
|
||||
X(background_dc_lr, "background_dc_lr", "background_dc_lr", "optimizer", "", "") \
|
||||
X(background_sh_lr, "background_sh_lr", "background_sh_lr", "optimizer", "", "") \
|
||||
X(bilagrid_lr, "bilagrid_lr", "bilagrid_lr", "optimizer", "", "") \
|
||||
X(bilagrid_lr_final, "bilagrid_lr_final", "bilagrid_lr_final", "optimizer", "", "") \
|
||||
X(bilagrid_lr_warmup, "bilagrid_lr_warmup", "bilagrid_lr_warmup", "optimizer", "", "") \
|
||||
X(bilagrid_depth_lr, "bilagrid_depth_lr", "bilagrid_depth_lr", "optimizer", "", "") \
|
||||
X(bilagrid_depth_lr_final, "bilagrid_depth_lr_final", "bilagrid_depth_lr_final", "optimizer", "", "") \
|
||||
X(bilagrid_depth_lr_warmup, "bilagrid_depth_lr_warmup", "bilagrid_depth_lr_warmup", "optimizer", "", "") \
|
||||
X(bilagrid_normal_lr, "bilagrid_normal_lr", "bilagrid_normal_lr", "optimizer", "", "") \
|
||||
X(bilagrid_normal_lr_final, "bilagrid_normal_lr_final", "bilagrid_normal_lr_final", "optimizer", "", "") \
|
||||
X(bilagrid_normal_lr_warmup, "bilagrid_normal_lr_warmup", "bilagrid_normal_lr_warmup", "optimizer", "", "") \
|
||||
X(bilagrid_adagrad_lr, "bilagrid_adagrad_lr", "bilagrid_adagrad_lr", "optimizer", "", "") \
|
||||
X(bilagrid_adagrad_depth_lr, "bilagrid_adagrad_depth_lr", "bilagrid_adagrad_depth_lr", "optimizer", "", "") \
|
||||
X(bilagrid_adagrad_normal_lr, "bilagrid_adagrad_normal_lr", "bilagrid_adagrad_normal_lr", "optimizer", "", "") \
|
||||
X(ppisp_lr, "ppisp_lr", "ppisp_lr", "optimizer", "", "") \
|
||||
X(ppisp_lr_final, "ppisp_lr_final", "ppisp_lr_final", "optimizer", "", "") \
|
||||
X(ppisp_lr_warmup, "ppisp_lr_warmup", "ppisp_lr_warmup", "optimizer", "", "") \
|
||||
X(ppisp_adagrad_lr, "ppisp_adagrad_lr", "ppisp_adagrad_lr", "optimizer", "", "") \
|
||||
X(camera_opt_lr, "camera_opt_lr", "camera_opt_lr", "optimizer", "", "") \
|
||||
X(camera_opt_lr_final, "camera_opt_lr_final", "camera_opt_lr_final", "optimizer", "", "") \
|
||||
X(camera_opt_lr_warmup, "camera_opt_lr_warmup", "camera_opt_lr_warmup", "optimizer", "", "") \
|
||||
/* end */
|
||||
|
||||
// Required fields (no Python default). Checked after flag parsing.
|
||||
#define SSPLAT_CONFIG_REQUIRED_FIELDS(X) \
|
||||
X(data) \
|
||||
/* end */
|
||||
|
||||
struct SsplatPresetInfo { const char* name; const char* help; };
|
||||
inline constexpr SsplatPresetInfo kSsplatPresets[] = {
|
||||
{"3dgs", "Generic method that works well for most datasets."},
|
||||
{"360-camera", "Preset for training on original distorted images captured by 360 cameras (e.g. Insta360, DJI Osmo). Recommended if your dataset contains fisheye images with a circle visible."},
|
||||
{"in-the-wild", "Preset for datasets consisting of internet images, with extreme lighting variation, with un-masked outliers, and/or shot with long focal lengths."},
|
||||
{"linear-color", "Preset for training splats in linear color spaces (e.g. ACEScg)."},
|
||||
{"synthetic", "Preset for training splats on synthetic datasets rendered with constant exposure."},
|
||||
{"meshing", "Preset for training splats for meshing. Use `spirulae-meshing` to convert trained splats to mesh."},
|
||||
{"academic-baseline", "Preset that replicates 3DGS MCMC as faithful as possible."},
|
||||
};
|
||||
|
||||
// Apply a preset's default overrides (tyro subcommand equivalent).
|
||||
// Returns false for an unknown preset name. "3dgs" is the base config.
|
||||
inline bool ssplat_apply_preset(SsplatConfig& c, const std::string& name) {
|
||||
if (name == "3dgs") {
|
||||
return true;
|
||||
}
|
||||
if (name == "360-camera") {
|
||||
c.warp_to_pinhole = true;
|
||||
c.mask_boundary_offset = -0.025f;
|
||||
c.primitive = "mip";
|
||||
c.long_axis_split_opacity_k = {0.5f, 0.6f, 15000.0f};
|
||||
c.input_depth_is_ray_depth = true;
|
||||
return true;
|
||||
}
|
||||
if (name == "in-the-wild") {
|
||||
c.center_method = "focus";
|
||||
c.outlier_threshold = 10.0f;
|
||||
c.load_depths = true;
|
||||
c.load_normals = true;
|
||||
c.mask_boundary_offset = -0.025f;
|
||||
c.densify_score_mode = "median";
|
||||
c.densify_loss_map_mode = "robust_edge_aware";
|
||||
c.densify_robust_edge_aware_quantile = 0.75f;
|
||||
c.ssim_lambda = 0.1f;
|
||||
c.rgb_distortion_reg = 0.1f;
|
||||
c.depth_distortion_reg = 0.01f;
|
||||
c.sh_degree_warmup_every = 0;
|
||||
c.long_axis_split_opacity_k = {0.5f, 0.6f, 30000.0f};
|
||||
c.noise_lr = 10.0f;
|
||||
c.noise_lr_final = 0.1f;
|
||||
c.erank_reg = 0.1f;
|
||||
c.means_lr = 5e-05f;
|
||||
c.means_lr_final = 1e-07f;
|
||||
return true;
|
||||
}
|
||||
if (name == "linear-color") {
|
||||
c.splat_color_gamut = "ACEScg";
|
||||
c.splat_color_is_linear = true;
|
||||
c.image_color_gamut = "Rec.2020";
|
||||
c.image_color_is_linear = false;
|
||||
c.background_mode = "noise";
|
||||
return true;
|
||||
}
|
||||
if (name == "synthetic") {
|
||||
c.min_init_fraction = 0.1f;
|
||||
c.use_bilateral_grid = false;
|
||||
c.use_ppisp = false;
|
||||
c.use_bilateral_grid_for_geometry = false;
|
||||
c.long_axis_split_opacity_k = {0.5f, 0.6f, 25000.0f};
|
||||
return true;
|
||||
}
|
||||
if (name == "meshing") {
|
||||
c.primitive = "3dgut";
|
||||
c.sh_degree = 0;
|
||||
c.sh_reg = 10.0f;
|
||||
c.overexposure_reg = 10.0f;
|
||||
c.background_mode = "noise";
|
||||
c.depth_distortion_reg = 0.01f;
|
||||
c.normal_distortion_reg = 0.01f;
|
||||
c.mean_median_depth_weight = 0.01f;
|
||||
c.median_depth_normal_reg_weight = 0.01f;
|
||||
c.normal_supervision_weight = 0.01f;
|
||||
c.median_normal_supervision_weight = 0.01f;
|
||||
c.median_render_normal_reg_weight = 0.01f;
|
||||
c.erank_reg = 0.01f;
|
||||
c.erank_reg_s3 = 0.01f;
|
||||
return true;
|
||||
}
|
||||
if (name == "academic-baseline") {
|
||||
c.eval_mode = "interval";
|
||||
c.eval_interval = 8;
|
||||
c.center_method = "gsplat";
|
||||
c.orientation_method = "gsplat";
|
||||
c.max_batch_per_epoch = 387420489;
|
||||
c.load_depths = false;
|
||||
c.load_normals = false;
|
||||
c.primitive = "3dgs";
|
||||
c.relative_scale = 1.0f;
|
||||
c.use_bilateral_grid = false;
|
||||
c.use_bilateral_grid_for_geometry = false;
|
||||
c.use_ppisp = false;
|
||||
c.use_revised_densification = false;
|
||||
c.densify_loss_map_mode = "none";
|
||||
c.use_long_axis_split = false;
|
||||
c.use_fused_proj_bwd_optim = false;
|
||||
c.quantization_level = 0;
|
||||
c.max_screen_size = std::numeric_limits<float>::infinity();
|
||||
c.max_world_size = std::numeric_limits<float>::infinity();
|
||||
c.suppress_initial_scales = false;
|
||||
c.scale_init = 0.1f;
|
||||
c.opacity_init = 0.5f;
|
||||
c.depth_distortion_reg = 0.0f;
|
||||
c.normal_reg_weight = 0.0f;
|
||||
c.alpha_reg_weight = 0.0f;
|
||||
c.alpha_loss_weight = 0.0f;
|
||||
c.alpha_loss_weight_under = 0.0f;
|
||||
c.erank_reg = 0.0f;
|
||||
c.erank_reg_s3 = 0.0f;
|
||||
c.quat_norm_reg = 0.0f;
|
||||
c.sh_reg = 0.0f;
|
||||
c.normal_supervision_weight = 0.0f;
|
||||
c.opacity_reg = 0.01f;
|
||||
c.scale_reg = 0.01f;
|
||||
c.max_steps = 30000;
|
||||
c.use_scale_agnostic_mean = false;
|
||||
c.use_per_splat_bias_correction = false;
|
||||
c.means_lr = 0.00016f;
|
||||
c.means_lr_final = 1.6e-06f;
|
||||
c.scales_lr = 0.005f;
|
||||
c.scales_lr_final = std::nullopt;
|
||||
c.quats_lr = 0.001f;
|
||||
c.opacities_lr = 0.05f;
|
||||
c.features_dc_lr = 0.0025f;
|
||||
c.features_sh_lr = 0.000125f;
|
||||
return true;
|
||||
}
|
||||
(void)c;
|
||||
return false;
|
||||
}
|
||||
@@ -1,7 +1,7 @@
|
||||
// ConfigUI.cpp -- see ConfigUI.h. Every widget below is expanded from the
|
||||
// generated SSPLAT_CONFIG_FIELDS X-macro; keep this file free of per-field
|
||||
// special cases (the point is that new Python config fields show up here
|
||||
// with zero GUI work).
|
||||
// SSPLAT_CONFIG_FIELDS X-macro; keep this file free of per-field special
|
||||
// cases (the point is that a new row in the field table shows up here with
|
||||
// zero GUI work).
|
||||
|
||||
#include "app/gui/ConfigUI.h"
|
||||
|
||||
@@ -248,10 +248,10 @@ bool draw_config_editor(SsplatConfig& cfg, const SsplatConfig& defaults,
|
||||
};
|
||||
|
||||
// Pass 1: per-group visible-field counts (groups are contiguous in the
|
||||
// generated table, so pass 2 can stream group headers).
|
||||
// field table, so pass 2 can stream group headers).
|
||||
int vis[5] = {0, 0, 0, 0, 0};
|
||||
#define SSPLAT_COUNT(member, cli_key, pyname, group, choices, help) \
|
||||
if (passes(cli_key, help, !(cfg.member == defaults.member))) \
|
||||
#define SSPLAT_COUNT(type, member, default_, group, choices, help) \
|
||||
if (passes(#member, help, !(cfg.member == defaults.member))) \
|
||||
vis[group_index(group)]++;
|
||||
SSPLAT_CONFIG_FIELDS(SSPLAT_COUNT)
|
||||
#undef SSPLAT_COUNT
|
||||
@@ -260,7 +260,7 @@ bool draw_config_editor(SsplatConfig& cfg, const SsplatConfig& defaults,
|
||||
bool any_changed = false;
|
||||
const char* cur_group = "";
|
||||
bool group_open = false;
|
||||
#define SSPLAT_DRAW(member, cli_key, pyname, group, choices, help) \
|
||||
#define SSPLAT_DRAW(type, member, default_, group, choices, help) \
|
||||
if (std::strcmp(cur_group, group) != 0) { \
|
||||
cur_group = group; \
|
||||
if (vis[group_index(group)] == 0) { \
|
||||
@@ -270,8 +270,8 @@ bool draw_config_editor(SsplatConfig& cfg, const SsplatConfig& defaults,
|
||||
group_open = ImGui::CollapsingHeader(group_label(group)); \
|
||||
} \
|
||||
} \
|
||||
if (group_open && passes(cli_key, help, !(cfg.member == defaults.member))) \
|
||||
any_changed |= field_row(cli_key, cfg.member, defaults.member, choices, help);
|
||||
if (group_open && passes(#member, help, !(cfg.member == defaults.member))) \
|
||||
any_changed |= field_row(#member, cfg.member, defaults.member, choices, help);
|
||||
SSPLAT_CONFIG_FIELDS(SSPLAT_DRAW)
|
||||
#undef SSPLAT_DRAW
|
||||
|
||||
|
||||
@@ -1,19 +1,19 @@
|
||||
#pragma once
|
||||
|
||||
// ConfigUI -- the "Advanced options" editor. All widgets are generated from
|
||||
// the SSPLAT_CONFIG_FIELDS X-macro (generated/cli_config.h), so every one of
|
||||
// ConfigUI -- the "Advanced options" editor. All widgets are expanded from
|
||||
// the SSPLAT_CONFIG_FIELDS X-macro (config/TrainConfig.h), so every one of
|
||||
// the training config fields is editable, with:
|
||||
// - grouping by Python sub-config (Run / Dataset / Data loading / Model /
|
||||
// - grouping by config group (Run / Dataset / Data loading / Model /
|
||||
// Optimizer), collapsed by default so novices are not overwhelmed
|
||||
// - a search box filtering by flag name and help text
|
||||
// - the full Python docstring as a hover tooltip (+ the preset default)
|
||||
// - the field's help text as a hover tooltip (+ the preset default)
|
||||
// - modified-from-preset highlighting and right-click "Reset to default"
|
||||
// - Literal[...] fields as dropdowns, Optional[...] as auto/override
|
||||
// - fields with `choices` as dropdowns, std::optional as auto/override
|
||||
//
|
||||
// New config fields added on the Python side appear here automatically after
|
||||
// re-running generate_cli_config.py -- no GUI change needed.
|
||||
// A row added to the field table appears here automatically -- no GUI change
|
||||
// needed.
|
||||
|
||||
#include "app/generated/cli_config.h"
|
||||
#include "config/TrainConfig.h"
|
||||
|
||||
namespace gui {
|
||||
|
||||
|
||||
@@ -58,10 +58,10 @@ const char* preset_help(const std::string& name) {
|
||||
}
|
||||
|
||||
// True when two configs parse to the same dataset: every dataparser-group
|
||||
// field (generated) plus the non-dataparser fields load_dataset() consumes.
|
||||
// field plus the non-dataparser fields load_dataset() consumes.
|
||||
bool parse_settings_equal(const SsplatConfig& a, const SsplatConfig& b) {
|
||||
bool eq = true;
|
||||
#define SSPLAT_CMP(member, cli_key, pyname, group, choices, help) \
|
||||
#define SSPLAT_CMP(type, member, default_, group, choices, help) \
|
||||
if (!std::strcmp(group, "dataparser")) eq = eq && (a.member == b.member);
|
||||
SSPLAT_CONFIG_FIELDS(SSPLAT_CMP)
|
||||
#undef SSPLAT_CMP
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
// viewport.
|
||||
|
||||
#include "backend/api/BackendRuntime.h"
|
||||
#include "app/generated/cli_config.h"
|
||||
#include "config/TrainConfig.h"
|
||||
#include "app/gui/ColmapRunner.h"
|
||||
#include "app/gui/ConfigUI.h"
|
||||
#include "app/gui/FileDialog.h"
|
||||
|
||||
@@ -10,7 +10,7 @@
|
||||
// What is bound, and why each piece:
|
||||
//
|
||||
// SsplatConfig the flattened training config. Generated from the
|
||||
// Python dataclasses by tools/codegen/generate_cli_config.py,
|
||||
// the hand-written table in src/config/TrainConfig.h,
|
||||
// so this binding is written against the generated
|
||||
// X-macro rather than a hand-listed field set -- add a
|
||||
// field to a Python dataclass, re-run the generator,
|
||||
@@ -46,15 +46,14 @@ void bind_trainer(py::module_& m) {
|
||||
// SsplatConfig -- generated field set, bound mechanically
|
||||
// -----------------------------------------------------------------
|
||||
py::class_<SsplatConfig> cfg_cls(m, "SsplatConfig", R"doc(
|
||||
The flattened native training config (src/app/generated/cli_config.h).
|
||||
The flattened native training config (src/config/TrainConfig.h).
|
||||
|
||||
Field names are the CLI keys with '_' separators; defaults are the Python
|
||||
dataclass defaults, baked in by the generator. Build one from a Python
|
||||
Field names are the CLI keys with '_' separators. Build one from a Python
|
||||
TrainerConfig with spirulae_splat.modules.native_trainer.to_native_config().
|
||||
)doc");
|
||||
cfg_cls.def(py::init<>());
|
||||
#define X(member, key, pyname, group, choices, help) \
|
||||
cfg_cls.def_readwrite(key, &SsplatConfig::member, help);
|
||||
#define X(type, member, default_, group, choices, help) \
|
||||
cfg_cls.def_readwrite(#member, &SsplatConfig::member, help);
|
||||
SSPLAT_CONFIG_FIELDS(X)
|
||||
#undef X
|
||||
|
||||
@@ -69,13 +68,14 @@ TrainerConfig with spirulae_splat.modules.native_trainer.to_native_config().
|
||||
"in-the-wild, linear-color, synthetic, meshing, academic-baseline).");
|
||||
|
||||
// (cli_key, pyname, group, choices, help) per field, straight off the
|
||||
// generated X-macro. Python uses (group, pyname) to find the value on its
|
||||
// own dataclass tree and cli_key to set it here -- so the name mapping
|
||||
// exists once, in the generator, instead of once more in Python.
|
||||
// field table. Python uses (group, pyname) to find the value on its own
|
||||
// dataclass tree and cli_key to set it here. pyname is the config.json
|
||||
// key, which is also what the Python dataclasses call the field.
|
||||
m.def("ssplat_config_fields", []() {
|
||||
py::list out;
|
||||
#define X(member, key, pyname, group, choices, help) \
|
||||
out.append(py::make_tuple(key, pyname, group, choices, help));
|
||||
#define X(type, member, default_, group, choices, help) \
|
||||
out.append(py::make_tuple(#member, ssplat_json_key(#member), group, \
|
||||
choices, help));
|
||||
SSPLAT_CONFIG_FIELDS(X)
|
||||
#undef X
|
||||
return out;
|
||||
|
||||
@@ -0,0 +1,601 @@
|
||||
#pragma once
|
||||
|
||||
// The training config: every flag `ssplat train` accepts, in one table.
|
||||
//
|
||||
// This file is the single source of truth. Adding a row to
|
||||
// SSPLAT_CONFIG_FIELDS makes the field appear in the CLI parser, in `--help`,
|
||||
// in the GUI's "All Options" editor, in the run's config.json and in
|
||||
// TrainerCore -- with no other edit. The struct is expanded from the same
|
||||
// table, so a field cannot exist in one and not the other.
|
||||
//
|
||||
// Row: X(type, member, default, group, choices, help)
|
||||
//
|
||||
// type one of the scalar types below, or SsplatVec3i / SsplatVec3f.
|
||||
// std::array<T, N> cannot appear here: its comma would split the
|
||||
// macro argument.
|
||||
// member the struct member AND, stringified, the CLI flag. The parser
|
||||
// treats '-' and '_' alike, so --sh-degree sets sh_degree.
|
||||
// default a constant expression. Vector defaults go through ssplat_v3i() /
|
||||
// ssplat_v3f() for the same comma reason.
|
||||
// group the section the flag is listed and nested under. Rows must stay
|
||||
// contiguous per group: --help and the GUI stream group headers as
|
||||
// they walk the table rather than sorting first.
|
||||
// choices '|'-separated list for string fields; "" is free-form; "none"
|
||||
// means the empty string is allowed and displays as `none`.
|
||||
// help one line, shown by --help and as the GUI tooltip. Keep the first
|
||||
// sentence self-contained -- --help truncates at the first ". ".
|
||||
|
||||
#include <array>
|
||||
#include <limits>
|
||||
#include <optional>
|
||||
#include <string>
|
||||
#include <string_view>
|
||||
|
||||
// Vector field types and their makers, so the table stays comma-free.
|
||||
using SsplatVec3i = std::array<int, 3>;
|
||||
using SsplatVec3f = std::array<float, 3>;
|
||||
constexpr SsplatVec3i ssplat_v3i(int a, int b, int c) { return {a, b, c}; }
|
||||
constexpr SsplatVec3f ssplat_v3f(float a, float b, float c) { return {a, b, c}; }
|
||||
|
||||
inline constexpr float kSsplatInf = std::numeric_limits<float>::infinity();
|
||||
|
||||
// The run's config.json keys the fields by flag name, with one exception:
|
||||
// datamanager's split_batch would collide with model.split_batch, so its flag
|
||||
// is dm_split_batch while the on-disk key stays split_batch. config.json is a
|
||||
// read-back format (`ssplat mesh` and --resume parse it), so this mapping is
|
||||
// compatibility, not policy -- a new field never needs an entry here.
|
||||
constexpr const char* ssplat_json_key(const char* flag) {
|
||||
return std::string_view(flag) == "dm_split_batch" ? "split_batch" : flag;
|
||||
}
|
||||
|
||||
|
||||
// ===========================================================================
|
||||
// The field table
|
||||
// ===========================================================================
|
||||
|
||||
#define SSPLAT_CONFIG_FIELDS(X) \
|
||||
\
|
||||
/* ==== trainer -- run control: output, checkpoints, viewer ==== */ \
|
||||
X(std::string, data, {}, "trainer", "", \
|
||||
"Path to dataset. Can be a Nerfstudio or a COLMAP dataset.") \
|
||||
X(std::string, resume, "", "trainer", "none", \
|
||||
"Resume training from a checkpoint. Pass a run output dir (latest step-*.ckpt is used) or a specific step-*.ckpt dir. The run's config.json supplies the architecture/model/data config (run-control flags like num_iterations / save cadence / viewer are still taken from the CLI); requires the checkpoint to have been written with save_full_checkpoint.") \
|
||||
X(std::string, output_dir_prefix, "outputs", "trainer", "", \
|
||||
"Prefix to output directory") \
|
||||
X(std::string, output_dir_name, "", "trainer", "none", \
|
||||
"Output directory name relative to output_dir_prefix. If not specified, will set a generic combining current timestamp and dataset name.") \
|
||||
X(int, steps_per_save, 2000, "trainer", "", \
|
||||
"Save checkpoint every this number of steps. If -1, save only at the end. If zero, never save (used in benchmark).") \
|
||||
X(bool, save_only_latest_checkpoint, true, "trainer", "", \
|
||||
"Whether to save only last checkpoint") \
|
||||
X(bool, save_full_checkpoint, false, "trainer", "", \
|
||||
"If True, each checkpoint's `state.tar` additionally includes the Resume slots -- world raw parameters (at max_num_splats) plus all optimizer state -- making the checkpoint sufficient to resume training, not just for inference. If False, only the Always slots (appearance/inference params) are saved alongside `splat.ply`.") \
|
||||
X(bool, save_eval_images, false, "trainer", "", \
|
||||
"Whether to save eval images at end of training") \
|
||||
X(int, num_iterations, 30000, "trainer", "", \
|
||||
"Number of training iterations") \
|
||||
X(int, viewer_port, 7007, "trainer", "", \
|
||||
"Port used by the web viewer") \
|
||||
X(bool, disable_viewer, false, "trainer", "", \
|
||||
"If True, ss_trainer skips starting the viewer thread. Used by ss_benchmark so each scene runs without competing for the viewer port.") \
|
||||
X(bool, keep_viewer_alive, true, "trainer", "", \
|
||||
"If True, ss_trainer keeps the process (and thus the viewer) running after training + eval finish, so the result can still be inspected in the browser. Press Ctrl-C to exit. Ignored when disable_viewer=True.") \
|
||||
\
|
||||
/* ==== dataparser -- where the dataset is and how poses are normalized ==== */ \
|
||||
X(std::string, data_format, "", "dataparser", "colmap|nerfstudio|metashape|none", \
|
||||
"Data format, leave None to auto detect") \
|
||||
X(std::string, colmap_recon_dir, "", "dataparser", "none", \
|
||||
"Path to COLMAP reconstruction relative to dataset directory (e.g. sparse/0). Will auto detect if not specified (picking the model with the most registered images when several exist).") \
|
||||
X(std::string, image_dir, "images", "dataparser", "", \
|
||||
"Path to images relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(std::string, mask_dir, "masks", "dataparser", "", \
|
||||
"Path to masks relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(std::string, depth_dir, "depths", "dataparser", "", \
|
||||
"Path to depth maps relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(std::string, normal_dir, "normals", "dataparser", "", \
|
||||
"Path to normal maps relative to dataset directory, used for COLMAP and Metashape datasets") \
|
||||
X(std::string, metashape_xml, "", "dataparser", "none", \
|
||||
"Path to the Metashape xml file. Will automatically detect if not specified.") \
|
||||
X(std::string, metashape_ply, "", "dataparser", "none", \
|
||||
"Path to the Metashape point export ply file. Will automatically detect if not specified.") \
|
||||
X(std::string, metashape_psx, "", "dataparser", "none", \
|
||||
"Path to Metashape PSX file, used to resolve file name ambiguity when there are multiple images with the same file name") \
|
||||
X(float, rescale_camera_to_fit, 0.0f, "dataparser", "", \
|
||||
"Whether to check if image resolution match camera resolution and scale camera intrinsics accordingly if not. Set this to a number to divide intrinsics by that number, e.g. Mip-NeRF 360 and Zip-NeRF with images_(2|4) Set this to True to detect resolution, e.g. tankt_db [CLI: 0 = off, -1 = auto-detect from image resolution, > 0 = divide intrinsics by this]") \
|
||||
X(std::string, downscale_rounding_mode, "floor", "dataparser", "floor|ceil|round", \
|
||||
"Rounding mode applied to camera width/height when dividing by `rescale_camera_to_fit`. Use `round` to match the convention used by most image downscalers (e.g. Mip-NeRF 360 images_(2|4|8)).") \
|
||||
X(float, scene_scale, 1.0f, "dataparser", "", \
|
||||
"How much to scale the region of interest by.") \
|
||||
X(std::string, orientation_method, "up", "dataparser", "pca|up|vertical|none|gsplat", \
|
||||
"The method to use for orientation.") \
|
||||
X(std::string, center_method, "poses", "dataparser", "poses|focus|none|gsplat", \
|
||||
"The method to use to center the poses.") \
|
||||
X(bool, auto_scale_poses, true, "dataparser", "", \
|
||||
"Whether to automatically scale the poses to fit in +/- 1 bounding box.") \
|
||||
X(float, outlier_threshold, kSsplatInf, "dataparser", "", \
|
||||
"Threshold to reject outlier camera poses.") \
|
||||
X(std::string, train_frame, "points", "dataparser", "normalized|camera|points", \
|
||||
"Coordinate frame in which splats are trained.") \
|
||||
X(std::string, eval_mode, "all", "dataparser", "fraction|filename|interval|all", \
|
||||
"The method to use for splitting the dataset into train and eval. Fraction splits based on a percentage for train and the remaining for eval. Filename splits based on filenames containing train/eval. Interval uses every nth frame for eval. All uses all the images for any split.") \
|
||||
X(float, train_split_fraction, 0.9f, "dataparser", "", \
|
||||
"The percentage of the dataset to use for training. Only used when eval_mode is train-split-fraction.") \
|
||||
X(int, eval_interval, 8, "dataparser", "", \
|
||||
"The interval between frames to use for eval. Only used when eval_mode is eval-interval.") \
|
||||
X(float, depth_unit_scale_factor, 0.001f, "dataparser", "", \
|
||||
"Scales the depth values to meters. Default value is 0.001 for a millimeter to meter conversion.") \
|
||||
X(float, validation_fraction, 0.0f, "dataparser", "", \
|
||||
"Use this fraction of training images for validation. Stop training when performance on validation images start to drop.") \
|
||||
\
|
||||
/* ==== datamanager -- image caching, masks, warping ==== */ \
|
||||
X(int, max_batch_per_epoch, 800, "datamanager", "", \
|
||||
"Maximum number of batches per epoch, used for configuring batch size") \
|
||||
X(bool, dm_split_batch, false, "datamanager", "", \
|
||||
"Whether to one large batch into many small batches to avoid OOM, at cost of slower training") \
|
||||
X(std::string, cache_images, "disk", "datamanager", "cpu-pageable|cpu|gpu|disk", \
|
||||
"Image cache location. If \"cpu\", caches on cpu. If \"gpu\", caches on device. If \"cpu-pageable\", cache on cpu pageable memory (saves RAM but may cause error if spill to swap memory). If \"disk\", cache on disk (limited support).") \
|
||||
X(bool, load_depths, true, "datamanager", "", \
|
||||
"Whether to load depth maps, if exist") \
|
||||
X(bool, load_normals, true, "datamanager", "", \
|
||||
"Whether to load normal maps, if exist") \
|
||||
X(float, mask_boundary_offset, 0.0f, "datamanager", "", \
|
||||
"Signed boundary offset applied to binarized masks at decode time, as a fraction of sqrt(W*H) of the decoded mask. Positive dilates (grows) foreground, negative erodes (shrinks). Runs on CPU during data loading via separable Felzenszwalb-Huttenlocher squared-Euclidean DT (exact, O(N) per row + col).") \
|
||||
X(bool, warp_to_pinhole, false, "datamanager", "", \
|
||||
"Whether to split a fisheye image into 5 undistorted pinhole images. Can sometimes give better quality and compatibility for dataset captured by fisheye/360 cameras.") \
|
||||
X(bool, warp_spherical_to_pinhole, true, "datamanager", "", \
|
||||
"Whether to split an equirectangular (spherical panorama) image into 6 pinhole cubemap faces for training. When True (default), equirectangular images are split into 6 undistorted pinhole sub-images (the historical behavior). When False, train directly on the equirectangular image using an equirectangular projection/unprojection plugged into the linear/UT projection pipeline (supports 3dgs/mip and 3dgut primitives). Direct equirectangular training does not support depth/normal supervision.") \
|
||||
X(bool, deblur_training_images, false, "datamanager", "", \
|
||||
"Whether to use a custom trained deep learning model to deblur images before training") \
|
||||
\
|
||||
/* ==== model -- the splat model, losses, densification, regularizers ==== */ \
|
||||
X(std::string, primitive, "3dgs", "model", "3dgs|mip|3dgut", \
|
||||
"Splat primitive to use") \
|
||||
X(int, sh_degree, 3, "model", "", \
|
||||
"Maximum degree of spherical harmonics to use.") \
|
||||
X(int, sh_degree_warmup_every, 1000, "model", "", \
|
||||
"Increase SH degree every this number of iterations") \
|
||||
X(std::string, background_mode, "black", "model", "black|noise|sh", \
|
||||
"Background mode, black per convention, noise to discourage transparency, sh for skybox.") \
|
||||
X(int, background_noise_warmup, 2000, "model", "", \
|
||||
"Number of steps to warmup background noise. This applies when background_mode is noise") \
|
||||
X(float, background_noise_pre_warmup, 0.25f, "model", "", \
|
||||
"Weight of background noise at start of training (0 to 1). Higher value reduce the chance of washing away splat opacities near the beginning of training.") \
|
||||
X(int, background_sh_degree, 4, "model", "", \
|
||||
"SH degree for background color, only used when background_mode is sh.") \
|
||||
X(std::optional<float>, relative_scale, std::nullopt, "model", "", \
|
||||
"Manually set scale when a scene is poorly scaled, i.e. increase this for large datasets. If not set, will use a scale agnostic optimizer. To prevent this, set it to 1.0.") \
|
||||
X(float, l1_weight, 1.0f, "model", "", \
|
||||
"Weight of L1 loss, default 1.0") \
|
||||
X(float, l2_weight, 0.0f, "model", "", \
|
||||
"Weight of L2 loss, default 0.0") \
|
||||
X(float, ssim_lambda, 0.2f, "model", "", \
|
||||
"Weight of ssim loss; 0.2 for academic baseline, higher for potentially more high-frequency details, lower for less blurry background in outdoor scenes") \
|
||||
X(float, l1_weight_y, 0.0f, "model", "", \
|
||||
"Weight of per-pixel BT.601 luma (Y) L1 loss.") \
|
||||
X(float, l2_weight_y, 0.0f, "model", "", \
|
||||
"Weight of per-pixel BT.601 luma (Y) L2 loss.") \
|
||||
X(float, l2_weight_u, 0.0f, "model", "", \
|
||||
"Weight of per-pixel BT.601 chroma U L2 loss.") \
|
||||
X(float, l2_weight_v, 0.0f, "model", "", \
|
||||
"Weight of per-pixel BT.601 chroma V L2 loss.") \
|
||||
X(int, num_loss_scales, 0, "model", "", \
|
||||
"Number of scales for image loss. For multi-scale loss, image is downscaled by 2 this number of times, and losses are averaged across scales. Improves convergence for high-resolution images.") \
|
||||
X(int, loss_scale_min_pixels, 1920, "model", "", \
|
||||
"If positive, overrides num_loss_scales per image based on resolution, in units of pixels. num_loss_scales is chosen so the smallest image dimension is halved down toward (but not below) this many pixels. e.g. with 2000: min dim 1999 -> num_loss_scales=0, 2000 -> 1, 4000 -> 2, 8000 -> 3, etc. Adapts per training step, so datasets with mixed image resolutions get the right count per image automatically.") \
|
||||
X(bool, use_camera_optimizer, false, "model", "", \
|
||||
"Whether to use camera optimizer Note: this only works well in patch batching mode") \
|
||||
X(bool, packed, true, "model", "", \
|
||||
"Pack projection outputs, reduce VRAM usage at large batch size but can be slightly slower") \
|
||||
X(bool, use_bvh, false, "model", "", \
|
||||
"Use BVH for splat-patch intersection test, may be faster when batching large number of small patches") \
|
||||
X(bool, use_fused_proj_bwd_optim, true, "model", "", \
|
||||
"Whether to use fused projection backward and optimizer. More memory efficient for large number of Gaussians, with slight performance hit.") \
|
||||
X(bool, split_batch, true, "model", "", \
|
||||
"Split the camera batch into one-camera sub-batches inside the C++ train step. Per-splat grads accumulate via atomicAdd across sub-batches; a single optim+densify pass at the end consumes the accumulator with grad_scale = 1/B. Drops peak VRAM for the immediate projection / rasterization buffers by roughly 1/B. Per-image grad magnitude vs regularization weight stays batch-size invariant. Not compatible with use_fused_proj_bwd_optim or with the warped train-step path.") \
|
||||
X(int, quantization_level, 1, "model", "", \
|
||||
"SH quantization level: a single int that selects one of two (param bits, optim bits) configurations. 0 = off : 32-bit param, fp32 optim 1 = light : 16-bit param, 8-bit packed optim (2 B / cell) Collapsing the prior independent param+optim bit controls into a single level minimizes the FPBO kernel instantiations.") \
|
||||
X(std::string, optimizer_offload, "", "model", "sh|all|none", \
|
||||
"Whether to offload optimizer momentum to CPU to save VRAM. This is only supported for Adam optimizer.") \
|
||||
X(int, resolution_schedule, 3000, "model", "", \
|
||||
"training starts at 1/d resolution, every n steps this is doubled") \
|
||||
X(int, num_downscales, 0, "model", "", \
|
||||
"At the beginning, resolution is 1/2^d, where d is this number") \
|
||||
X(bool, use_mcmc, true, "model", "", \
|
||||
"Must be True for 3DGS methods.") \
|
||||
X(bool, preallocate_splat_tensors, true, "model", "", \
|
||||
"Whether to pre-allocate Gaussian attribute tensors to cap_max to avoid OOM during densification") \
|
||||
X(int, cap_max, 1000000, "model", "", \
|
||||
"maximum number of splats, dataset-specific tuning required") \
|
||||
X(float, min_init_fraction, 0.0f, "model", "", \
|
||||
"minimum fraction of splats out of cap_max at initialization") \
|
||||
X(int, refine_every, 100, "model", "", \
|
||||
"Densify every this number of steps") \
|
||||
X(int, refine_start_iter, 500, "model", "", \
|
||||
"Start densification at this number of steps") \
|
||||
X(int, refine_stop_num_iter, 5000, "model", "", \
|
||||
"Stop densification at this number of steps before maximum number of training iterations") \
|
||||
X(int, refine_stop_iter, 25000, "model", "", \
|
||||
"Densification runs until max(this, num_iterations - refine_stop_num_iter). Without this floor, runs shorter than refine_stop_num_iter would never densify at all (num_iterations - refine_stop_num_iter goes negative), which confuses users.") \
|
||||
X(float, noise_lr, 80.0f, "model", "", \
|
||||
"MCMC-like noise injection magnitude at start of training") \
|
||||
X(float, noise_lr_final, 0.8f, "model", "", \
|
||||
"MCMC-like noise injection magnitude at end of training") \
|
||||
X(float, min_opacity, 0.005f, "model", "", \
|
||||
"Minimum Gaussian opacity before relocation") \
|
||||
X(float, growth_factor, 1.05f, "model", "", \
|
||||
"Multiply number of splats by this number at each densification step") \
|
||||
X(bool, use_revised_densification, true, "model", "", \
|
||||
"Whether to use revised densification instead of original MCMC.") \
|
||||
X(std::string, densify_score_mode, "mean", "model", "mean|max|median|geom", \
|
||||
"How to accumulate per-splat scores across iterations for densification. \"mean\": running mean of |w|. \"max\": running max of |w|. \"median\": running median of |w| (approximation). \"geom\": running geometric mean of |w|.") \
|
||||
X(float, densify_score_blend_world_grad, 0.0f, "model", "", \
|
||||
"Blend weight `w` in [0, 1] between the image-space loss score and the world-space gradient score for densification. The per-step score is (image-space accum_weight)^(1-w) * (||dL/dmean_world|| * max post-exp world scale)^w. The world-grad term favors world-space-large splats (e.g. distant background in unbounded outdoor scenes) that the image-space score under-weights; the geometric blend is invariant to each score's global scale so no cross-normalization is needed. 0 (default) = image-space score only, identical cost and behavior to before. 1 = world-grad score only; the per-pixel densification loss map (densify_loss_map_mode) is skipped entirely. In between, both scores are computed (one extra float per splat of VRAM).") \
|
||||
X(std::string, densify_loss_map_mode, "ssim_structure", "model", "none|loss_full|ssim_full|ssim_cs|ssim_structure|edge_aware|robust_edge_aware", \
|
||||
"What gets accumulated into the per-pixel densification loss map. The loss map is read by raster bwd to weight the per-splat accum_weight. Only active when use_revised_densification. Modes: \"none\": no loss map (uniform alpha*T accumulation). \"loss_full\": per-pixel L1/L2 + auxiliary supervisory terms + full SSIM (luminance*contrast*structure). \"ssim_full\": full SSIM only. \"ssim_cs\": contrast*structure SSIM (no luminance). \"ssim_structure\": structure-only SSIM, biases toward pattern/edge mismatches and ignores brightness/contrast errors. \"edge_aware\": canny edge magnitude of GT rgb (Plenoxels-style, https://arxiv.org/abs/2603.08661). Biases densification toward GT edges directly, regardless of how well the splats already reconstruct them. \"robust_edge_aware\": RobustNeRF-style Tukey biweight on the BT.601 luma of |render - GT|, capped at the per-image `densify_robust_edge_aware_quantile`, then canny. Near-zero where the render already matches GT, zeroed past the quantile cutoff so distractor pixels (people/cars/operator) don't pull splats toward them, and luminance-shift tolerant since a global DC residual has no spatial gradient. For num_loss_scales > 0 the map is computed per scale and the per-scale results are averaged (matches the multi-scale loss accumulation). Affects loss_map only; training gradients and scalar losses unchanged.") \
|
||||
X(float, densify_robust_edge_aware_quantile, 0.9f, "model", "", \
|
||||
"Per-image quantile of the luma residual used as the Tukey biweight cutoff in `robust_edge_aware` mode. Pixels whose residual exceeds this quantile get zero densification weight. Lower values are more aggressive outlier rejection (good for distractor-heavy datasets); higher are more permissive (good for clean datasets where real edges may produce large residuals). Ignored unless `densify_loss_map_mode == \"robust_edge_aware\"`.") \
|
||||
X(bool, use_long_axis_split, true, "model", "", \
|
||||
"whether to use long-axis split described in https://arxiv.org/abs/2508.12313 for relocation and sample add. When combined with use_revised_densification, this can give less blurry background details for unbounded outdoor scenes.") \
|
||||
X(SsplatVec3f, long_axis_split_opacity_k, ssplat_v3f(0.6f, 0.6f, 4500.0f), "model", "", \
|
||||
"opacity split factor `k` for long-axis split, as (initial, final, warmup_steps). Each split child keeps opacity `logit^-1(k / (1 + exp(-logit_opacity) - k))`; `k` is linearly scheduled from `initial` to `final` over the first `warmup_steps` training steps, then held at `final`. Larger `k` preserves more opacity per child (denser, sharper); smaller `k` fades children faster.") \
|
||||
X(float, relocate_screen_size, kSsplatInf, "model", "", \
|
||||
"if a gaussian is more than this fraction of screen space, relocate it Useful for fisheye with 3DGUT, may drop PSNR for conventional cameras For likely better quality, use max_screen_size instead") \
|
||||
X(float, max_screen_size, 0.3f, "model", "", \
|
||||
"if a gaussian is more than this fraction of screen space, clip scale and increase opacity Intended to be an MCMC-friendly alternative of relocate_screen_size") \
|
||||
X(float, max_screen_size_clip_hardness, 1.5f, "model", "", \
|
||||
"clip hardness for Gaussians with large screen space size, between 1 and infinity, larger is harder") \
|
||||
X(float, max_world_size, kSsplatInf, "model", "", \
|
||||
"if a gaussian is more than this of world space, clip scale Useful if you see huge floaters at a distance in large indoor space") \
|
||||
X(int, reset_alpha_every, 30, "model", "", \
|
||||
"Every this many refinement steps, reset the alpha. Only applies for opaque triangle splatting.") \
|
||||
X(bool, use_bilateral_grid, true, "model", "", \
|
||||
"If True, use bilateral grid to handle the ISP changes in the image space. This technique was introduced in the paper 'Bilateral Guided Radiance Field Processing' (https://bilarfpro.github.io/).") \
|
||||
X(SsplatVec3i, bilagrid_shape, ssplat_v3i(16, 16, 8), "model", "", \
|
||||
"Shape of the bilateral grid, typically `16 16 8`, or `8 8 4` for scenes with low-texture surfaces.") \
|
||||
X(std::string, bilagrid_type, "ppisp", "model", "affine|ppisp|loglinear", \
|
||||
"What the bilateral grid predicts. affine: 4x3 matrix per original bilateral grid. ppisp: PPISP exposure and color parameters, generally gives less color shift but can be less numerically stable. loglinear: 3x3 linear transformation matrix with log-encoded diagonals, balances color shift and numerical stability.") \
|
||||
X(bool, use_bilateral_grid_for_geometry, true, "model", "", \
|
||||
"If True, use bilateral grid for depth and normal (e.g. AI generated biased ones)") \
|
||||
X(SsplatVec3i, bilagrid_shape_geometry, ssplat_v3i(8, 8, 4), "model", "", \
|
||||
"Shape of the bilateral grid for depth and normal (X, Y, W)") \
|
||||
X(bool, use_adagrad_bilagrid_optim, true, "model", "", \
|
||||
"Use AdaGrad (lr_decay=0, weight_decay=0, initial_accumulator_value=0, eps=1e-15) instead of Adam for all bilateral-grid parameters (RGB + depth + normal). When True, the bilagrid LR fields read from OptimizerConfig switch to ``bilagrid_adagrad_*_lr``. Bilagrid bit depths are coupled to `quantization_level`: level 0 = fp32; level 1 = 16-bit value + 8x2-bit optimizer state across all three bilagrid types.") \
|
||||
X(float, bilagrid_tv_loss_weight, 10.0f, "model", "", \
|
||||
"Total variation loss weight for bilateral grid used for radiance") \
|
||||
X(float, color_shift_reg_weight, 0.0f, "model", "", \
|
||||
"Color-shift regularizer for the combined bilagrid + PPISP color transform. Penalizes the dataset-wide mean of sign(post - pre) per channel, where `pre` is the splat-side rendered RGB (input to whichever of bilagrid / PPISP runs first) and `post` is the final image fed to the photometric loss: R = w * ||EMA[mean_p sign(post - pre)]||^2. The gradient is injected on the POST-transforms v_render_rgb buffer and flows through each transform's vjp into its parameters (and as a small leak into the splats). Active when at least one of bilagrid_rgb / PPISP is enabled; 0 disables. Typical values: 0.01--1.0.") \
|
||||
X(int, color_shift_reg_ema_period, 750, "model", "", \
|
||||
"EMA decay length for the color-shift regularizer, in BATCHES. beta = max(0, 1 - 1/period). Should be roughly the number of batches per epoch so the EMA estimates the dataset-wide mean. Ignored when color_shift_reg_weight = 0.") \
|
||||
X(float, bilagrid_tv_loss_weight_geometry, 10.0f, "model", "", \
|
||||
"Total variation loss weight for bilateral grid used for geometry") \
|
||||
X(bool, use_ppisp, true, "model", "", \
|
||||
"If True, use the PPISP model (https://research.nvidia.com/labs/sil/projects/ppisp/) to handle per-pixel color distortions.") \
|
||||
X(std::string, ppisp_param_type, "no_crf", "model", "original|rqs|no_crf", \
|
||||
"Parameterization for PPISP. \"original\" implements the original paper, \"rqs\" uses a parameterization that is more friendly to optimization and can produce better results in darker areas, \"no_crf\" omits the tone-curve (CRF) stage entirely and just clamps colors to [0,1] after exposure/vignetting/color correction.") \
|
||||
X(bool, use_adagrad_ppisp_optim, true, "model", "", \
|
||||
"Use unscheduled AdaGrad (lr_decay=0, weight_decay=0, initial_accumulator_value=0, eps=1e-15) instead of Adam for the PPISP parameter table. When True, the PPISP LR reads from ``OptimizerConfig.ppisp_adagrad_lr`` (constant) instead of the scheduled ``ppisp_lr``. No quantization path either way.") \
|
||||
X(bool, apply_ppisp_before_bilagrid, true, "model", "", \
|
||||
"When True, the PPISP forward runs BEFORE the RGB bilagrid (and PPISP backward runs AFTER bilagrid backward), i.e. render -> PPISP -> bilagrid -> loss. Otherwise: render -> bilagrid -> PPISP -> loss. Only meaningful when both ``use_bilateral_grid`` and ``use_ppisp`` are enabled.") \
|
||||
X(float, ppisp_reg_exposure_mean, 1.0f, "model", "", \
|
||||
"Encourage exposure mean ~ 0 to resolve SH <-> exposure ambiguity in PPISP.") \
|
||||
X(float, ppisp_reg_vig_center, 0.02f, "model", "", \
|
||||
"Encourage vignetting optical center near image center in PPISP.") \
|
||||
X(float, ppisp_reg_vig_non_pos, 0.01f, "model", "", \
|
||||
"Penalize positive vignetting alpha coefficients in PPISP (should be <= 0).") \
|
||||
X(float, ppisp_reg_vig_channel_var, 0.1f, "model", "", \
|
||||
"Encourage similar vignetting across RGB channels in PPISP.") \
|
||||
X(float, ppisp_reg_color_mean, 1.0f, "model", "", \
|
||||
"Encourage color correction mean ~ 0 across frames in PPISP.") \
|
||||
X(float, ppisp_reg_crf_channel_var, 0.1f, "model", "", \
|
||||
"Encourage similar CRF parameters across RGB channels in PPISP.") \
|
||||
X(bool, image_color_is_linear, false, "model", "", \
|
||||
"Whether to assume training images are in linear color space.") \
|
||||
X(std::string, image_color_gamut, "", "model", "ACES2065-1|ACEScg|Rec.2020|AdobeRGB|DCI-P3|none", \
|
||||
"Color gamut of input images. If None, Rec.709 will be used. Note that tonemap is not applied.") \
|
||||
X(std::optional<bool>, splat_color_is_linear, std::nullopt, "model", "", \
|
||||
"Whether to train splats in linear color space. If None, will use same as images.") \
|
||||
X(std::string, splat_color_gamut, "", "model", "Rec.709|ACES2065-1|ACEScg|Rec.2020|AdobeRGB|DCI-P3|none", \
|
||||
"Color gamut of trained splats. If None, will use same as images. Note that tonemap is not applied.") \
|
||||
X(std::optional<bool>, convert_initial_point_cloud_color, std::nullopt, "model", "", \
|
||||
"If True, this will assume color in initial point cloud is sRGB, and convert if images are in a linear or wide-gamut color space.") \
|
||||
X(std::optional<float>, scale_init, std::nullopt, "model", "", \
|
||||
"Initial scale. If not set, auto decide") \
|
||||
X(std::optional<float>, opacity_init, std::nullopt, "model", "", \
|
||||
"Initial opacity. If not set, auto decide") \
|
||||
X(bool, suppress_initial_scales, false, "model", "", \
|
||||
"Whether to suppress scales during initialization to discourage large floaters in vacant areas") \
|
||||
X(float, scale_regularization_weight, 0.0f, "model", "", \
|
||||
"If enabled, a scale regularization introduced in PhysGauss (https://xpandora.github.io/PhysGaussian/) is used for reducing huge spikey gaussians.") \
|
||||
X(float, max_gauss_ratio, 10.0f, "model", "", \
|
||||
"Threshold of ratio of gaussian max to min scale before applying regularization loss from the PhysGaussian paper") \
|
||||
X(float, depth_distortion_reg, 0.0f, "model", "", \
|
||||
"Weight for depth distortion regularizer") \
|
||||
X(float, normal_distortion_reg, 0.0f, "model", "", \
|
||||
"Weight for normal distortion regularizer") \
|
||||
X(float, rgb_distortion_reg, 0.0f, "model", "", \
|
||||
"Weight for rgb distortion regularizer") \
|
||||
X(int, distortion_reg_warmup, 6000, "model", "", \
|
||||
"warmup steps for depth regularizer, regularization weight ramps up") \
|
||||
X(float, normal_reg_weight, 0.04f, "model", "", \
|
||||
"Weight for normal regularizer") \
|
||||
X(int, normal_reg_warmup, 6000, "model", "", \
|
||||
"warmup steps for normal regularizer, regularization weight ramps up") \
|
||||
X(float, alpha_reg_weight, 0.0f, "model", "", \
|
||||
"Weight for alpha regularizer (encourage alpha to go to either 0 or 1)") \
|
||||
X(int, alpha_reg_warmup, 12000, "model", "", \
|
||||
"warmup steps for alpha regularizer, regularization weight ramps up") \
|
||||
X(int, reg_warmup_length, 0, "model", "", \
|
||||
"Warmup steps for depth, normal, and alpha regularizers. only apply regularizers after this many steps.") \
|
||||
X(bool, apply_loss_for_mask, false, "model", "", \
|
||||
"Set this to False to use masks to ignore distractors (e.g. people and cars, area outside fisheye circle, over exposure) Set this to True to remove background (e.g. sky, background outside centered object)") \
|
||||
X(bool, enable_sky_masking, true, "model", "", \
|
||||
"If enabled, sky from depth map will be used for masking. Alpha loss will always be applied for sky.") \
|
||||
X(float, alpha_loss_weight, 0.01f, "model", "", \
|
||||
"Loss weight for alpha, applies when rendered alpha is above reference alpha") \
|
||||
X(float, alpha_loss_weight_under, 0.0f, "model", "", \
|
||||
"Loss weight for alpha, applies when rendered alpha is below reference alpha") \
|
||||
X(float, opacity_reg, 0.01f, "model", "", \
|
||||
"Encourage low opacity to aid densification, per MCMC.") \
|
||||
X(float, scale_reg, 0.01f, "model", "", \
|
||||
"Encourage low scale, per MCMC.") \
|
||||
X(float, opacity_decay, 0.0f, "model", "", \
|
||||
"Decay opacity to aid densification, per MRNF.") \
|
||||
X(float, scale_decay, 0.0f, "model", "", \
|
||||
"Decay scale to aid densification, per MRNF.") \
|
||||
X(float, erank_reg, 0.0f, "model", "", \
|
||||
"erank regularization weight, for 3DGS only - see *Effective Rank Analysis and Regularization for Enhanced 3D Gaussian Splatting, Hyung et al.*") \
|
||||
X(float, erank_reg_s3, 0.0f, "model", "", \
|
||||
"erank regularization weight for smallest dimension, for 3DGS only") \
|
||||
X(float, quat_norm_reg, 0.01f, "model", "", \
|
||||
"Weight to regularize quaternion norm to identity") \
|
||||
X(float, sh_reg, 0.001f, "model", "", \
|
||||
"Regularize SH magnitude to find a balance between bilagrid/PPISP and improve generalizability.") \
|
||||
X(float, overexposure_reg, 0.0f, "model", "", \
|
||||
"Image-space L2 penalty on rendered RGB outside [0, 1]: L = w * mean(max(-x, x-1, 0)^2) over all pixels and channels of the raw rendered image (pre-bilagrid / pre-PPISP / pre-color-space). When non-zero a dedicated CUDA kernel adds dL/dx directly into the rendered-RGB gradient; the scalar loss value is never materialized.") \
|
||||
X(int, supervision_warmup, 0, "model", "", \
|
||||
"Start using foundation model depth at this number of steps") \
|
||||
X(float, depth_supervision_weight, 0.0f, "model", "", \
|
||||
"Weight for depth supervision by comparing rendered depth with depth predicted by a foundation model Warn that this can reduce quality if AI generated depth is heavily biased") \
|
||||
X(bool, input_depth_is_ray_depth, false, "model", "", \
|
||||
"Whether the input/supervision depth maps store ray depth (Euclidean distance along the camera ray) rather than linear (z) depth. The rasterizer renders ray depth, so when this is False (the common case, e.g. most foundation-model depths) the GT depth is converted from linear to ray depth in place on the GPU before the depth bilateral grid / loss. Set True for depth maps already in ray depth, e.g. >180deg fisheye captures where linear depth is ill-defined.") \
|
||||
X(float, normal_supervision_weight, 0.01f, "model", "", \
|
||||
"Weight for normal supervision by comparing normal from rendered depth with normal from depth predicted by a foundation model") \
|
||||
X(float, mean_median_depth_weight, 0.0f, "model", "", \
|
||||
"L1 between the mean (expected) depth and the median depth, where both are nonzero.") \
|
||||
X(float, median_depth_normal_reg_weight, 0.0f, "model", "", \
|
||||
"normal_loss between the normal from the median depth and the normal from the mean depth.") \
|
||||
X(float, median_normal_supervision_weight, 0.0f, "model", "", \
|
||||
"normal_loss between the normal from the median depth and the reference (foundation-model) normal.") \
|
||||
X(float, median_render_normal_reg_weight, 0.0f, "model", "", \
|
||||
"normal_loss between the normal from the median depth and the rendered normal (placeholder until render_normal exists).") \
|
||||
X(int, median_warmup, 6000, "model", "", \
|
||||
"Linear warmup length (steps) shared by all four median-depth loss weights: each ramps 0->full over this many steps.") \
|
||||
X(std::string, overfit_score_aggregation_mode, "min", "model", "max|min|mean", \
|
||||
"Mode to aggregate multiple overfitting objectives. Use max for more aggressive early stopping, min for more conservative early stopping, and mean for something in between.") \
|
||||
X(int, validation_loss_average_window, 500, "model", "", \
|
||||
"Window to calculate moving average validation loss for early stop") \
|
||||
X(int, early_stop_patience, 1000, "model", "", \
|
||||
"Stop training if overfitting score remains positive for this number of iterations") \
|
||||
X(int, early_stop_warmup, 12000, "model", "", \
|
||||
"Warmup steps for early stop, will not early stop before this number of steps Recommend setting this number no less than regularization warmups") \
|
||||
\
|
||||
/* ==== optimizer -- learning rates and schedules ==== */ \
|
||||
X(std::optional<int>, max_steps, std::nullopt, "optimizer", "", \
|
||||
"") \
|
||||
X(bool, use_scale_agnostic_mean, true, "optimizer", "", \
|
||||
"") \
|
||||
X(bool, use_per_splat_bias_correction, true, "optimizer", "", \
|
||||
"") \
|
||||
X(float, means_lr, 0.000128f, "optimizer", "", \
|
||||
"") \
|
||||
X(std::optional<float>, means_lr_final, 1.6e-06f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, scales_lr, 0.02f, "optimizer", "", \
|
||||
"") \
|
||||
X(std::optional<float>, scales_lr_final, 0.005f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, quats_lr, 0.0015f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, opacities_lr, 0.025f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, features_dc_lr, 0.005f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, features_sh_lr, 0.00025f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, background_dc_lr, 0.0025f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, background_sh_lr, 0.0005f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, bilagrid_lr, 0.002f, "optimizer", "", \
|
||||
"") \
|
||||
X(std::optional<float>, bilagrid_lr_final, 0.0001f, "optimizer", "", \
|
||||
"") \
|
||||
X(int, bilagrid_lr_warmup, 1000, "optimizer", "", \
|
||||
"") \
|
||||
X(float, bilagrid_depth_lr, 0.002f, "optimizer", "", \
|
||||
"") \
|
||||
X(std::optional<float>, bilagrid_depth_lr_final, 0.0001f, "optimizer", "", \
|
||||
"") \
|
||||
X(int, bilagrid_depth_lr_warmup, 2000, "optimizer", "", \
|
||||
"") \
|
||||
X(float, bilagrid_normal_lr, 0.0005f, "optimizer", "", \
|
||||
"") \
|
||||
X(std::optional<float>, bilagrid_normal_lr_final, 4e-05f, "optimizer", "", \
|
||||
"") \
|
||||
X(int, bilagrid_normal_lr_warmup, 2000, "optimizer", "", \
|
||||
"") \
|
||||
X(float, bilagrid_adagrad_lr, 0.04f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, bilagrid_adagrad_depth_lr, 0.04f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, bilagrid_adagrad_normal_lr, 0.01f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, ppisp_lr, 0.002f, "optimizer", "", \
|
||||
"") \
|
||||
X(std::optional<float>, ppisp_lr_final, 2e-05f, "optimizer", "", \
|
||||
"") \
|
||||
X(int, ppisp_lr_warmup, 500, "optimizer", "", \
|
||||
"") \
|
||||
X(float, ppisp_adagrad_lr, 0.1f, "optimizer", "", \
|
||||
"") \
|
||||
X(float, camera_opt_lr, 0.0001f, "optimizer", "", \
|
||||
"") \
|
||||
X(std::optional<float>, camera_opt_lr_final, 5e-07f, "optimizer", "", \
|
||||
"") \
|
||||
X(int, camera_opt_lr_warmup, 1000, "optimizer", "", \
|
||||
"") \
|
||||
/* end */
|
||||
|
||||
|
||||
// ===========================================================================
|
||||
// The config struct, expanded from the table above
|
||||
// ===========================================================================
|
||||
|
||||
struct SsplatConfig {
|
||||
#define SSPLAT_DECLARE_FIELD(type, member, default_, group, choices, help) \
|
||||
type member = default_;
|
||||
SSPLAT_CONFIG_FIELDS(SSPLAT_DECLARE_FIELD)
|
||||
#undef SSPLAT_DECLARE_FIELD
|
||||
};
|
||||
|
||||
// Fields whose default ({}) is not a usable value. Checked after flag
|
||||
// parsing, so --help still works without them.
|
||||
#define SSPLAT_CONFIG_REQUIRED_FIELDS(X) \
|
||||
X(data) \
|
||||
/* end */
|
||||
|
||||
|
||||
// ===========================================================================
|
||||
// Presets -- named bundles of default overrides, selected as
|
||||
// `ssplat train <preset>`. "3dgs" is the base config and applies nothing.
|
||||
// ===========================================================================
|
||||
|
||||
struct SsplatPresetInfo { const char* name; const char* help; };
|
||||
inline constexpr SsplatPresetInfo kSsplatPresets[] = {
|
||||
{"3dgs", "Generic method that works well for most datasets."},
|
||||
{"360-camera", "Preset for training on original distorted images captured by 360 cameras (e.g. Insta360, DJI Osmo). Recommended if your dataset contains fisheye images with a circle visible."},
|
||||
{"in-the-wild", "Preset for datasets consisting of internet images, with extreme lighting variation, with un-masked outliers, and/or shot with long focal lengths."},
|
||||
{"linear-color", "Preset for training splats in linear color spaces (e.g. ACEScg)."},
|
||||
{"synthetic", "Preset for training splats on synthetic datasets rendered with constant exposure."},
|
||||
{"meshing", "Preset for training splats for meshing. Use `spirulae-meshing` to convert trained splats to mesh."},
|
||||
{"academic-baseline", "Preset that replicates 3DGS MCMC as faithful as possible."},
|
||||
};
|
||||
|
||||
// Returns false for an unknown preset name.
|
||||
inline bool ssplat_apply_preset(SsplatConfig& c, const std::string& name) {
|
||||
if (name == "3dgs") {
|
||||
return true;
|
||||
}
|
||||
if (name == "360-camera") {
|
||||
c.warp_to_pinhole = true;
|
||||
c.mask_boundary_offset = -0.025f;
|
||||
c.primitive = "mip";
|
||||
c.long_axis_split_opacity_k = {0.5f, 0.6f, 15000.0f};
|
||||
c.input_depth_is_ray_depth = true;
|
||||
return true;
|
||||
}
|
||||
if (name == "in-the-wild") {
|
||||
c.center_method = "focus";
|
||||
c.outlier_threshold = 10.0f;
|
||||
c.load_depths = true;
|
||||
c.load_normals = true;
|
||||
c.mask_boundary_offset = -0.025f;
|
||||
c.densify_score_mode = "median";
|
||||
c.densify_loss_map_mode = "robust_edge_aware";
|
||||
c.densify_robust_edge_aware_quantile = 0.75f;
|
||||
c.ssim_lambda = 0.1f;
|
||||
c.rgb_distortion_reg = 0.1f;
|
||||
c.depth_distortion_reg = 0.01f;
|
||||
c.sh_degree_warmup_every = 0;
|
||||
c.long_axis_split_opacity_k = {0.5f, 0.6f, 30000.0f};
|
||||
c.noise_lr = 10.0f;
|
||||
c.noise_lr_final = 0.1f;
|
||||
c.erank_reg = 0.1f;
|
||||
c.means_lr = 5e-05f;
|
||||
c.means_lr_final = 1e-07f;
|
||||
return true;
|
||||
}
|
||||
if (name == "linear-color") {
|
||||
c.splat_color_gamut = "ACEScg";
|
||||
c.splat_color_is_linear = true;
|
||||
c.image_color_gamut = "Rec.2020";
|
||||
c.image_color_is_linear = false;
|
||||
c.background_mode = "noise";
|
||||
return true;
|
||||
}
|
||||
if (name == "synthetic") {
|
||||
c.min_init_fraction = 0.1f;
|
||||
c.use_bilateral_grid = false;
|
||||
c.use_ppisp = false;
|
||||
c.use_bilateral_grid_for_geometry = false;
|
||||
c.long_axis_split_opacity_k = {0.5f, 0.6f, 25000.0f};
|
||||
return true;
|
||||
}
|
||||
if (name == "meshing") {
|
||||
c.primitive = "3dgut";
|
||||
c.sh_degree = 0;
|
||||
c.sh_reg = 10.0f;
|
||||
c.overexposure_reg = 10.0f;
|
||||
c.background_mode = "noise";
|
||||
c.depth_distortion_reg = 0.01f;
|
||||
c.normal_distortion_reg = 0.01f;
|
||||
c.mean_median_depth_weight = 0.01f;
|
||||
c.median_depth_normal_reg_weight = 0.01f;
|
||||
c.normal_supervision_weight = 0.01f;
|
||||
c.median_normal_supervision_weight = 0.01f;
|
||||
c.median_render_normal_reg_weight = 0.01f;
|
||||
c.erank_reg = 0.01f;
|
||||
c.erank_reg_s3 = 0.01f;
|
||||
return true;
|
||||
}
|
||||
if (name == "academic-baseline") {
|
||||
c.eval_mode = "interval";
|
||||
c.eval_interval = 8;
|
||||
c.center_method = "gsplat";
|
||||
c.orientation_method = "gsplat";
|
||||
c.max_batch_per_epoch = 387420489;
|
||||
c.load_depths = false;
|
||||
c.load_normals = false;
|
||||
c.primitive = "3dgs";
|
||||
c.relative_scale = 1.0f;
|
||||
c.use_bilateral_grid = false;
|
||||
c.use_bilateral_grid_for_geometry = false;
|
||||
c.use_ppisp = false;
|
||||
c.use_revised_densification = false;
|
||||
c.densify_loss_map_mode = "none";
|
||||
c.use_long_axis_split = false;
|
||||
c.use_fused_proj_bwd_optim = false;
|
||||
c.quantization_level = 0;
|
||||
c.max_screen_size = kSsplatInf;
|
||||
c.max_world_size = kSsplatInf;
|
||||
c.suppress_initial_scales = false;
|
||||
c.scale_init = 0.1f;
|
||||
c.opacity_init = 0.5f;
|
||||
c.depth_distortion_reg = 0.0f;
|
||||
c.normal_reg_weight = 0.0f;
|
||||
c.alpha_reg_weight = 0.0f;
|
||||
c.alpha_loss_weight = 0.0f;
|
||||
c.alpha_loss_weight_under = 0.0f;
|
||||
c.erank_reg = 0.0f;
|
||||
c.erank_reg_s3 = 0.0f;
|
||||
c.quat_norm_reg = 0.0f;
|
||||
c.sh_reg = 0.0f;
|
||||
c.normal_supervision_weight = 0.0f;
|
||||
c.opacity_reg = 0.01f;
|
||||
c.scale_reg = 0.01f;
|
||||
c.max_steps = 30000;
|
||||
c.use_scale_agnostic_mean = false;
|
||||
c.use_per_splat_bias_correction = false;
|
||||
c.means_lr = 0.00016f;
|
||||
c.means_lr_final = 1.6e-06f;
|
||||
c.scales_lr = 0.005f;
|
||||
c.scales_lr_final = std::nullopt;
|
||||
c.quats_lr = 0.001f;
|
||||
c.opacities_lr = 0.05f;
|
||||
c.features_dc_lr = 0.0025f;
|
||||
c.features_sh_lr = 0.000125f;
|
||||
return true;
|
||||
}
|
||||
(void)c;
|
||||
return false;
|
||||
}
|
||||
@@ -139,11 +139,11 @@ def _eq(a, b) -> bool:
|
||||
def test_native_config_matches_dataclass_defaults(class_name, preset):
|
||||
"""to_native_config(PresetClass()) == SsplatConfig() + apply_preset(name).
|
||||
|
||||
Both sides claim to encode the same defaults: the generator bakes them
|
||||
into cli_config.h at codegen time, the adapter reads them off the live
|
||||
dataclasses. If they disagree, either the generator is stale or a preset
|
||||
branch is wrong -- and every native run would silently use a different
|
||||
config than the CLI flags describe.
|
||||
Both sides claim to encode the same defaults: src/config/TrainConfig.h is
|
||||
the source of truth, the adapter reads them off the live dataclasses. If
|
||||
they disagree, the dataclass copy has drifted or a preset branch is wrong
|
||||
-- and every native run would silently use a different config than the CLI
|
||||
flags describe.
|
||||
"""
|
||||
cls = getattr(trainer_mod, class_name)
|
||||
assert preset_name(cls(data=Path("/nonexistent"))) == preset
|
||||
|
||||
@@ -1,463 +0,0 @@
|
||||
#!/usr/bin/env python3
|
||||
"""Generate the standalone CLI trainer's config header from the Python
|
||||
training config dataclasses.
|
||||
|
||||
Source of truth: TrainerConfig (+ presets) in modules/trainer.py and the
|
||||
nested SpirulaeSplatDataParserConfig / SpirulaeSplatDataManagerConfig /
|
||||
SpirulaeSplatModelConfig / OptimizerConfig dataclasses. This script parses
|
||||
them with `ast` (no torch import, so it runs on a fresh checkout before
|
||||
csrc.so exists) and emits
|
||||
src/app/generated/cli_config.h
|
||||
containing:
|
||||
- struct SsplatConfig: every training config field, flattened, with the
|
||||
Python defaults baked in;
|
||||
- SSPLAT_CONFIG_FIELDS(X): an X-macro over (member, cli_key, group,
|
||||
choices, help) that the CLI's generic parser/help printer expands;
|
||||
- ssplat_apply_preset(): one branch per tyro preset subcommand, assigning
|
||||
exactly the fields whose defaults the preset class overrides.
|
||||
|
||||
Flattening: nested config names are dropped (--model.sh-degree becomes
|
||||
--sh-degree). Name collisions across groups must be resolved in RENAMES
|
||||
below; the script errors on any unlisted collision so a new conflicting
|
||||
Python field cannot silently shadow an existing flag.
|
||||
|
||||
Type mapping:
|
||||
int/float/bool/str/Path -> int / float / bool / std::string
|
||||
Optional[int|float] -> std::optional<int|float> ("none" on CLI)
|
||||
Optional[str|Path] -> std::string, "" = None
|
||||
Literal[str...] (opt. None) -> std::string + choices ("none" for None)
|
||||
Literal[True, False, None] -> std::optional<bool>
|
||||
Tuple[...] -> std::array<int|float, N> (N CLI values)
|
||||
Union[bool, int] -> float via TYPE_OVERRIDES (see entry)
|
||||
"""
|
||||
|
||||
import ast
|
||||
import os
|
||||
import sys
|
||||
from dataclasses import dataclass, field as dc_field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
REPO = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
OUT_PATH = os.path.join(
|
||||
REPO, "src/app/generated/cli_config.h")
|
||||
|
||||
# (group, source file, class name), in flattening order.
|
||||
CONFIG_SOURCES = [
|
||||
("trainer", "spirulae_splat/modules/trainer.py", "TrainerConfig"),
|
||||
("dataparser", "spirulae_splat/modules/dataparser.py", "SpirulaeSplatDataParserConfig"),
|
||||
("datamanager", "spirulae_splat/modules/datamanager.py", "SpirulaeSplatDataManagerConfig"),
|
||||
("model", "spirulae_splat/modules/model.py", "SpirulaeSplatModelConfig"),
|
||||
("optimizer", "spirulae_splat/modules/optimizer.py", "OptimizerConfig"),
|
||||
]
|
||||
|
||||
# TrainerConfig fields that hold nested config objects (flattened separately).
|
||||
NESTED_CONFIG_FIELDS = {"dataparser", "datamanager", "model", "optimizer"}
|
||||
|
||||
# Cross-group name collisions -> flattened CLI/struct name. The script
|
||||
# errors on collisions not listed here.
|
||||
RENAMES = {
|
||||
# model.split_batch is the engine grad-accumulation path; the
|
||||
# datamanager flag is the legacy Python data-path OOM workaround
|
||||
# (a no-op with the C++ DataManager).
|
||||
("datamanager", "split_batch"): "dm_split_batch",
|
||||
}
|
||||
|
||||
# Fields whose annotation the generic mapper can't (or shouldn't) handle:
|
||||
# (group, name) -> (kind, help suffix appended to the docstring).
|
||||
TYPE_OVERRIDES = {
|
||||
# Union[bool, int]: True = probe the image resolution, int = fixed factor.
|
||||
("dataparser", "rescale_camera_to_fit"): (
|
||||
("Float", 0, None),
|
||||
" [CLI: 0 = off, -1 = auto-detect from image resolution, > 0 = divide intrinsics by this]"),
|
||||
}
|
||||
|
||||
# Preset subcommands, matching ss_trainer.py's tyro mapping. "3dgs" is the
|
||||
# base TrainerConfig (no overrides) and the default when no preset is given.
|
||||
PRESETS = [
|
||||
("3dgs", "TrainerConfig"),
|
||||
("360-camera", "TrainerConfig360Camera"),
|
||||
("in-the-wild", "TrainerConfigInTheWild"),
|
||||
("linear-color", "TrainerConfigLinear"),
|
||||
("synthetic", "TrainerConfigSynthetic"),
|
||||
("meshing", "TrainerConfigMeshing"),
|
||||
("academic-baseline", "TrainerConfigAcademicBaseline"),
|
||||
]
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# AST helpers
|
||||
# ===========================================================================
|
||||
|
||||
class SkipField(Exception):
|
||||
pass
|
||||
|
||||
|
||||
REQUIRED = object()
|
||||
|
||||
|
||||
def safe_eval(node):
|
||||
"""Evaluate a default-value expression without importing the module.
|
||||
Supports the literal subset the configs actually use."""
|
||||
if isinstance(node, ast.Constant):
|
||||
return node.value
|
||||
if isinstance(node, (ast.Tuple, ast.List)):
|
||||
return tuple(safe_eval(e) for e in node.elts)
|
||||
if isinstance(node, ast.UnaryOp) and isinstance(node.op, ast.USub):
|
||||
return -safe_eval(node.operand)
|
||||
if isinstance(node, ast.BinOp):
|
||||
a, b = safe_eval(node.left), safe_eval(node.right)
|
||||
if isinstance(node.op, ast.Add): return a + b
|
||||
if isinstance(node.op, ast.Sub): return a - b
|
||||
if isinstance(node.op, ast.Mult): return a * b
|
||||
if isinstance(node.op, ast.Div): return a / b
|
||||
if isinstance(node.op, ast.Pow): return a ** b
|
||||
if isinstance(node, ast.Call):
|
||||
fname = node.func.id if isinstance(node.func, ast.Name) else None
|
||||
if fname == "float" and len(node.args) == 1:
|
||||
return float(safe_eval(node.args[0]))
|
||||
if fname == "int" and len(node.args) == 1:
|
||||
return int(safe_eval(node.args[0]))
|
||||
if fname == "Path" and len(node.args) == 1:
|
||||
return str(safe_eval(node.args[0]))
|
||||
if fname == "field":
|
||||
raise SkipField() # nested config default_factory
|
||||
raise ValueError(f"unsupported default expression: {ast.dump(node)}")
|
||||
|
||||
|
||||
@dataclass
|
||||
class FieldSpec:
|
||||
group: str
|
||||
pyname: str # name inside its Python config class
|
||||
cname: str = "" # flattened struct member / CLI key (after renames)
|
||||
kind: str = "" # Int/Float/Bool/String/OptInt/OptFloat/OptBool/TupleI/TupleF
|
||||
arity: int = 0 # tuple arity
|
||||
choices: tuple = () # allowed strings for Literal fields ("" entries excluded)
|
||||
allow_none: bool = False # string field where "none" -> "" is legal
|
||||
default: object = None
|
||||
required: bool = False
|
||||
help: str = ""
|
||||
|
||||
|
||||
def classify_annotation(node, group, name):
|
||||
"""Return (kind, arity, choices) for an annotation AST node."""
|
||||
if (group, name) in TYPE_OVERRIDES:
|
||||
return TYPE_OVERRIDES[(group, name)][0]
|
||||
if isinstance(node, ast.Name):
|
||||
base = {"int": "Int", "float": "Float", "bool": "Bool",
|
||||
"str": "String", "Path": "String"}.get(node.id)
|
||||
if base:
|
||||
return (base, 0, None)
|
||||
raise ValueError(f"{group}.{name}: unsupported annotation {node.id}")
|
||||
if isinstance(node, ast.Subscript):
|
||||
outer = node.value.id if isinstance(node.value, ast.Name) else None
|
||||
inner = node.slice
|
||||
if outer == "Optional":
|
||||
k, _, _ = classify_annotation(inner, group, name)
|
||||
if k == "Int": return ("OptInt", 0, None)
|
||||
if k == "Float": return ("OptFloat", 0, None)
|
||||
if k == "String": return ("String", 0, "none") # "" = None
|
||||
raise ValueError(f"{group}.{name}: unsupported Optional[{k}]")
|
||||
if outer == "Literal":
|
||||
elts = inner.elts if isinstance(inner, ast.Tuple) else [inner]
|
||||
vals = [safe_eval(e) for e in elts]
|
||||
if set(vals) <= {True, False, None}:
|
||||
return ("OptBool", 0, None)
|
||||
strs = tuple(v for v in vals if isinstance(v, str))
|
||||
has_none = any(v is None for v in vals)
|
||||
if len(strs) + has_none == len(vals):
|
||||
return ("String", 0, (strs, has_none))
|
||||
raise ValueError(f"{group}.{name}: unsupported Literal {vals}")
|
||||
if outer == "Tuple":
|
||||
elts = inner.elts if isinstance(inner, ast.Tuple) else [inner]
|
||||
kinds = [classify_annotation(e, group, name)[0] for e in elts]
|
||||
if all(k == "Int" for k in kinds):
|
||||
return ("TupleI", len(kinds), None)
|
||||
if all(k in ("Int", "Float") for k in kinds):
|
||||
return ("TupleF", len(kinds), None)
|
||||
raise ValueError(f"{group}.{name}: unsupported Tuple {kinds}")
|
||||
if outer == "Union":
|
||||
elts = inner.elts if isinstance(inner, ast.Tuple) else [inner]
|
||||
names = [e.id if isinstance(e, ast.Name) else None for e in elts]
|
||||
raise ValueError(
|
||||
f"{group}.{name}: Union[{names}] needs a TYPE_OVERRIDES entry")
|
||||
raise ValueError(f"{group}.{name}: unsupported annotation {ast.dump(node)}")
|
||||
|
||||
|
||||
def collapse_doc(doc):
|
||||
return " ".join(doc.split())
|
||||
|
||||
|
||||
def find_class(tree, class_name):
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.ClassDef) and node.name == class_name:
|
||||
return node
|
||||
raise ValueError(f"class {class_name} not found")
|
||||
|
||||
|
||||
def extract_fields(tree, group, class_name):
|
||||
"""Ordered FieldSpec list for one config dataclass."""
|
||||
cls = find_class(tree, class_name)
|
||||
fields = []
|
||||
body = cls.body
|
||||
for i, stmt in enumerate(body):
|
||||
if not (isinstance(stmt, ast.AnnAssign) and isinstance(stmt.target, ast.Name)):
|
||||
continue
|
||||
name = stmt.target.id
|
||||
if group == "trainer" and name in NESTED_CONFIG_FIELDS:
|
||||
continue
|
||||
# Docstring: string literal expression immediately following.
|
||||
doc = ""
|
||||
if i + 1 < len(body):
|
||||
nxt = body[i + 1]
|
||||
if (isinstance(nxt, ast.Expr) and isinstance(nxt.value, ast.Constant)
|
||||
and isinstance(nxt.value.value, str)):
|
||||
doc = collapse_doc(nxt.value.value)
|
||||
|
||||
try:
|
||||
kind, arity, extra = classify_annotation(stmt.annotation, group, name)
|
||||
except SkipField:
|
||||
continue
|
||||
choices, allow_none = (), False
|
||||
if kind == "String" and extra == "none":
|
||||
allow_none = True
|
||||
elif kind == "String" and isinstance(extra, tuple):
|
||||
choices, allow_none = extra
|
||||
|
||||
if (group, name) in TYPE_OVERRIDES:
|
||||
doc += TYPE_OVERRIDES[(group, name)][1]
|
||||
|
||||
default, required = None, False
|
||||
if stmt.value is None:
|
||||
required = True
|
||||
else:
|
||||
try:
|
||||
default = safe_eval(stmt.value)
|
||||
except SkipField:
|
||||
continue
|
||||
|
||||
fields.append(FieldSpec(
|
||||
group=group, pyname=name, kind=kind, arity=arity,
|
||||
choices=choices, allow_none=allow_none,
|
||||
default=default, required=required, help=doc))
|
||||
return fields
|
||||
|
||||
|
||||
def extract_preset_overrides(tree, class_name):
|
||||
"""{(group, pyname): value} for a preset subclass's default_factory
|
||||
keyword overrides. Base TrainerConfig -> {}."""
|
||||
if class_name == "TrainerConfig":
|
||||
cls = find_class(tree, class_name)
|
||||
return {}, collapse_doc(ast.get_docstring(cls) or "")
|
||||
cls = find_class(tree, class_name)
|
||||
doc = collapse_doc(ast.get_docstring(cls) or "")
|
||||
overrides = {}
|
||||
for stmt in cls.body:
|
||||
if not (isinstance(stmt, ast.AnnAssign) and isinstance(stmt.target, ast.Name)):
|
||||
continue
|
||||
group = stmt.target.id
|
||||
if group not in NESTED_CONFIG_FIELDS:
|
||||
raise ValueError(f"{class_name}: unexpected preset field {group}")
|
||||
# field(default_factory=lambda: SomeConfig(kw=..., ...))
|
||||
call = stmt.value
|
||||
assert isinstance(call, ast.Call) and call.func.id == "field", \
|
||||
f"{class_name}.{group}: expected field(default_factory=lambda: ...)"
|
||||
factory = next(kw.value for kw in call.keywords
|
||||
if kw.arg == "default_factory")
|
||||
assert isinstance(factory, ast.Lambda), \
|
||||
f"{class_name}.{group}: expected a lambda default_factory"
|
||||
inner = factory.body
|
||||
assert isinstance(inner, ast.Call), \
|
||||
f"{class_name}.{group}: expected a config-constructor call"
|
||||
for kw in inner.keywords:
|
||||
overrides[(group, kw.arg)] = safe_eval(kw.value)
|
||||
return overrides, doc
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# C++ emission
|
||||
# ===========================================================================
|
||||
|
||||
def c_str(s):
|
||||
return '"' + s.replace("\\", "\\\\").replace('"', '\\"') + '"'
|
||||
|
||||
|
||||
def float_lit(v):
|
||||
v = float(v)
|
||||
if v != v:
|
||||
raise ValueError("NaN default not supported")
|
||||
if v == float("inf"):
|
||||
return "std::numeric_limits<float>::infinity()"
|
||||
if v == float("-inf"):
|
||||
return "-std::numeric_limits<float>::infinity()"
|
||||
s = repr(v)
|
||||
return s + ("f" if ("." in s or "e" in s or "E" in s) else ".0f")
|
||||
|
||||
|
||||
def cpp_value(f: FieldSpec, v):
|
||||
if f.kind == "Int":
|
||||
return str(int(v))
|
||||
if f.kind == "Float":
|
||||
# bool default for Union[bool,int] overrides: False -> 0, True -> -1 (auto)
|
||||
if isinstance(v, bool):
|
||||
v = -1.0 if v else 0.0
|
||||
return float_lit(v)
|
||||
if f.kind == "Bool":
|
||||
return "true" if v else "false"
|
||||
if f.kind == "String":
|
||||
return c_str("" if v is None else str(v))
|
||||
if f.kind == "OptInt":
|
||||
return "std::nullopt" if v is None else str(int(v))
|
||||
if f.kind == "OptFloat":
|
||||
return "std::nullopt" if v is None else float_lit(v)
|
||||
if f.kind == "OptBool":
|
||||
return "std::nullopt" if v is None else ("true" if v else "false")
|
||||
if f.kind == "TupleI":
|
||||
return "{" + ", ".join(str(int(x)) for x in v) + "}"
|
||||
if f.kind == "TupleF":
|
||||
return "{" + ", ".join(float_lit(x) for x in v) + "}"
|
||||
raise ValueError(f.kind)
|
||||
|
||||
|
||||
def cpp_type(f: FieldSpec):
|
||||
return {
|
||||
"Int": "int", "Float": "float", "Bool": "bool",
|
||||
"String": "std::string",
|
||||
"OptInt": "std::optional<int>",
|
||||
"OptFloat": "std::optional<float>",
|
||||
"OptBool": "std::optional<bool>",
|
||||
"TupleI": f"std::array<int, {f.arity}>",
|
||||
"TupleF": f"std::array<float, {f.arity}>",
|
||||
}[f.kind]
|
||||
|
||||
|
||||
def choices_str(f: FieldSpec):
|
||||
parts = list(f.choices)
|
||||
if f.allow_none:
|
||||
parts.append("none")
|
||||
return "|".join(parts)
|
||||
|
||||
|
||||
def write_if_changed(path, text):
|
||||
old = None
|
||||
if os.path.exists(path):
|
||||
with open(path) as fp:
|
||||
old = fp.read()
|
||||
if old != text:
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
with open(path, "w") as fp:
|
||||
fp.write(text)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def main():
|
||||
trees = {}
|
||||
all_fields = []
|
||||
for group, rel, class_name in CONFIG_SOURCES:
|
||||
path = os.path.join(REPO, rel)
|
||||
with open(path) as fp:
|
||||
trees[group] = ast.parse(fp.read())
|
||||
all_fields += extract_fields(trees[group], group, class_name)
|
||||
|
||||
# Flatten with collision detection.
|
||||
by_cname = {}
|
||||
for f in all_fields:
|
||||
f.cname = RENAMES.get((f.group, f.pyname), f.pyname)
|
||||
if f.cname in by_cname:
|
||||
other = by_cname[f.cname]
|
||||
raise SystemExit(
|
||||
f"error: flattened name collision '{f.cname}' between "
|
||||
f"{other.group}.{other.pyname} and {f.group}.{f.pyname}; "
|
||||
f"add a RENAMES entry in {os.path.basename(__file__)}")
|
||||
by_cname[f.cname] = f
|
||||
|
||||
# Presets (parsed from trainer.py).
|
||||
presets = []
|
||||
for preset_name, class_name in PRESETS:
|
||||
overrides, doc = extract_preset_overrides(trees["trainer"], class_name)
|
||||
items = []
|
||||
for (group, pyname), value in overrides.items():
|
||||
key = RENAMES.get((group, pyname), pyname)
|
||||
if key not in by_cname or by_cname[key].group != group:
|
||||
raise SystemExit(
|
||||
f"error: preset {preset_name} overrides unknown field "
|
||||
f"{group}.{pyname}")
|
||||
items.append((by_cname[key], value))
|
||||
presets.append((preset_name, doc, items))
|
||||
|
||||
# ---- Emit ----
|
||||
L = []
|
||||
L.append("#pragma once")
|
||||
L.append("")
|
||||
L.append("// AUTO-GENERATED by tools/codegen/generate_cli_config.py -- DO NOT EDIT.")
|
||||
L.append("// Source of truth: the Python training config dataclasses")
|
||||
L.append("// (TrainerConfig + presets, SpirulaeSplatDataParserConfig,")
|
||||
L.append("// SpirulaeSplatDataManagerConfig, SpirulaeSplatModelConfig,")
|
||||
L.append("// OptimizerConfig). Re-run the generator after editing those.")
|
||||
L.append("")
|
||||
L.append("#include <array>")
|
||||
L.append("#include <limits>")
|
||||
L.append("#include <optional>")
|
||||
L.append("#include <string>")
|
||||
L.append("")
|
||||
L.append("struct SsplatConfig {")
|
||||
cur_group = None
|
||||
for f in all_fields:
|
||||
if f.group != cur_group:
|
||||
cur_group = f.group
|
||||
src = dict((g, c) for g, _, c in CONFIG_SOURCES)[cur_group]
|
||||
L.append(f" // ==== {cur_group} ({src}) ====")
|
||||
default = "{}" if f.required else cpp_value(f, f.default)
|
||||
req = " // REQUIRED" if f.required else ""
|
||||
L.append(f" {cpp_type(f)} {f.cname} = {default};{req}")
|
||||
L.append("};")
|
||||
L.append("")
|
||||
L.append("// X(member, cli_key, pyname, group, choices, help). cli_key uses '_';")
|
||||
L.append("// the parser treats '-' and '_' as equivalent. pyname is the field's")
|
||||
L.append("// original name inside its Python config class (differs from cli_key")
|
||||
L.append("// only for RENAMES entries); group is the Python sub-config it lives")
|
||||
L.append("// in. choices is a '|' list for string fields ('' = free-form);")
|
||||
L.append("// 'none' selects the empty string.")
|
||||
L.append("#define SSPLAT_CONFIG_FIELDS(X) \\")
|
||||
for f in all_fields:
|
||||
L.append(f" X({f.cname}, {c_str(f.cname)}, {c_str(f.pyname)}, "
|
||||
f"{c_str(f.group)}, {c_str(choices_str(f))}, {c_str(f.help)}) \\")
|
||||
L.append(" /* end */")
|
||||
L.append("")
|
||||
L.append("// Required fields (no Python default). Checked after flag parsing.")
|
||||
req_names = [f.cname for f in all_fields if f.required]
|
||||
L.append("#define SSPLAT_CONFIG_REQUIRED_FIELDS(X) \\")
|
||||
for name in req_names:
|
||||
L.append(f" X({name}) \\")
|
||||
L.append(" /* end */")
|
||||
L.append("")
|
||||
L.append("struct SsplatPresetInfo { const char* name; const char* help; };")
|
||||
L.append("inline constexpr SsplatPresetInfo kSsplatPresets[] = {")
|
||||
for preset_name, doc, _ in presets:
|
||||
L.append(f" {{{c_str(preset_name)}, {c_str(doc)}}},")
|
||||
L.append("};")
|
||||
L.append("")
|
||||
L.append("// Apply a preset's default overrides (tyro subcommand equivalent).")
|
||||
L.append("// Returns false for an unknown preset name. \"3dgs\" is the base config.")
|
||||
L.append("inline bool ssplat_apply_preset(SsplatConfig& c, const std::string& name) {")
|
||||
for preset_name, _, items in presets:
|
||||
L.append(f" if (name == {c_str(preset_name)}) {{")
|
||||
for f, value in items:
|
||||
L.append(f" c.{f.cname} = {cpp_value(f, value)};")
|
||||
L.append(" return true;")
|
||||
L.append(" }")
|
||||
L.append(" (void)c;")
|
||||
L.append(" return false;")
|
||||
L.append("}")
|
||||
L.append("")
|
||||
|
||||
changed = write_if_changed(OUT_PATH, "\n".join(L))
|
||||
n_presets = len(presets)
|
||||
print(f"generate_cli_config: {len(all_fields)} fields, {n_presets} presets"
|
||||
f" -> {os.path.relpath(OUT_PATH, REPO)}"
|
||||
f" ({'updated' if changed else 'unchanged'})")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user