Merge pull request #21 from Alexankharin/video-world-experiments

Publish camera-comfyUI to the ComfyUI Registry (+ 4D world nodes)
This commit is contained in:
Alexander Kharin
2026-07-13 11:09:33 +02:00
committed by GitHub
19 changed files with 9466 additions and 33 deletions
+10
View File
@@ -0,0 +1,10 @@
# Excluded from the ComfyUI Registry archive (not from git).
demo_images/
notebooks/
docs/
screenshot1.ply
__pycache__/
models/
.github/
Makefile
install.sh
+28
View File
@@ -0,0 +1,28 @@
name: Publish to Comfy registry
on:
workflow_dispatch:
push:
branches:
- main
paths:
- "pyproject.toml"
permissions:
contents: read
jobs:
publish-node:
name: Publish Custom Node to registry
runs-on: ubuntu-latest
if: ${{ github.repository_owner == 'Alexankharin' }}
steps:
- name: Check out code
uses: actions/checkout@v4
with:
# The SHARP submodule must be materialized so [tool.comfy].includes
# can pack it into the published archive.
submodules: recursive
- name: Publish Custom Node
uses: Comfy-Org/publish-node-action@main
with:
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
+3
View File
@@ -0,0 +1,3 @@
[submodule "submodules/ml-sharpt"]
path = submodules/ml-sharpt
url = https://github.com/apple/ml-sharp
+1227
View File
File diff suppressed because it is too large Load Diff
+2499
View File
File diff suppressed because it is too large Load Diff
+29
View File
@@ -0,0 +1,29 @@
MIT License
Copyright (c) 2026 Alexander Kharin
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE.
---
Note: the bundled directory `submodules/ml-sharpt` contains Apple's ml-sharp
project and is licensed separately under the terms in
`submodules/ml-sharpt/LICENSE` (source) and `submodules/ml-sharpt/LICENSE_MODEL`
(model weights, research-only). The MIT license above does not apply to that
directory.
+88 -2
View File
@@ -14,6 +14,7 @@
* [Installation](#installation)
* [Node Categories](#node-categories)
* [Node Reference](#node-reference)
* [Video → 4D World](#video--4d-world)
* [Workflows](#workflows)
* [Example Workflows](#example-workflows)
* [Contributing](#contributing)
@@ -35,6 +36,14 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
## Installation
### Option A — ComfyUI Manager (recommended)
The node pack is published to the [ComfyUI Registry](https://registry.comfy.org) as **`camera-comfyui`** (publisher `alexk`). In ComfyUI, open **Manager → Custom Nodes Manager**, search for **camera-comfyUI**, and click **Install**, then restart ComfyUI. The registry package bundles the SHARP submodule and installs the base Python requirements automatically; optional CUDA-specific extras (`gsplat`, `vggt`) still follow the manual steps below.
> **Maintainers:** releases are automated — bumping `version` in `pyproject.toml` on `main` triggers `.github/workflows/publish_action.yml`, which publishes the new version to the registry (requires the `REGISTRY_ACCESS_TOKEN` repo secret).
### Option B — Manual install (git)
1. **Clone** into your ComfyUI custom nodes folder:
```bash
@@ -55,6 +64,13 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
* *Optional:* `open3d` for GUI point cloud tools.
**Optional dependencies** (only needed for specific nodes):
* **gsplat** — CUDA-accelerated Gaussian splat rasterizer. Required by `SplatPolish` and used as the fast render backend for `RenderSplat` / `RenderSplats4D*`. Needs a CUDA GPU and a matching PyTorch build: `pip install gsplat`.
* **vggt** — camera pose + depth estimation (`VideoPoseEstimator`). Install with `pip install vggt` (or `pip install git+https://github.com/facebookresearch/vggt.git`), or clone [facebookresearch/vggt](https://github.com/facebookresearch/vggt) as a sibling folder in your ComfyUI root. The `facebook/VGGT-1B` weights (~5 GB) download via `huggingface_hub` on first use.
* **CoTracker3** — point tracking for `EstimateTracks`. No manual install: it is fetched automatically via `torch.hub` on first use.
* **SHARP** — image→splat prediction (`ImageToSplat`, `FisheyeToGaussian`, `VideoToFusedSplats`, `SplatTrajectoryEnricher`). Ships as the existing git submodule at `submodules/ml-sharpt` ([apple/ml-sharp](https://github.com/apple/ml-sharp)) — run `git submodule update --init` after cloning.
4. **Additional Nodes** (for certain workflows):
* Clone the following repositories directly into your `custom_nodes` folder:
@@ -94,13 +110,33 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
* ### Point Cloud Nodes
* `DepthToPointCloud`, `TransformPointCloud`, `ProjectPointCloud`, `PointCloudUnion`
* `PointCloudCleaner`, `LoadPointCloud`, `SavePointCloud`, `ProjectAndClean`
* `PointCloudCleaner`, `LoadPointCloud`, `SavePointCloud`, `ProjectAndClean`, `DepthEdgeFilter`
* ### Trajectory Nodes
* `CameraMotionNode`, `CameraInterpolationNode`, `CameraTrajectoryNode`
* `SaveTrajectory`, `LoadTrajectory`, `PointcloudTrajectoryEnricher`
* ### Gaussian Splat Nodes
* `LoadPlySplat`, `SavePlySplat`, `ImageToSplat`, `FisheyeToGaussian`
* `RotateSplats`, `MergeSplats`, `FuseSplats`, `RenderSplat`
* `VideoToFusedSplats`, `SplatPolish`
* ### 4D Gaussian Splat Nodes
* `MotionMaskFromDepth`, `EstimateTracks`, `TracksToTrajectories`, `SplitSplatsByMask`
* `BuildSplats4D`, `RenderSplats4DFrame`, `RenderSplats4DVideo`
* `SaveSplats4D`, `LoadSplats4D`
* ### Pose Nodes
* `VideoPoseEstimator`, `TrajectoryInvert`, `TrajectoryCompose`
* ### World Nodes
* `DepthScaleAnchor`, `SplatTrajectoryEnricher`, `SphereSplatSeed`
---
## Node Reference
@@ -130,6 +166,51 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
| `VideoCameraMotionSequence` | Processes video frames and depth maps along a camera trajectory, generating reprojected outputs. |
| `DepthFramesToVideo` | Converts a sequence of depth maps into video frame tensors for saving. |
| `VideoMetricDepthEstimate` | Estimates metric depth for a sequence of frames using VideoDepthAnything. |
| `DepthEdgeFilter` | Detects "flying pixel" depth discontinuities and outputs a validity mask (1.0 = valid). |
| `LoadPlySplat` | Loads a 3D Gaussian Splatting `.ply` file into a `GSPLAT` object. |
| `SavePlySplat` | Saves a `GSPLAT` to the ComfyUI output directory as a `.ply` file. |
| `ImageToSplat` | Predicts Gaussian splats from a single image using SHARP. |
| `FisheyeToGaussian` | Reprojects a fisheye view to multiple pinhole angles, predicts splats, rotates and merges them. |
| `RotateSplats` | Applies a 4×4 transform matrix to a splat cloud. |
| `MergeSplats` | Concatenates two `GSPLAT` objects into one. |
| `FuseSplats` | Fuses two splat clouds with weighted voxel merging (keep/discard/average/smart modes). |
| `RenderSplat` | Renders a splat cloud from a camera pose into an image + mask. |
| `VideoToFusedSplats` | Runs SHARP on video keyframes, scale-aligns to metric depth, filters dynamic pixels, and fuses all keyframes into one world-frame splat cloud. |
| `SplatPolish` | Optimizes a world-frame splat cloud against posed video frames (L1 + D-SSIM) using gsplat's differentiable rasterizer. |
| `MotionMaskFromDepth` | Detects dynamic pixels from a depth+pose sequence (1.0 = moving). |
| `EstimateTracks` | Runs CoTracker3 on a video; returns tracks `[T,N,2]` (pixels) and visibility `[T,N]`. |
| `TracksToTrajectories` | Unprojects 2D tracks with depth and camera poses into world-space 3D trajectories `[T,M,3]`. |
| `SplitSplatsByMask` | Projects splat centers into a 2D mask and splits the cloud into inside/outside parts. |
| `BuildSplats4D` | Builds a 4D splat scene: each canonical splat follows a kNN blend of track control-point motions. |
| `RenderSplats4DFrame` | Evaluates the 4D scene at a single time value and renders it from a given camera. |
| `RenderSplats4DVideo` | Interpolates the camera path, sweeps time from start to end, and renders each frame. |
| `SaveSplats4D` | Saves a `GSPLAT4D` scene as an `.npz` archive (plus optional per-frame PLYs). |
| `LoadSplats4D` | Loads a `GSPLAT4D` scene from an `.npz` archive. |
| `VideoPoseEstimator` | VGGT-based per-frame camera poses `[T,4,4]`, depth maps, FOV and depth confidence from a video clip. |
| `TrajectoryInvert` | Inverts each 4×4 pose (world-to-camera ↔ camera-to-world). |
| `TrajectoryCompose` | Per-frame matrix product `A @ B`; a single 4×4 input broadcasts over the other. |
| `DepthScaleAnchor` | Robustly aligns a depth map to a reference depth via disparity-domain scale(+shift). |
| `SplatTrajectoryEnricher` | Expands a splat world along a trajectory: render, outpaint holes with Flux, lift with SHARP, scale-align, smart-stitch. |
| `SphereSplatSeed` | Converts an equirectangular panorama into a Gaussian sphere seeding a 360° world. |
---
## Video → 4D World
Turn a monocular video into a navigable 4D (3D + time) Gaussian splat scene and re-render it from any novel camera trajectory. The reference workflow is **`workflows/video_to_4d_world.json`**; the stages are:
1. **Pose & depth (VGGT)** — `VideoPoseEstimator` estimates per-frame world-to-camera poses `[T,4,4]`, depth maps, FOV and depth confidence from the input frames. Since the depth maps are Z-depths, run `ZDepthToRayDepthNode` before any node that expects ray depth (see caveats below). `DepthEdgeFilter` can additionally mask out flying pixels at depth discontinuities.
2. **Motion masking** — `MotionMaskFromDepth` warps depth between frames using the estimated poses and flags pixels whose residual is too large as dynamic (moving objects vs. static background).
3. **Static splat fusion + polish** — `VideoToFusedSplats` runs SHARP on keyframes, keeps only static pixels (via the motion mask), scale-aligns each keyframe to metric depth, transforms splats into the world frame and fuses them incrementally. `SplatPolish` then fine-tunes the fused cloud photometrically against the posed video frames.
4. **Tracked dynamic 4D Gaussians** — `EstimateTracks` (CoTracker3) tracks a dense point grid across the video; `TracksToTrajectories` lifts the tracks to world-space 3D using depth + poses; `SplitSplatsByMask` separates dynamic splats from the static background; `BuildSplats4D` binds the dynamic canonical splats to track control points via kNN blending, producing a `GSPLAT4D` scene.
5. **Render a novel trajectory** — build any new camera path (e.g. `CameraInterpolationNode`, `TrajectoryCompose` to retarget relative to a source pose) and render with `RenderSplats4DVideo` (or single frames with `RenderSplats4DFrame`). Save/reload scenes with `SaveSplats4D` / `LoadSplats4D`.
### Caveats
* **Z-depth vs ray depth**: depth estimators (including `VideoPoseEstimator`) output Z-depth; point-cloud and splat lifting nodes expect ray depth. Insert `ZDepthToRayDepthNode` where needed, or geometry will bow at wide FOVs.
* **`SplatPolish` requires gsplat + CUDA**: without them it can fall back to the differentiable torch renderer at reduced resolution, which is extremely slow (minutes per 100 iterations).
* **`EstimateTracks` downloads CoTracker3 via `torch.hub` on first use** — expect a one-time download and allow network access.
* **`VideoPoseEstimator` downloads `facebook/VGGT-1B` (~5 GB)** on first use via `huggingface_hub`.
---
@@ -150,6 +231,8 @@ A set of JSON workflows illustrating typical use cases. Each workflow lives in `
| **Pointcloud\_walker.json** | GUI‐based camera control via Open3D |
| **sbs180\_workflow.json** | Generate stereo (side-by-side) wide-angle/fisheye/equirectangular stereo pairs from a high-res input |
| **video_camera.json** | Camera trajectory movement workflow using `wan-vace` for video inpainting. |
| **video_to_4d_world\.json** | Video → 4D world: VGGT poses/depth → motion masking → fused static splats + polish → tracked dynamic 4D Gaussians → novel-trajectory render. |
| **video_to_4d_walkable_world\.json** | Video → 4D WALKABLE world (test-friendly defaults): polished static splats enriched along a walk trajectory (`SplatTrajectoryEnricher`, Flux outpaint + SHARP) → 4D scene → walk-through render + `.ply`/`.npz` exports for free walking in external 3DGS viewers. |
---
@@ -268,9 +351,12 @@ Contributions welcome! Please open issues or PRs to add features, improve docs,
* [x] Add processing to pointcloud or depthmap to remove outlier and lonely points at depth borders.
* [x] Use built-in comfyUI mask type an image.
* [x] Unite nodes into groups to simplify workflows.
* [ ] Create a single workflow for view synthesis.
* [x] Create a single workflow for view synthesis (`video_to_4d_world.json`).
* [x] Implement easier and more flexible camera control - more complex camera movements with more than 2 points.
* [x] Add more examples and documentation for each node.
* [x] Add pointcloud union
* [x] Fix imports for renamed folders (e.g., inpainting_flux)
* [x] Integrate camera movement pipeline with video models (e.g., wan2.1) for smooth, high-quality inpainting along camera trajectories.
* [ ] Compressed export format for 4D scenes (current `.npz` stores raw tensors).
* [ ] SAM2-based refinement of motion masks (current masks come from depth-warp residuals only).
* [ ] Fisheye/equirectangular rendering through gsplat (e.g., via cubemap render + reprojection); the fast CUDA path is currently pinhole-only.
+23 -2
View File
@@ -4,6 +4,27 @@ from .metric_depth_nodes import NODE_CLASS_MAPPINGS as NCM3
from .flux_fisheye_filling_nodes import NODE_CLASS_MAPPINGS as NCM4
from .complex_nodes import NODE_CLASS_MAPPINGS as NCM5
from .video_nodes import NODE_CLASS_MAPPINGS as NCM6
NODE_CLASS_MAPPINGS = {**NCM1, **NCM2, **NCM3, **NCM4, **NCM5, **NCM6}
from .GS_nodes import NODE_CLASS_MAPPINGS as NCM7
__all__ = ["NODE_CLASS_MAPPINGS"]
# Optional node packs: a missing/broken optional dependency must never kill the
# whole extension (mirrors how video_nodes degrades when video_depth_anything
# is unavailable).
try:
from .GS4D_nodes import NODE_CLASS_MAPPINGS as NCM8
except Exception as _exc:
print(f"[camera-comfyUI] Warning: GS4D_nodes could not be loaded, 4D splat nodes disabled: {_exc}")
NCM8 = {}
try:
from .pose_nodes import NODE_CLASS_MAPPINGS as NCM9
except Exception as _exc:
print(f"[camera-comfyUI] Warning: pose_nodes could not be loaded, pose estimation nodes disabled: {_exc}")
NCM9 = {}
try:
from .world_nodes import NODE_CLASS_MAPPINGS as NCM10
except Exception as _exc:
print(f"[camera-comfyUI] Warning: world_nodes could not be loaded, world-building nodes disabled: {_exc}")
NCM10 = {}
NODE_CLASS_MAPPINGS = {**NCM1, **NCM2, **NCM3, **NCM4, **NCM5, **NCM6, **NCM7, **NCM8, **NCM9, **NCM10}
__all__ = ["NODE_CLASS_MAPPINGS"]
+127
View File
@@ -0,0 +1,127 @@
# LingBot-World 2.0 → 4D video: analysis & integration report
*Research date: 2026-07-13. LingBot-World 2.0 was released 2026-07-09, four days before this report.*
## TL;DR
**LingBot-World 2.0 is not a 3D/4D model — it is a camera-pose- and action-conditioned autoregressive video generator.** It outputs only pixels and maintains no explicit geometry. But it has exactly the property that makes a video-generation model useful for 4D reconstruction: **you command the camera trajectory (poses + intrinsics) of every generated frame**, so every output video is a *posed* video. That turns it into a controllable multi-view video factory whose output can be lifted into 4D Gaussian splats by the existing `video_to_4d_world.json` pipeline in this repo — with the pose-estimation step optionally replaced by the commanded poses.
Feasibility verdicts:
| Question | Verdict |
| --- | --- |
| 4D video from a 3D scene (splat/mesh) | **Yes, indirectly** — render the 3D scene to a seed image, then LingBot animates + explores it. 3D enters only as a rendered start frame; there is no native 3D conditioning. |
| 4D Gaussian-splat video from its output | **Feasible and first-party-endorsed** — the LingBot-World paper itself demonstrates reconstructing its generated videos into point clouds with VGGT-class models, the same VGGT this repo already uses. |
| Drop-in ComfyUI use today | **Not yet** — 14B Wan2.2-based weights, no quantized release for v2, no wrapper support yet ([kijai/WanVideoWrapper#1920](https://github.com/kijai/ComfyUI-WanVideoWrapper/issues/1920), [Comfy-Org/ComfyUI#12154](https://github.com/Comfy-Org/ComfyUI/issues/12154)); reference inference is 8×GPU `torchrun`. |
| Commercial use | **v2: no** (CC BY-NC-SA 4.0). **v1: yes** (Apache 2.0). This alone may decide which version to build on. |
---
## 1. What LingBot-World 2.0 actually is
**Repos & papers**
- v2 (current): [Robbyant/lingbot-world-v2](https://github.com/Robbyant/lingbot-world-v2) — "Infinite Worlds with Versatile Interactions", tech report [arXiv:2607.07534](https://arxiv.org/abs/2607.07534), weights [robbyant/lingbot-world-v2-14b-causal-fast](https://huggingface.co/robbyant/lingbot-world-v2-14b-causal-fast). Released 2026-07-09 by Robbyant (embodied-AI subsidiary of Ant Group).
- v1 (deprecated but still useful): [Robbyant/lingbot-world](https://github.com/robbyant/lingbot-world) — "Advancing Open-source World Models", [arXiv:2601.20540](https://arxiv.org/abs/2601.20540), weights `robbyant/lingbot-world-base-cam` / `-base-act` / `-fast`. Released 2026-01-29.
**Architecture (verified against code + paper)**
- Built on **Wan2.2 i2v-A14B**: a two-expert MoE video diffusion model, ~28B total parameters with **14B active** per denoising step (high-noise expert for global structure, low-noise for detail). Ships the Wan2.1 VAE and umT5-XXL text encoder.
- v2 converts it to **causal, chunk-by-chunk autoregressive generation**: latents are generated `chunk_size` latent frames at a time against a **KV cache** with **sink tokens** and a **local attention window** (`run_fast.sh` uses `--local_attn_size 18 --sink_size 6`). A **MoBA mask** ("Mixture of Bidirectional and Autoregressive Attention Mask") mixes bidirectional attention into teacher forcing to stop the long-horizon quality collapse that plagues autoregressive video. Result: the paper demonstrates an **uninterrupted hour-long session with no perceptible quality decay**.
- Two inference modes: `causal_fast` (distilled few-step; drives **720p @ 60 fps** in their real-time deployment) and `causal_pretrain` (40-step CFG; checkpoint still marked TODO). A single-GPU **1.3B variant is described in the paper but not released**.
**Conditioning inputs — the part that matters for 4D** (from `wan/image2video.py` + `wan/utils/cam_utils.py`)
- **Seed image** (`--image`) + **text prompt**: the world is initialized from one image and a background description. This is the *only* way content enters — no 3D input of any kind.
- **Camera trajectory**: `poses.npy` `[T,4,4]` **camera-to-world, OpenCV convention** + `intrinsics.npy` `[T,4]` = `[fx,fy,cx,cy]`. Converted to per-pixel **Plücker ray embeddings** (`get_plucker_embeddings`), folded into the latent grid and injected per-chunk into the DiT (AdaLN per the tech report). Relative poses are translation-normalized (`compute_relative_poses`), and `interpolate_camera_poses` (SLERP) is provided.
- **Keyboard actions**: `wasd_action.npy` (movement) / `ijkl_action.npy` (view) as multi-hot vectors concatenated onto the Plücker conditioning. v2 adds character actions (attack, archery, spell-cast, shoot, jump, glide) and **chunk-wise text events** (weather, entity spawning, time-of-day), plus a VLM-driven "pilot/director" agentic harness.
- v1 README explicitly recommends **[NVIDIA ViPE](https://github.com/nv-tlabs/vipe)** to extract `poses.npy`/`intrinsics.npy` from an *existing real video* — i.e., the official video→control-signal bridge.
**Inference & hardware**
```bash
torchrun --nproc_per_node=8 generate.py --task i2v-A14B --size 480*832 \
--frame_num 361 --ckpt_dir lingbot-world-v2-14b-causal-fast \
--image examples/03/image.jpg --action_path examples/03 \
--infer_mode causal_fast --dit_fsdp --t5_fsdp --ulysses_size 8 \
--local_attn_size 18 --sink_size 6
```
- Reference: 8×GPU (FSDP + Ulysses sequence parallel), 480×832, 361 frames (`frame_num` must be 4n+1). Single-GPU runs auto-enable `--offload_model` (T5/DiT swapped to CPU between stages) — expect 80GB-class VRAM for comfortable 14B bf16 inference; there is **no quantized v2 release yet**. v1 has a community **4-bit quant** and `--t5_cpu`, and supports up to 961 frames (~1 min @ 16 fps).
- Requirements: `torch >= 2.4.0`, `flash_attn`.
**License** — v2 code *and* weights are **CC BY-NC-SA 4.0 (non-commercial, share-alike)**; v1 is **Apache 2.0**. Anything commercial built on v2 outputs is off the table; v1 remains the commercially safe option at lower quality/horizon.
---
## 2. Can it turn 3D into 4D video?
**Yes, with the 3D scene entering as a rendered image, not as geometry.** The paper is explicit that the world "is initialized from an initial image and its background description" — there is no splat/mesh/point-cloud conditioning path, and the model "operates without an explicit notion of geometry."
The working recipe, using nodes already in this repo:
1. **Render a seed view** of your static 3D asset: `LoadPlySplat` → `RenderSplat` (or a mesh render) at 832×480+, from a pose with good scene coverage.
2. **Author the camera trajectory you want** in the splat's own coordinate frame (`CameraInterpolationNode` / `CameraTrajectoryNode`), convert to camera-to-world OpenCV `poses.npy` + `intrinsics.npy`.
3. **Feed image + poses + actions/text-events to LingBot-World.** The model animates the scene (wind, characters, weather, spawned entities via text events) while following your camera — i.e., it *invents plausible dynamics* for your static 3D scene. This is "3D → 4D video" in the sense of *generating* the time dimension, not simulating it: physics is learned and imperfect, and the output will drift from your 3D asset's exact geometry the further the camera goes from the seed view.
4. **Optionally lift the result back to 4D splats** (section 3) so the animated version of your scene becomes re-renderable from any camera.
Caveat on fidelity: only the seed frame is constrained by your 3D input. Occluded/unseen regions are hallucinated. For higher fidelity to the source scene you can seed successive generations from renders at multiple poses and stitch — the same strategy `SplatTrajectoryEnricher` already uses with Flux outpainting, but with LingBot providing temporally coherent *video* instead of stills.
---
## 3. Feasibility: 4D Gaussian-splat video from LingBot output
**This is the strongest part of the story.** Three findings, all verified against primary sources:
1. **Posed video for free.** Because generation is conditioned on `poses.npy`/`intrinsics.npy`, every generated frame comes with a commanded camera. A monocular real video gives you poses only after VGGT/COLMAP estimation; LingBot gives you the trajectory you asked for. (Treat commanded poses as *approximate* — the model follows them but is not geometrically exact; see limitations.)
2. **First-party evidence that reconstruction works.** The LingBot-World paper itself demonstrates: *"by leveraging large-scale 3D reconstruction foundation models [lin2025depth, wang2025vggt], we can further convert the generated video sequences into high-quality scene point clouds"*, with point clouds showing *"strong spatial coherence across frames"* (Fig. 16, [arXiv:2601.20540](https://arxiv.org/html/2601.20540v1)). That is literally VGGT — the model behind this repo's `VideoPoseEstimator` — applied to LingBot output by its own authors.
3. **Long-horizon consistency is the v2 headline.** Landmarks stay structurally intact after being out of view for up to ~60 s (v1) and v2 extends coherent generation to hour scale with no perceptible decay. Long consistent orbits are exactly what splat optimization needs.
**How it maps onto known video-to-4D paradigms:**
- **CAT4D-style** ([arXiv:2411.18613](https://arxiv.org/abs/2411.18613)): camera/time-disentangled video diffusion → deformable 3DGS optimization. LingBot is not time-disentangled (you cannot freeze time and move the camera — camera and time advance together in one causal stream), so you *cannot* get true simultaneous multi-view of a dynamic instant from a single run.
- **Monocular 4D lifting** (this repo's pipeline): works on any single posed video — LingBot output qualifies directly and improves on real footage by letting you *choose* a camera path that orbits/parallaxes around the action, which is the single biggest quality lever for monocular 4D reconstruction.
- **Multi-run multi-view**: re-running with the same seed image but different trajectories gives multiple views of the *same static scene* but **different sampled dynamics** (different seeds/action outcomes per run) — usable for static splat fusion, **not** for dynamic 4D supervision. Keep dynamics within one continuous run.
**Bottom line:** treat LingBot-World as a *trajectory-controllable monocular video source* feeding the existing 4D pipeline; don't expect synchronized multi-view rigs out of it.
---
## 4. Concrete pipeline: video → 4D video / 4D splats
### Path A — real video in, 4D world out, LingBot as the world extender
Your existing `video_to_4d_world.json` already handles real-video → 4D. LingBot adds value where that pipeline is weakest: viewpoints the source video never saw.
1. **Base 4D scene from the real video** (existing flow): `VideoPoseEstimator` (VGGT poses/depth) → `ZDepthToRayDepthNode` → `MotionMaskFromDepth` → `VideoToFusedSplats` + `SplatPolish` (static) → `EstimateTracks`/`TracksToTrajectories`/`SplitSplatsByMask`/`BuildSplats4D` (dynamic) → `GSPLAT4D`.
2. **Extract control signals from the same video** with ViPE (officially recommended) or reuse the VGGT poses: `VideoPoseEstimator` outputs world-to-camera `[T,4,4]` → `TrajectoryInvert` → camera-to-world OpenCV → export `poses.npy` + `intrinsics.npy` (VGGT's FOV output gives `fx,fy`; `cx,cy` = image center). *(Small new node needed: `TrajectoryToNpyExport` — trivial, ~20 lines.)*
3. **Continue the world where the video ends**: last real frame = LingBot seed image; author an exploration trajectory (orbit, dolly, walk) continuing from the last real pose; generate 361+ frames.
4. **Lift the generated segment** through the same stage-1 flow and **fuse into the base scene**: `FuseSplats`/`MergeSplats` for statics (scale-anchor with `DepthScaleAnchor` against the base scene's depth), separate `BuildSplats4D` time range for new dynamics. Result: a 4D world larger than the source footage.
### Path B — single image or 3D scene in, 4D splat video out
1. **Seed**: any image, or a render of an existing splat (`RenderSplat`) / mesh.
2. **Trajectory design**: slow orbit or arc around the subject + gentle forward motion — maximize parallax, avoid pure rotation (no baseline → no geometry). Keep FOV fixed; write `poses.npy`/`intrinsics.npy` (c2w, OpenCV; translations get normalized internally, so keep the trajectory scale moderate and re-anchor metric scale later with `DepthScaleAnchor`).
3. **Generate** with `causal_fast`, 480×832, 361 frames; drive dynamics with keyboard/character actions and chunk-wise text events ("a horse gallops through", "rain starts").
4. **Reconstruct** — two pose options:
- *Trust-but-verify (recommended)*: run `VideoPoseEstimator` on the generated frames anyway; compare with commanded poses (`TrajectoryCompose` of one with `TrajectoryInvert` of the other should be ≈ identity); use VGGT's poses for reconstruction, commanded poses as sanity check. This absorbs the model's camera-following error.
- *Fast path*: use commanded poses directly, skip VGGT pose estimation, still run its depth head (or `VideoMetricDepthEstimate`) for the depth maps the lifting nodes need.
5. **Lift to 4D**: identical to the existing workflow — motion mask → static fusion (`VideoToFusedSplats` + `SplatPolish`) → tracks (`EstimateTracks` is CoTracker3, works fine on generated footage) → `BuildSplats4D` → `RenderSplats4DVideo` along any novel camera path → `SaveSplats4D`.
### Integration notes for camera-comfyUI
- **Coordinate conventions align well**: LingBot uses OpenCV c2w + `[fx,fy,cx,cy]`, this repo's `TRAJECTORY` is 4×4 matrices with `TrajectoryInvert`/`TrajectoryCompose` already available. Needed glue: (a) `TrajectoryToNpyExport` / `NpyToTrajectory` nodes, (b) optionally a `LingBotGenerate` node wrapping `generate.py` via subprocess for remote/8-GPU boxes — running 14B in-process inside ComfyUI is not realistic today.
- **ComfyUI-native inference isn't there yet**: WanVideoWrapper/ComfyUI support for LingBot checkpoints is an open request blocked on VRAM/quantization ([#1920](https://github.com/kijai/ComfyUI-WanVideoWrapper/issues/1920), [#12154](https://github.com/Comfy-Org/ComfyUI/issues/12154)). Because it's Wan2.2-architecture, wrapper support and GGUF/FP8 quants are likely to appear quickly; the causal KV-cache/sink/MoBA inference loop is custom, so a naive Wan2.2 loader won't reproduce long-horizon behavior.
- **Pragmatic hardware ladder**: (1) today, single-image experiments on v1 `base-cam` 4-bit quant (Apache 2.0, 480p, camera-pose conditioned — same poses.npy interface) on a 24 GB GPU; (2) v2 14B on a rented 8×A100/H100 node or single 80 GB GPU with offload; (3) wait for the announced 1.3B v2 release for true single-GPU local use.
### Known limitations
- **No geometry inside the model** — all 3D/4D structure comes from post-hoc reconstruction; physics is "imperfect" by the authors' own admission.
- **Camera-following error**: commanded poses ≠ achieved poses exactly (Plücker conditioning is a soft constraint; translations are normalized, so absolute scale is undefined) — always re-anchor scale and consider re-estimating poses.
- **Dynamics are not repeatable across runs** — multi-view supervision of a dynamic instant is impossible; design single continuous runs whose camera moves *around* the action.
- **480×832 native offline resolution** (720p is the real-time streaming mode) — plan on splat-space upscaling or `SplatPolish` against upscaled frames.
- **Generated-content artifacts** (texture shimmer, occasional object morphing) become floaters/ghosts in splat space — the existing `MotionMaskFromDepth` + `DepthEdgeFilter` + `PointCloudCleaner` stack mitigates this, and track-validity filtering in `TracksToTrajectories` matters more than with real footage.
- **License**: v2 is CC BY-NC-SA 4.0 — non-commercial only, share-alike. Use v1 (Apache 2.0) for anything with commercial intent.
---
## Sources
Primary: [lingbot-world-v2 repo](https://github.com/Robbyant/lingbot-world-v2) · [v2 tech report arXiv:2607.07534](https://arxiv.org/abs/2607.07534) · [v2 weights (HF)](https://huggingface.co/robbyant/lingbot-world-v2-14b-causal-fast) · [lingbot-world v1 repo](https://github.com/robbyant/lingbot-world) · [v1 paper arXiv:2601.20540](https://arxiv.org/abs/2601.20540) · [v1 cam weights (HF)](https://huggingface.co/robbyant/lingbot-world-base-cam) · code files `generate.py`, `wan/image2video.py`, `wan/utils/cam_utils.py`, `run_fast.sh` (read directly).
Secondary: [Robbyant press release (2026-07-09)](https://www.businesswire.com/news/home/20260708757367/en/Robbyant-Unveils-LingBot-World-2.0-Pioneering-Hour-Long-Real-Time-Generation-in-World-Models) · [v1 release (2026-01-28)](https://www.businesswire.com/news/home/20260128459962/en/Robbyant-Open-Sources-LingBot-World-a-World-Model-for-Millisecond-Level-Real-Time-Interaction) · [CAT4D arXiv:2411.18613](https://arxiv.org/abs/2411.18613) · [ViPE](https://github.com/nv-tlabs/vipe) · ComfyUI support threads [WanVideoWrapper#1920](https://github.com/kijai/ComfyUI-WanVideoWrapper/issues/1920), [ComfyUI#12154](https://github.com/Comfy-Org/ComfyUI/issues/12154).
*Method note: claims were gathered by a fan-out research pass (18 sources, 90 raw claims, 25 adversarially verified: 14 confirmed 3-0, 3 refuted, 8 verification-errored) plus direct reading of both repos' inference code and both arXiv papers. The two load-bearing claims whose automated verification errored (v1's video→point-cloud demonstration; the unreleased 1.3B variant) were re-verified manually against the arXiv HTML.*
+484
View File
@@ -0,0 +1,484 @@
"""Standalone CPU smoke test for the 4D-world node stack (no ComfyUI, no CUDA,
no model downloads).
Run with:
python notebooks/smoke_test_4d.py
Stubs `folder_paths` via sys.modules injection so the repo modules import
outside the ComfyUI runtime, then functionally exercises the NEW code paths
with small synthetic data:
1. interpolate_se3 (pointcloud_nodes, contract C1)
2. render_gaussians (GS_nodes, contract C2) shapes + empty case
3. render_gaussians fast anisotropic footprint
4. GaussianSplats4D.at_time (GS4D_nodes, contract C3)
5. BuildSplats4D kNN track binding
6. SplitSplatsByMask
7. MotionMaskFromDepth
8. align_depth_scale (world_nodes, contract C4) + DepthEdgeFilter
9. FuseSplats weighted voxel fusion
10. SphereSplatSeed pano -> splat sphere -> render round-trip
"""
import math
import os
import sys
import tempfile
import traceback
import types
# --------------------------------------------------------------------------- #
# Environment setup: repo on sys.path + folder_paths stub (before repo imports)
# --------------------------------------------------------------------------- #
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
if REPO_ROOT not in sys.path:
sys.path.insert(0, REPO_ROOT)
_TMP_DIR = tempfile.mkdtemp(prefix="smoke_test_4d_")
def _stub_get_save_image_path(filename_prefix, output_dir, *args, **kwargs):
os.makedirs(output_dir, exist_ok=True)
return output_dir, filename_prefix, 0, "", filename_prefix
_fp_stub = types.ModuleType("folder_paths")
_fp_stub.get_input_directory = lambda: _TMP_DIR
_fp_stub.get_output_directory = lambda: _TMP_DIR
_fp_stub.get_temp_directory = lambda: _TMP_DIR
_fp_stub.get_save_image_path = _stub_get_save_image_path
_fp_stub.get_annotated_filepath = lambda name: os.path.join(_TMP_DIR, name)
_fp_stub.exists_annotated_filepath = lambda name: os.path.exists(os.path.join(_TMP_DIR, name))
_fp_stub.get_filename_list = lambda folder: []
_fp_stub.models_dir = _TMP_DIR
sys.modules["folder_paths"] = _fp_stub
import numpy as np # noqa: E402
import torch # noqa: E402
import GS_nodes # noqa: E402
import GS4D_nodes # noqa: E402
import pointcloud_nodes # noqa: E402
import world_nodes # noqa: E402
GaussianSplats = GS_nodes.GaussianSplats
torch.manual_seed(0)
np.random.seed(0)
# --------------------------------------------------------------------------- #
# Helpers
# --------------------------------------------------------------------------- #
def make_splats(
xyz: torch.Tensor,
sigma: float = 0.05,
color: tuple = None,
opacity_logit: float = 4.0,
) -> GaussianSplats:
"""Isotropic sh_order-0 splats at the given positions."""
n = xyz.shape[0]
if color is None:
rgb = torch.rand(n, 3)
else:
rgb = torch.tensor(color, dtype=torch.float32).view(1, 3).expand(n, 3)
C0 = 0.28209479177387814
return GaussianSplats(
xyz=xyz.float(),
scale=torch.full((n, 3), math.log(sigma)),
rotation=torch.tensor([1.0, 0.0, 0.0, 0.0]).view(1, 4).expand(n, 4).contiguous(),
opacity=torch.full((n, 1), float(opacity_logit)),
f_dc=((rgb - 0.5) / C0).contiguous(),
f_rest=torch.zeros(n, 0),
sh_order=0,
)
def rot_x(deg: float) -> torch.Tensor:
a = math.radians(deg)
return torch.tensor(
[[1, 0, 0], [0, math.cos(a), -math.sin(a)], [0, math.sin(a), math.cos(a)]],
dtype=torch.float32,
)
def rot_y(deg: float) -> torch.Tensor:
a = math.radians(deg)
return torch.tensor(
[[math.cos(a), 0, math.sin(a)], [0, 1, 0], [-math.sin(a), 0, math.cos(a)]],
dtype=torch.float32,
)
def make_pose(R: torch.Tensor, t) -> torch.Tensor:
M = torch.eye(4)
M[:3, :3] = R
M[:3, 3] = torch.tensor(t, dtype=torch.float32)
return M
IDENTITY_4X4 = torch.eye(4)
# --------------------------------------------------------------------------- #
# Tests
# --------------------------------------------------------------------------- #
def test_01_interpolate_se3():
poses = torch.stack(
[
make_pose(torch.eye(3), [0.0, 0.0, 0.0]),
make_pose(rot_y(90.0), [1.0, 2.0, 3.0]),
make_pose(rot_y(90.0) @ rot_x(45.0), [-1.0, 0.0, 2.0]),
]
)
out = pointcloud_nodes.interpolate_se3(poses, 10)
assert out.shape == (10, 4, 4), f"shape {tuple(out.shape)}"
eye = torch.eye(3)
for i in range(10):
R = out[i, :3, :3]
ortho_err = (R @ R.T - eye).abs().max().item()
det = torch.det(R).item()
assert ortho_err < 1e-4, f"step {i}: R@R.T deviates from I by {ortho_err}"
assert abs(det - 1.0) < 1e-4, f"step {i}: det(R)={det}"
assert torch.allclose(out[i, 3], torch.tensor([0.0, 0.0, 0.0, 1.0]), atol=1e-6)
assert (out[0] - poses[0]).abs().max().item() < 1e-4, "start pose mismatch"
assert (out[-1] - poses[-1]).abs().max().item() < 1e-4, "end pose mismatch"
# K == 1 repeats.
rep = pointcloud_nodes.interpolate_se3(poses[:1], 5)
assert rep.shape == (5, 4, 4)
assert (rep - poses[0]).abs().max().item() < 1e-6
def test_02_render_gaussians_shapes_and_empty():
n, H, W = 200, 48, 64
xyz = torch.stack(
[
torch.rand(n) * 2.0 - 1.0,
torch.rand(n) * 2.0 - 1.0,
torch.rand(n) * 3.0 + 2.0,
],
dim=-1,
)
splats = make_splats(xyz, sigma=0.05)
for projection, fov in (("PINHOLE", 90.0), ("EQUIRECTANGULAR", 360.0)):
image, mask, disparity = GS_nodes.render_gaussians(
splats, IDENTITY_4X4, projection, fov, W, H,
render_mode="fast", device="cpu",
)
assert image.shape == (1, H, W, 3), f"{projection} image {tuple(image.shape)}"
assert mask.shape == (H, W), f"{projection} mask {tuple(mask.shape)}"
assert disparity.shape == (1, H, W, 1), f"{projection} disparity {tuple(disparity.shape)}"
assert torch.isfinite(image).all() and torch.isfinite(disparity).all()
assert float(mask.min()) >= 0.0 and float(mask.max()) <= 1.0 + 1e-6
assert float(mask.sum()) > 0.0, f"{projection}: nothing rendered"
# Empty case: every splat strictly behind a pinhole camera (known past bug:
# early return used to yield only 2 outputs).
behind = make_splats(xyz * torch.tensor([1.0, 1.0, -1.0]), sigma=0.05)
result = GS_nodes.render_gaussians(
behind, IDENTITY_4X4, "PINHOLE", 90.0, W, H,
render_mode="fast", device="cpu",
)
assert isinstance(result, tuple) and len(result) == 3, f"empty render returned {len(result)} outputs"
image, mask, disparity = result
assert image.shape == (1, H, W, 3)
assert mask.shape == (H, W)
assert disparity.shape == (1, H, W, 1)
assert float(mask.sum()) == 0.0
def test_03_fast_mode_anisotropy():
H = W = 128
ang = math.radians(45.0) / 2.0
splats = GaussianSplats(
xyz=torch.tensor([[0.0, 0.0, 3.0]]),
scale=torch.log(torch.tensor([[0.5, 0.01, 0.01]])),
rotation=torch.tensor([[math.cos(ang), 0.0, 0.0, math.sin(ang)]]), # 45 deg about +z
opacity=torch.tensor([[6.0]]),
f_dc=torch.zeros(1, 3),
f_rest=torch.zeros(1, 0),
sh_order=0,
)
image, mask, disparity = GS_nodes.render_gaussians(
splats, IDENTITY_4X4, "PINHOLE", 60.0, W, H,
render_mode="fast", max_radius=64, device="cpu",
)
assert float(mask.sum()) > 0.0, "elongated splat rendered nothing"
# Alpha-weighted pixel covariance of the footprint.
ys, xs = torch.meshgrid(
torch.arange(H, dtype=torch.float32), torch.arange(W, dtype=torch.float32),
indexing="ij",
)
w = mask.flatten()
wsum = w.sum()
mx = (w * xs.flatten()).sum() / wsum
my = (w * ys.flatten()).sum() / wsum
dx = xs.flatten() - mx
dy = ys.flatten() - my
cxx = (w * dx * dx).sum() / wsum
cyy = (w * dy * dy).sum() / wsum
cxy = (w * dx * dy).sum() / wsum
cov = torch.tensor([[cxx, cxy], [cxy, cyy]])
evals, evecs = torch.linalg.eigh(cov)
ratio = float(evals[1] / evals[0].clamp(min=1e-8))
assert ratio > 2.0, f"footprint not elongated: eigenvalue ratio {ratio:.2f}"
# Principal axis should be near 45 degrees (rotation honored).
major = evecs[:, 1]
angle = math.degrees(math.atan2(float(major[1]), float(major[0]))) % 180.0
assert abs(angle - 45.0) < 15.0, f"major axis at {angle:.1f} deg, expected ~45"
def test_04_at_time():
T = 5
canonical = make_splats(torch.tensor([[0.0, 0.0, 2.0], [0.0, 1.0, 3.0]]))
static = make_splats(torch.tensor([[5.0, 5.0, 5.0]]))
start = torch.tensor([[0.0, 0.0, 2.0], [0.0, 1.0, 3.0]])
end = torch.tensor([[1.0, 0.0, 2.0], [0.0, -1.0, 3.0]])
ts = torch.linspace(0.0, 1.0, T)
trajectories = torch.stack([start + (end - start) * t for t in ts]) # [5,2,3]
s4d = GS4D_nodes.GaussianSplats4D(
static=static, canonical=canonical, trajectories=trajectories, times=ts,
)
mid = s4d.at_time(0.5)
assert len(mid) == 3, f"count {len(mid)} != dynamic+static (3)"
# Concat order is [static, dynamic].
assert torch.allclose(mid.xyz[0], static.xyz[0], atol=1e-6)
expected_mid = 0.5 * (start + end)
assert torch.allclose(mid.xyz[1:], expected_mid, atol=1e-5), (
f"midpoint mismatch: {mid.xyz[1:]} vs {expected_mid}"
)
lo = s4d.at_time(-1.0)
hi = s4d.at_time(2.0)
assert torch.allclose(lo.xyz[1:], start, atol=1e-5), "t<range should clamp to first step"
assert torch.allclose(hi.xyz[1:], end, atol=1e-5), "t>range should clamp to last step"
def test_05_build_splats4d():
T = 5
ts = torch.linspace(0.0, 1.0, T)
# Two control tracks moving apart along x.
track_a = torch.stack([torch.tensor([-1.0 - 2.0 * t, 0.0, 2.0]) for t in ts])
track_b = torch.stack([torch.tensor([1.0 + 2.0 * t, 0.0, 2.0]) for t in ts])
trajectories3d = torch.stack([track_a, track_b], dim=1) # [T,2,3]
canonical = make_splats(torch.tensor([[-1.05, 0.0, 2.0], [1.05, 0.0, 2.0]]))
node = GS4D_nodes.BuildSplats4D()
(s4d,) = node.build_splats4d(
canonical=canonical,
trajectories3d=trajectories3d,
reference_index=0,
knn=1,
rbf_gamma=0.0,
device="cpu",
)
traj = s4d.trajectories
assert traj.shape == (T, 2, 3), f"trajectories shape {tuple(traj.shape)}"
# Reference timestep: splats stay at their canonical positions.
assert torch.allclose(traj[0], canonical.xyz, atol=1e-5)
# Each splat follows its nearest track's displacement direction.
disp0 = traj[-1, 0] - traj[0, 0]
disp1 = traj[-1, 1] - traj[0, 1]
assert disp0[0] < -1.0, f"splat 0 should move -x with track A, moved {disp0.tolist()}"
assert disp1[0] > 1.0, f"splat 1 should move +x with track B, moved {disp1.tolist()}"
assert torch.allclose(traj[-1, 0], torch.tensor([-3.05, 0.0, 2.0]), atol=1e-4)
assert torch.allclose(traj[-1, 1], torch.tensor([3.05, 0.0, 2.0]), atol=1e-4)
def test_06_split_splats_by_mask():
H = W = 32
mask = torch.zeros(H, W)
mask[:, : W // 2] = 1.0 # left half white
# 10 splats projecting into the left half (x<0), 10 into the right half,
# 5 behind the camera.
jitter = torch.linspace(-0.1, 0.1, 10)
left = torch.stack([torch.full((10,), -0.5) + jitter * 0.1, jitter, torch.full((10,), 2.0)], dim=-1)
right = torch.stack([torch.full((10,), 0.5) + jitter * 0.1, jitter, torch.full((10,), 2.0)], dim=-1)
behind = torch.stack([jitter[:5], jitter[:5], torch.full((5,), -2.0)], dim=-1)
splats = make_splats(torch.cat([left, right, behind], dim=0))
node = GS4D_nodes.SplitSplatsByMask()
inside, outside = node.split_splats(
splats=splats,
mask=mask,
projection="PINHOLE",
horizontal_fov=90.0,
threshold=0.5,
camera_matrix=None,
device="cpu",
)
assert len(inside) == 10, f"inside count {len(inside)} != 10"
assert len(outside) == 15, f"outside count {len(outside)} != 15 (10 right + 5 behind)"
assert (inside.xyz[:, 0] < 0).all(), "inside splats should be the x<0 group"
def test_07_motion_mask_from_depth():
T, H, W = 6, 32, 32
depth = torch.full((T, H, W), 5.0)
r0, r1 = 8, 16
for t in range(T):
depth[t, r0:r1, r0:r1] = 3.0 + 0.4 * t # depth-changing square patch
poses = torch.eye(4).unsqueeze(0).expand(T, 4, 4).contiguous()
node = GS4D_nodes.MotionMaskFromDepth()
(mask,) = node.motion_mask(
depth_seq=depth,
trajectory=poses,
input_projection="PINHOLE",
input_horizontal_fov=90.0,
threshold=0.10,
frame_gap=2,
dilate=0,
device="cpu",
)
assert mask.shape == (T, H, W), f"mask shape {tuple(mask.shape)}"
patch = mask[:, r0:r1, r0:r1]
background = mask.clone()
background[:, r0:r1, r0:r1] = 0.0
patch_mean = float(patch.mean())
bg_sum = float(background.sum())
assert patch_mean > 0.9, f"moving square under-detected: mean {patch_mean:.3f}"
assert bg_sum == 0.0, f"static plane falsely flagged: {bg_sum} pixels"
def test_08_align_depth_scale_and_depth_edge_filter():
H = W = 32
new_depth = torch.rand(H, W) * 9.0 + 1.0
# ref disparity = 0.5 * new disparity + 0.1 (i.e. ref = 2*new before shift).
true_scale, true_shift = 0.5, 0.1
ref_depth = 1.0 / (true_scale / new_depth + true_shift)
valid = torch.ones(H, W)
aligned, scale, shift = world_nodes.align_depth_scale(
new_depth, ref_depth, valid, mode="scale_shift"
)
assert abs(scale - true_scale) / true_scale < 0.05, f"scale {scale} vs {true_scale}"
assert abs(shift - true_shift) / true_shift < 0.05, f"shift {shift} vs {true_shift}"
rel_err = float(((aligned - ref_depth).abs() / ref_depth).max())
assert rel_err < 0.01, f"aligned depth off by {rel_err:.4f} (rel)"
# DepthEdgeFilter: a vertical step edge must be masked out, flat kept.
depth = torch.full((H, W), 1.0)
depth[:, W // 2 :] = 5.0
node = pointcloud_nodes.DepthEdgeFilter()
(valid_mask,) = node.filter_edges(depth, relative_threshold=0.05, dilate=1)
assert valid_mask.shape == (H, W)
edge_cols = valid_mask[:, W // 2 - 1 : W // 2 + 1]
assert float(edge_cols.max()) == 0.0, "step-edge pixels not masked out"
assert float(valid_mask[:, : W // 2 - 3].min()) == 1.0, "flat left region wrongly masked"
assert float(valid_mask[:, W // 2 + 3 :].min()) == 1.0, "flat right region wrongly masked"
def test_09_fuse_splats():
n = 20
voxel = 0.5
base = torch.stack(
[
torch.arange(n, dtype=torch.float32) * voxel + 0.15,
torch.full((n,), 0.15),
torch.full((n,), 0.15),
],
dim=-1,
)
cloud_a = make_splats(base)
cloud_b = make_splats(base + 0.2) # same voxels as A (0.15+0.2 < 0.5)
node = GS_nodes.FuseSplats()
(fused,) = node.fuse_splats(cloud_a, cloud_b, voxel, "smart", 1.0, 1.0, device="cpu")
assert len(fused) < len(cloud_a) + len(cloud_b), (
f"voxel fuse did not reduce: {len(fused)} vs {len(cloud_a) + len(cloud_b)}"
)
assert len(fused) == n, f"expected one splat per voxel ({n}), got {len(fused)}"
# Strong weight_a pulls fused positions onto cloud A.
(fused_w,) = node.fuse_splats(cloud_a, cloud_b, voxel, "average", 1000.0, 1.0, device="cpu")
assert len(fused_w) == n
d_a = torch.cdist(fused_w.xyz, cloud_a.xyz).min(dim=1).values
d_b = torch.cdist(fused_w.xyz, cloud_b.xyz).min(dim=1).values
assert float(d_a.max()) < 0.01, f"fused positions not near cloud A (max dist {float(d_a.max()):.4f})"
assert (d_a < d_b).all(), "weight_a=1000 should pull fused splats toward cloud A"
def test_10_sphere_splat_seed():
H, W = 64, 128
stride = 2
color = (0.2, 0.6, 0.9)
pano = torch.tensor(color).view(1, 1, 1, 3).expand(1, H, W, 3).contiguous()
node = world_nodes.SphereSplatSeed()
(splats,) = node.seed_sphere(
image=pano,
horizontal_fov=360.0,
radius=5.0,
splat_scale_frac=1.5,
stride=stride,
device="cpu",
)
expected = (H // stride) * (W // stride)
assert abs(len(splats) - expected) <= max(4, expected // 20), (
f"splat count {len(splats)} far from expected ~{expected}"
)
image, mask, disparity = GS_nodes.render_gaussians(
splats, IDENTITY_4X4, "PINHOLE", 60.0, 64, 64,
render_mode="fast", device="cpu",
)
assert float(mask.sum()) > 0.0, "pinhole render of the sphere seed is empty"
solid = mask > 0.9
assert bool(solid.any()), "no confidently covered pixels in the render"
rendered = image[0][solid] # [K,3]
target = torch.tensor(color)
err = (rendered.mean(dim=0) - target).abs().max().item()
assert err < 0.05, f"color round-trip failed: rendered mean {rendered.mean(dim=0).tolist()} vs {color}"
# --------------------------------------------------------------------------- #
# Runner
# --------------------------------------------------------------------------- #
TESTS = [
test_01_interpolate_se3,
test_02_render_gaussians_shapes_and_empty,
test_03_fast_mode_anisotropy,
test_04_at_time,
test_05_build_splats4d,
test_06_split_splats_by_mask,
test_07_motion_mask_from_depth,
test_08_align_depth_scale_and_depth_edge_filter,
test_09_fuse_splats,
test_10_sphere_splat_seed,
]
def main() -> int:
passed = 0
failed = []
for test in TESTS:
name = test.__name__
try:
test()
except Exception:
failed.append(name)
print(f"[FAIL] {name}")
traceback.print_exc()
else:
passed += 1
print(f"[ ok ] {name}")
print(f"\n{passed}/{len(TESTS)} tests passed")
if failed:
print("Failed:", ", ".join(failed))
return 1
return 0
if __name__ == "__main__":
sys.exit(main())
File diff suppressed because one or more lines are too long
+214 -6
View File
@@ -7,7 +7,10 @@ import os
import folder_paths
import logging
import hashlib
from kornia.filters import median_blur
try:
from kornia.filters import median_blur
except ImportError: # kornia is optional; median_blur is not used in this module
median_blur = None
from tqdm import tqdm
# Try importing open3d and its visualization modules; log a warning if not found
@@ -136,6 +139,116 @@ def project_first_hit(volume_sparse: torch.Tensor) -> Tuple[torch.Tensor, torch.
return rgba.permute(2, 0, 1), first_hit.any(dim=2)
# ==== SE(3) trajectory interpolation ==== #
def _rotmat_to_quat_wxyz(R: torch.Tensor) -> torch.Tensor:
"""
Convert a batch of rotation matrices [K,3,3] to unit quaternions [K,4] (wxyz).
Uses Shepperd's method for numerical robustness. K is expected to be small
(trajectory waypoints), so a Python loop is acceptable.
"""
quats = []
for i in range(R.shape[0]):
m = R[i]
trace = m[0, 0] + m[1, 1] + m[2, 2]
if trace > 0.0:
s = torch.sqrt(trace + 1.0) * 2.0
w = 0.25 * s
x = (m[2, 1] - m[1, 2]) / s
y = (m[0, 2] - m[2, 0]) / s
z = (m[1, 0] - m[0, 1]) / s
elif m[0, 0] > m[1, 1] and m[0, 0] > m[2, 2]:
s = torch.sqrt(1.0 + m[0, 0] - m[1, 1] - m[2, 2]) * 2.0
w = (m[2, 1] - m[1, 2]) / s
x = 0.25 * s
y = (m[0, 1] + m[1, 0]) / s
z = (m[0, 2] + m[2, 0]) / s
elif m[1, 1] > m[2, 2]:
s = torch.sqrt(1.0 + m[1, 1] - m[0, 0] - m[2, 2]) * 2.0
w = (m[0, 2] - m[2, 0]) / s
x = (m[0, 1] + m[1, 0]) / s
y = 0.25 * s
z = (m[1, 2] + m[2, 1]) / s
else:
s = torch.sqrt(1.0 + m[2, 2] - m[0, 0] - m[1, 1]) * 2.0
w = (m[1, 0] - m[0, 1]) / s
x = (m[0, 2] + m[2, 0]) / s
y = (m[1, 2] + m[2, 1]) / s
z = 0.25 * s
quats.append(torch.stack([w, x, y, z]))
q = torch.stack(quats, dim=0)
return q / q.norm(dim=-1, keepdim=True).clamp(min=1e-12)
def _quat_wxyz_to_rotmat(q: torch.Tensor) -> torch.Tensor:
"""Convert unit quaternions [N,4] (wxyz) to rotation matrices [N,3,3]."""
q = q / q.norm(dim=-1, keepdim=True).clamp(min=1e-12)
w, x, y, z = q.unbind(-1)
R = torch.stack([
1 - 2 * (y * y + z * z), 2 * (x * y - w * z), 2 * (x * z + w * y),
2 * (x * y + w * z), 1 - 2 * (x * x + z * z), 2 * (y * z - w * x),
2 * (x * z - w * y), 2 * (y * z + w * x), 1 - 2 * (x * x + y * y),
], dim=-1).reshape(*q.shape[:-1], 3, 3)
return R
def _quat_slerp(q0: torch.Tensor, q1: torch.Tensor, alpha: torch.Tensor) -> torch.Tensor:
"""
Spherical linear interpolation between quaternion batches q0, q1 [N,4] (wxyz)
with per-element interpolation factors alpha [N]. Falls back to normalized
lerp when the quaternions are nearly parallel.
"""
dot = (q0 * q1).sum(dim=-1, keepdim=True)
q1 = torch.where(dot < 0.0, -q1, q1) # shortest arc
dot = dot.abs().clamp(max=1.0)
a = alpha.reshape(-1, 1).to(q0.dtype)
theta = torch.acos(dot)
sin_theta = torch.sin(theta)
near_parallel = sin_theta < 1e-6
denom = sin_theta.clamp(min=1e-12)
w0 = torch.where(near_parallel, 1.0 - a, torch.sin((1.0 - a) * theta) / denom)
w1 = torch.where(near_parallel, a, torch.sin(a * theta) / denom)
q = w0 * q0 + w1 * q1
return q / q.norm(dim=-1, keepdim=True).clamp(min=1e-12)
def interpolate_se3(trajectory: torch.Tensor, num_steps: int) -> torch.Tensor:
"""trajectory [K,4,4] -> [num_steps,4,4]. Piecewise: quaternion SLERP on R, lerp on t.
K==1 -> repeat. Must return valid rotation matrices (orthonormal)."""
if isinstance(trajectory, np.ndarray):
trajectory = torch.from_numpy(trajectory)
trajectory = trajectory.float()
if trajectory.dim() == 2:
trajectory = trajectory.unsqueeze(0)
if trajectory.dim() != 3 or trajectory.shape[-2:] != (4, 4):
raise ValueError(f"interpolate_se3 expects trajectory of shape [K,4,4], got {tuple(trajectory.shape)}")
if num_steps < 1:
raise ValueError(f"interpolate_se3 requires num_steps >= 1, got {num_steps}")
K = trajectory.shape[0]
if K == 1:
return trajectory.expand(num_steps, 4, 4).clone()
R = trajectory[:, :3, :3]
t = trajectory[:, :3, 3]
q = _rotmat_to_quat_wxyz(R)
# Enforce hemisphere continuity along the waypoint sequence so piecewise
# SLERP always takes the shortest arc between consecutive poses.
for k in range(1, K):
if (q[k] * q[k - 1]).sum() < 0.0:
q[k] = -q[k]
idxs = torch.linspace(0, K - 1, num_steps, device=trajectory.device)
lower = idxs.floor().long().clamp(max=K - 2)
upper = lower + 1
alpha = (idxs - lower.float())
q_interp = _quat_slerp(q[lower], q[upper], alpha)
t_interp = t[lower] * (1.0 - alpha).unsqueeze(-1) + t[upper] * alpha.unsqueeze(-1)
out = torch.eye(4, dtype=trajectory.dtype, device=trajectory.device).repeat(num_steps, 1, 1)
out[:, :3, :3] = _quat_wxyz_to_rotmat(q_interp)
out[:, :3, 3] = t_interp
return out
# ==== Node Definitions ==== #
class DepthToPointCloud:
"""
@@ -794,8 +907,9 @@ class CameraMotionNode:
class CameraInterpolationNode:
"""
Wrap two 4×4 poses into a trajectory tensor.
Outputs only `trajectory` (shape 2×4×4).
Interpolate between two 4×4 poses into a trajectory tensor using proper
SE(3) interpolation (quaternion SLERP on rotation, lerp on translation).
Outputs `trajectory` (shape num_steps×4×4, default 2×4×4).
"""
@classmethod
@@ -804,7 +918,10 @@ class CameraInterpolationNode:
"required": {
"initial_matrix": ("MAT_4X4",),
"final_matrix": ("MAT_4X4",),
}
},
"optional": {
"num_steps": ("INT", {"default": 2, "min": 2, "max": 4096, "tooltip": "Number of poses in the output trajectory, SE(3)-interpolated between the two matrices."}),
},
}
RETURN_TYPES = ("TENSOR",)
RETURN_NAMES = ("trajectory",)
@@ -815,14 +932,15 @@ class CameraInterpolationNode:
self,
initial_matrix: torch.Tensor,
final_matrix: torch.Tensor,
num_steps: int = 2,
) -> Tuple[torch.Tensor]:
# stack into a (2,4,4) trajectory
# convert to tensor if needed
if isinstance(initial_matrix, np.ndarray):
initial_matrix = torch.from_numpy(initial_matrix).float()
if isinstance(final_matrix, np.ndarray):
final_matrix = torch.from_numpy(final_matrix).float()
traj = torch.stack([initial_matrix, final_matrix], dim=0)
keyframes = torch.stack([initial_matrix.float(), final_matrix.float()], dim=0)
traj = interpolate_se3(keyframes, num_steps)
return (traj,)
@@ -1250,6 +1368,95 @@ class LoadTrajectory:
return f"Invalid trajectory file: {trajectory_file}"
return True
class DepthEdgeFilter:
"""
Detect "flying pixel" depth discontinuities and output a validity mask.
A pixel is flagged as an edge where |depth gradient| / depth exceeds
`relative_threshold`; edges are optionally dilated. Returns a MASK with
1.0 where the depth is valid (NOT a flying-pixel edge) and 0.0 on edges.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
# Depth: [H,W] or [T,H,W], trailing channel dim of 1 accepted
"depth": ("TENSOR", {"shape_hint": [None, None, None]}),
"relative_threshold": ("FLOAT", {"default": 0.05, "min": 0.0, "max": 10.0, "step": 0.005, "tooltip": "Mark a pixel as edge where |depth gradient| / depth exceeds this value."}),
"dilate": ("INT", {"default": 1, "min": 0, "max": 64, "tooltip": "Grow detected edges by this many pixels (max-pool dilation)."}),
},
"optional": {
"mask": ("MASK", {"tooltip": "Optional validity mask ANDed with the edge-filter result."}),
},
}
RETURN_TYPES = ("MASK",)
RETURN_NAMES = ("valid_mask",)
FUNCTION = "filter_edges"
CATEGORY = "Camera/PointCloud"
def filter_edges(
self,
depth: torch.Tensor,
relative_threshold: float,
dilate: int,
mask: torch.Tensor = None,
) -> Tuple[torch.Tensor]:
d = depth
if isinstance(d, np.ndarray):
d = torch.from_numpy(d)
d = d.float()
# Accept [H,W], [H,W,1], [T,H,W], [T,H,W,1]
if d.dim() == 4 and d.shape[-1] == 1:
d = d[..., 0]
elif d.dim() == 3 and d.shape[-1] == 1:
d = d[..., 0]
squeeze_batch = False
if d.dim() == 2:
d = d.unsqueeze(0)
squeeze_batch = True
if d.dim() != 3:
raise ValueError(f"DepthEdgeFilter expects depth of shape [H,W] or [T,H,W] (trailing 1 ok), got {tuple(depth.shape)}")
eps = 1e-8
# Forward differences along x and y; propagate each difference to both
# neighbouring pixels so both sides of a discontinuity are flagged.
dx = (d[:, :, 1:] - d[:, :, :-1]).abs()
dy = (d[:, 1:, :] - d[:, :-1, :]).abs()
gx = torch.zeros_like(d)
gx[:, :, :-1] = dx
gx[:, :, 1:] = torch.maximum(gx[:, :, 1:], dx)
gy = torch.zeros_like(d)
gy[:, :-1, :] = dy
gy[:, 1:, :] = torch.maximum(gy[:, 1:, :], dy)
grad = torch.maximum(gx, gy)
edge = (grad / d.abs().clamp(min=eps)) > relative_threshold
if dilate > 0:
k = 2 * int(dilate) + 1
edge = F.max_pool2d(edge.float().unsqueeze(1), kernel_size=k, stride=1, padding=int(dilate)).squeeze(1) > 0.5
valid = (~edge).float()
if mask is not None:
m = mask
if isinstance(m, np.ndarray):
m = torch.from_numpy(m)
m = m.float().to(valid.device)
if m.dim() == 4 and m.shape[-1] == 1:
m = m[..., 0]
if m.dim() == 2:
m = m.unsqueeze(0)
if m.shape[0] == 1 and valid.shape[0] > 1:
m = m.expand(valid.shape[0], -1, -1)
if m.shape[-2:] != valid.shape[-2:]:
m = F.interpolate(m.unsqueeze(1), size=valid.shape[-2:], mode="nearest").squeeze(1)
valid = valid * (m > 0.5).float()
if squeeze_batch:
valid = valid[0]
return (valid,)
NODE_CLASS_MAPPINGS = {
"DepthToPointCloud": DepthToPointCloud,
"TransformPointCloud": TransformPointCloud,
@@ -1264,4 +1471,5 @@ NODE_CLASS_MAPPINGS = {
"PointCloudCleaner": PointCloudCleaner,
"SaveTrajectory": SaveTrajectory,
"LoadTrajectory": LoadTrajectory,
"DepthEdgeFilter": DepthEdgeFilter,
}
+461
View File
@@ -0,0 +1,461 @@
"""Camera pose estimation nodes.
Provides:
- VideoPoseEstimator: VGGT-based per-frame camera pose + depth + intrinsics
estimation from a video clip.
- TrajectoryInvert / TrajectoryCompose: small utility nodes for wiring
trajectory tensors ([K, 4, 4] world-to-camera matrices) in graphs.
Coordinate convention (matches the rest of this repo): camera frame is
+X right, +Y down, +Z forward; trajectory matrices are 4x4 world-to-camera
(`cam_pts = world_pts @ R.T + t`). VGGT outputs OpenCV-convention
camera-from-world extrinsics, which match this convention directly.
"""
import math
import os
import sys
from typing import Any, Dict, Optional, Tuple
import numpy as np
import torch
import torch.nn.functional as F
from tqdm import tqdm
try:
import folder_paths
except ImportError: # Allow notebook usage outside ComfyUI
class _FolderPathsStub:
def __getattr__(self, name):
raise ModuleNotFoundError(
"folder_paths is unavailable; this feature requires the ComfyUI runtime."
)
folder_paths = _FolderPathsStub()
_here = os.path.dirname(os.path.abspath(__file__))
# climb up 2 levels: camera-comfyUI -> custom_nodes -> ComfyUI
COMFYUI_ROOT = os.path.abspath(os.path.join(_here, os.pardir, os.pardir))
DEVICE_CHOICES = ["auto", "cpu", "cuda"]
# Module-level model cache: {device_str: model}
_VGGT_MODEL_CACHE: Dict[str, Any] = {}
# --------------------------------------------------------------------------- #
# SE(3) interpolation (contract C1). Prefer the shared implementation from
# pointcloud_nodes; fall back to a local copy so this file works standalone.
# --------------------------------------------------------------------------- #
def _matrix_to_quaternion(R: torch.Tensor) -> torch.Tensor:
"""Convert a single 3x3 rotation matrix to a wxyz quaternion."""
R = R.to(torch.float64)
m00, m01, m02 = R[0, 0], R[0, 1], R[0, 2]
m10, m11, m12 = R[1, 0], R[1, 1], R[1, 2]
m20, m21, m22 = R[2, 0], R[2, 1], R[2, 2]
trace = m00 + m11 + m22
if trace > 0.0:
s = torch.sqrt(trace + 1.0) * 2.0
w = 0.25 * s
x = (m21 - m12) / s
y = (m02 - m20) / s
z = (m10 - m01) / s
elif (m00 > m11) and (m00 > m22):
s = torch.sqrt(1.0 + m00 - m11 - m22) * 2.0
w = (m21 - m12) / s
x = 0.25 * s
y = (m01 + m10) / s
z = (m02 + m20) / s
elif m11 > m22:
s = torch.sqrt(1.0 + m11 - m00 - m22) * 2.0
w = (m02 - m20) / s
x = (m01 + m10) / s
y = 0.25 * s
z = (m12 + m21) / s
else:
s = torch.sqrt(1.0 + m22 - m00 - m11) * 2.0
w = (m10 - m01) / s
x = (m02 + m20) / s
y = (m12 + m21) / s
z = 0.25 * s
q = torch.stack([w, x, y, z])
return (q / q.norm().clamp(min=1e-12)).to(torch.float32)
def _quaternion_to_matrix(q: torch.Tensor) -> torch.Tensor:
"""Convert a wxyz quaternion to a 3x3 rotation matrix."""
q = q / q.norm().clamp(min=1e-12)
w, x, y, z = q[0], q[1], q[2], q[3]
return torch.stack([
torch.stack([1 - 2 * (y * y + z * z), 2 * (x * y - w * z), 2 * (x * z + w * y)]),
torch.stack([2 * (x * y + w * z), 1 - 2 * (x * x + z * z), 2 * (y * z - w * x)]),
torch.stack([2 * (x * z - w * y), 2 * (y * z + w * x), 1 - 2 * (x * x + y * y)]),
])
def _quat_slerp(q0: torch.Tensor, q1: torch.Tensor, alpha: float) -> torch.Tensor:
"""Spherical linear interpolation between two wxyz quaternions."""
q0 = q0 / q0.norm().clamp(min=1e-12)
q1 = q1 / q1.norm().clamp(min=1e-12)
dot = torch.dot(q0, q1)
if dot < 0.0: # take the short path on the quaternion hypersphere
q1 = -q1
dot = -dot
dot = dot.clamp(-1.0, 1.0)
if dot > 0.9995: # nearly parallel: lerp + renormalize is numerically safer
q = (1.0 - alpha) * q0 + alpha * q1
return q / q.norm().clamp(min=1e-12)
theta = torch.acos(dot)
sin_theta = torch.sin(theta)
w0 = torch.sin((1.0 - alpha) * theta) / sin_theta
w1 = torch.sin(alpha * theta) / sin_theta
q = w0 * q0 + w1 * q1
return q / q.norm().clamp(min=1e-12)
def _interpolate_se3_fallback(trajectory: torch.Tensor, num_steps: int) -> torch.Tensor:
"""trajectory [K,4,4] -> [num_steps,4,4]. Piecewise: quaternion SLERP on R, lerp on t.
K==1 -> repeat. Returns valid (orthonormal) rotation matrices. Matches contract C1."""
traj = torch.as_tensor(trajectory, dtype=torch.float32)
if traj.dim() == 2:
traj = traj.unsqueeze(0)
if traj.dim() != 3 or traj.shape[-2:] != (4, 4):
raise ValueError(f"Expected trajectory of shape [K,4,4], got {tuple(traj.shape)}")
K = traj.shape[0]
if K == 1:
return traj.expand(num_steps, 4, 4).clone()
quats = torch.stack([_matrix_to_quaternion(traj[i, :3, :3]) for i in range(K)])
trans = traj[:, :3, 3]
positions = torch.linspace(0.0, float(K - 1), num_steps)
out = []
for pos in positions:
lower = int(torch.floor(pos).clamp(max=K - 2))
upper = lower + 1
alpha = float(pos) - lower
q = _quat_slerp(quats[lower], quats[upper], alpha)
t = (1.0 - alpha) * trans[lower] + alpha * trans[upper]
M = torch.eye(4, dtype=torch.float32)
M[:3, :3] = _quaternion_to_matrix(q)
M[:3, 3] = t
out.append(M)
return torch.stack(out, dim=0)
try:
from .pointcloud_nodes import interpolate_se3
except Exception:
try:
from pointcloud_nodes import interpolate_se3
except Exception:
interpolate_se3 = _interpolate_se3_fallback
# --------------------------------------------------------------------------- #
# VGGT lazy import helpers
# --------------------------------------------------------------------------- #
def _import_vggt() -> Tuple[Any, Any]:
"""Lazily import VGGT. Tries the pip package first, then a sibling clone
at COMFYUI_ROOT/vggt (mirroring how video_nodes.py handles Video-Depth-Anything)."""
try:
from vggt.models.vggt import VGGT
from vggt.utils.pose_enc import pose_encoding_to_extri_intri
return VGGT, pose_encoding_to_extri_intri
except ImportError:
pass
vggt_clone_path = os.path.join(COMFYUI_ROOT, "vggt")
if os.path.isdir(vggt_clone_path) and vggt_clone_path not in sys.path:
sys.path.insert(0, vggt_clone_path)
try:
from vggt.models.vggt import VGGT
from vggt.utils.pose_enc import pose_encoding_to_extri_intri
return VGGT, pose_encoding_to_extri_intri
except ImportError as exc:
raise ModuleNotFoundError(
"VGGT is not installed. Install it with `pip install vggt` (or "
"`pip install git+https://github.com/facebookresearch/vggt.git`), or clone "
f"https://github.com/facebookresearch/vggt into {vggt_clone_path!r}. "
"It also requires `huggingface_hub` to download the facebook/VGGT-1B weights."
) from exc
def _get_vggt_model(device: torch.device) -> Any:
"""Load (and cache) the VGGT-1B model on the requested device."""
key = str(device)
if key not in _VGGT_MODEL_CACHE:
VGGT, _ = _import_vggt()
print(f"[pose_nodes] Loading facebook/VGGT-1B onto {key} (first call downloads ~5GB weights)...")
model = VGGT.from_pretrained("facebook/VGGT-1B")
model = model.to(device).eval()
_VGGT_MODEL_CACHE[key] = model
return _VGGT_MODEL_CACHE[key]
def _vggt_preprocess(frames: torch.Tensor, resolution: int, device: torch.device) -> torch.Tensor:
"""[T,H,W,3] float 0..1 -> [1,T,3,Hp,Wp] with max dim == resolution (both dims
divisible by 14, the VGGT patch size), aspect ratio preserved."""
T, H, W, _ = frames.shape
imgs = frames.permute(0, 3, 1, 2).to(device=device, dtype=torch.float32)
if imgs.max() > 1.5: # defensively handle 0..255 inputs
imgs = imgs / 255.0
scale = float(resolution) / float(max(H, W))
new_h = max(14, int(round(H * scale / 14.0)) * 14)
new_w = max(14, int(round(W * scale / 14.0)) * 14)
if (new_h, new_w) != (H, W):
imgs = F.interpolate(imgs, size=(new_h, new_w), mode="bilinear", align_corners=False)
return imgs.clamp(0.0, 1.0).unsqueeze(0) # [1,T,3,Hp,Wp]
class VideoPoseEstimator:
"""
Estimates per-frame camera poses (world-to-camera [T,4,4]), metric-ish depth
maps, depth confidence and the horizontal FOV from a video clip using
facebook/VGGT-1B.
VGGT extrinsics use the OpenCV camera convention (+X right, +Y down,
+Z forward, camera-from-world), which matches this repo's trajectory
convention, so the matrices are returned as-is (padded to 4x4).
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
# Video frames: Tensor [T, H, W, 3] float 0..1
"frames": ("IMAGE", {"shape_hint": [None, None, None, 3]}),
"max_frames": ("INT", {
"default": 64, "min": 1, "max": 1024,
"tooltip": "If the clip has more frames than this, it is stride-subsampled "
"for VGGT and the poses are SE(3)-interpolated back to full length "
"(depth/confidence use nearest-frame fill).",
}),
"resolution": ("INT", {
"default": 518, "min": 98, "max": 1036,
"tooltip": "Max image dimension fed to VGGT (rounded to a multiple of 14).",
}),
"device": (DEVICE_CHOICES, {"default": "auto"}),
}
}
RETURN_TYPES = ("TENSOR", "TENSOR", "FLOAT", "TENSOR")
RETURN_NAMES = ("trajectory", "depths", "horizontal_fov", "confidence")
FUNCTION = "estimate_poses"
CATEGORY = "Camera/Pose"
DESCRIPTION = (
"VGGT camera pose + depth estimation. Outputs world-to-camera trajectory [T,4,4], "
"depth maps [T,H,W] at the input resolution, mean horizontal FOV (degrees) and "
"per-pixel depth confidence [T,H,W]."
)
def estimate_poses(
self,
frames: torch.Tensor,
max_frames: int = 64,
resolution: int = 518,
device: str = "auto",
) -> Tuple[torch.Tensor, torch.Tensor, float, torch.Tensor]:
if frames.dim() != 4 or frames.shape[-1] != 3:
raise ValueError(f"Expected frames of shape [T,H,W,3], got {tuple(frames.shape)}")
T_full, H, W, _ = frames.shape
if device == "auto":
dev = torch.device("cuda" if torch.cuda.is_available() else "cpu")
elif device == "cuda":
if not torch.cuda.is_available():
raise ValueError("CUDA requested but not available.")
dev = torch.device("cuda")
else:
dev = torch.device("cpu")
# Stride-subsample overly long clips, keeping the frame mapping so that
# poses can be interpolated back afterwards.
if T_full > max_frames:
sub_indices = torch.linspace(0, T_full - 1, max_frames).round().long().unique()
print(
f"[VideoPoseEstimator] WARNING: clip has {T_full} frames > max_frames={max_frames}; "
f"running VGGT on {sub_indices.numel()} stride-subsampled frames. Poses are "
"SE(3)-interpolated back to full length; depth/confidence use nearest-frame fill. "
"Increase max_frames for exact per-frame estimates."
)
proc_frames = frames[sub_indices]
else:
sub_indices = None
proc_frames = frames
images = _vggt_preprocess(proc_frames, resolution, dev) # [1,S,3,Hp,Wp]
S, Hp, Wp = images.shape[1], images.shape[-2], images.shape[-1]
_, pose_encoding_to_extri_intri = _import_vggt()
model = _get_vggt_model(dev)
try:
with torch.no_grad():
if dev.type == "cuda":
capability = torch.cuda.get_device_capability(dev)
amp_dtype = torch.bfloat16 if capability[0] >= 8 else torch.float16
with torch.autocast(device_type="cuda", dtype=amp_dtype):
aggregated_tokens_list, ps_idx = model.aggregator(images)
else:
aggregated_tokens_list, ps_idx = model.aggregator(images)
# Camera + depth heads run in full precision (per the official VGGT example).
pose_enc = model.camera_head(aggregated_tokens_list)[-1]
extrinsic, intrinsic = pose_encoding_to_extri_intri(pose_enc, images.shape[-2:])
depth_map, depth_conf = model.depth_head(aggregated_tokens_list, images, ps_idx)
except torch.cuda.OutOfMemoryError as exc:
raise RuntimeError(
f"VGGT ran out of GPU memory on {S} frames at {Wp}x{Hp}. "
"Lower max_frames and/or resolution, or set device='cpu' (slow)."
) from exc
# ---- Trajectory: pad OpenCV world-to-camera [S,3,4] to [S,4,4] ---- #
extrinsic = extrinsic.squeeze(0).to(torch.float32).cpu() # [S,3,4]
trajectory = torch.eye(4, dtype=torch.float32).unsqueeze(0).repeat(extrinsic.shape[0], 1, 1)
trajectory[:, :3, :4] = extrinsic
# ---- Horizontal FOV from intrinsics (resolution-invariant fx/W ratio) ---- #
intrinsic = intrinsic.squeeze(0).to(torch.float32).cpu() # [S,3,3]
fx = intrinsic[:, 0, 0].clamp(min=1e-6)
hfov_per_frame = 2.0 * torch.atan(0.5 * float(Wp) / fx) # radians, at processing width
# Aspect ratio is preserved during preprocessing, so fx/W is the same at
# the original width and the FOV needs no conversion.
horizontal_fov = float(torch.rad2deg(hfov_per_frame).mean())
# ---- Depth + confidence, resized back to the input resolution ---- #
depth = depth_map.squeeze(0).to(torch.float32).cpu() # [S,Hp,Wp,1] (or [S,Hp,Wp])
if depth.dim() == 4 and depth.shape[-1] == 1:
depth = depth.squeeze(-1)
conf = depth_conf.squeeze(0).to(torch.float32).cpu() # [S,Hp,Wp]
if conf.dim() == 4 and conf.shape[-1] == 1:
conf = conf.squeeze(-1)
# ---- Convert VGGT z-depth to RADIAL ray depth ---- #
# VGGT's depth head predicts z-depth (its unprojection is
# x = (u - cx) * d / fx, z = d), while every consumer in this repo
# (pointcloud *_depth_to_XYZ helpers, MotionMaskFromDepth,
# TracksToTrajectories, the GS4D helpers) multiplies unit ray directions
# by depth, i.e. expects RADIAL distance. Multiply by the per-pixel ray
# norm sqrt(1 + ((u-cx)/fx)^2 + ((v-cy)/fy)^2) using the per-frame
# intrinsics at the VGGT processing resolution.
fx_pf = intrinsic[:, 0, 0].clamp(min=1e-6).view(-1, 1, 1) # [S,1,1]
fy_pf = intrinsic[:, 1, 1].clamp(min=1e-6).view(-1, 1, 1)
cx_pf = intrinsic[:, 0, 2].view(-1, 1, 1)
cy_pf = intrinsic[:, 1, 2].view(-1, 1, 1)
uu = torch.arange(Wp, dtype=torch.float32).view(1, 1, -1)
vv = torch.arange(Hp, dtype=torch.float32).view(1, -1, 1)
xn = (uu - cx_pf) / fx_pf
yn = (vv - cy_pf) / fy_pf
depth = depth * torch.sqrt(1.0 + xn * xn + yn * yn)
if (Hp, Wp) != (H, W):
depth = F.interpolate(depth.unsqueeze(1), size=(H, W), mode="bilinear", align_corners=False).squeeze(1)
conf = F.interpolate(conf.unsqueeze(1), size=(H, W), mode="bilinear", align_corners=False).squeeze(1)
# ---- If subsampled, expand back to the full frame count ---- #
if sub_indices is not None:
# Subsample indices are (near-)uniform over [0, T_full-1], so uniform
# SE(3) resampling reconstructs per-frame poses well.
trajectory = interpolate_se3(trajectory, T_full)
all_t = torch.arange(T_full).unsqueeze(1) # [T_full,1]
nearest = (sub_indices.unsqueeze(0) - all_t).abs().argmin(dim=1) # [T_full]
depth = depth[nearest]
conf = conf[nearest]
return (trajectory, depth, horizontal_fov, conf)
class TrajectoryInvert:
"""
Inverts each 4x4 matrix in a trajectory tensor, converting between
world-to-camera and camera-to-world conventions.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
# Trajectory: Tensor [K, 4, 4] (a single [4, 4] matrix also works)
"trajectory": ("TENSOR", {"shape_hint": [None, 4, 4]}),
}
}
RETURN_TYPES = ("TENSOR",)
RETURN_NAMES = ("trajectory",)
FUNCTION = "invert"
CATEGORY = "Camera/Pose"
DESCRIPTION = "Inverts each 4x4 pose (world-to-camera <-> camera-to-world)."
def invert(self, trajectory: torch.Tensor) -> Tuple[torch.Tensor]:
traj = torch.as_tensor(trajectory, dtype=torch.float32)
squeeze = traj.dim() == 2
if squeeze:
traj = traj.unsqueeze(0)
if traj.dim() != 3 or traj.shape[-2:] != (4, 4):
raise ValueError(f"Expected trajectory of shape [K,4,4], got {tuple(trajectory.shape)}")
# Rigid-body inverse: R -> R.T, t -> -R.T @ t (numerically stabler than
# a generic matrix inverse for SE(3) poses).
R = traj[:, :3, :3]
t = traj[:, :3, 3:4]
Rt = R.transpose(1, 2)
inv = torch.eye(4, dtype=traj.dtype).unsqueeze(0).repeat(traj.shape[0], 1, 1)
inv[:, :3, :3] = Rt
inv[:, :3, 3:4] = -Rt @ t
if squeeze:
inv = inv.squeeze(0)
return (inv,)
class TrajectoryCompose:
"""
Composes two trajectories per frame: out_k = A_k @ B_k. Either input may be
a single [4,4] matrix, which is broadcast against the other. Useful for
retargeting novel camera paths relative to a source pose (e.g. compose a
relative path with the inverse of source pose 0).
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
# Left operand: Tensor [K, 4, 4] or [4, 4]
"trajectory_a": ("TENSOR", {"shape_hint": [None, 4, 4]}),
# Right operand: Tensor [K, 4, 4] or [4, 4]
"trajectory_b": ("TENSOR", {"shape_hint": [None, 4, 4]}),
}
}
RETURN_TYPES = ("TENSOR",)
RETURN_NAMES = ("trajectory",)
FUNCTION = "compose"
CATEGORY = "Camera/Pose"
DESCRIPTION = "Per-frame matrix product A @ B; a single 4x4 input broadcasts over the other."
def compose(self, trajectory_a: torch.Tensor, trajectory_b: torch.Tensor) -> Tuple[torch.Tensor]:
A = torch.as_tensor(trajectory_a, dtype=torch.float32)
B = torch.as_tensor(trajectory_b, dtype=torch.float32)
both_single = A.dim() == 2 and B.dim() == 2
if A.dim() == 2:
A = A.unsqueeze(0)
if B.dim() == 2:
B = B.unsqueeze(0)
if A.dim() != 3 or A.shape[-2:] != (4, 4):
raise ValueError(f"Expected trajectory_a of shape [K,4,4] or [4,4], got {tuple(trajectory_a.shape)}")
if B.dim() != 3 or B.shape[-2:] != (4, 4):
raise ValueError(f"Expected trajectory_b of shape [K,4,4] or [4,4], got {tuple(trajectory_b.shape)}")
if A.shape[0] != B.shape[0] and A.shape[0] != 1 and B.shape[0] != 1:
raise ValueError(
f"Trajectory lengths do not broadcast: {A.shape[0]} vs {B.shape[0]} "
"(they must match, or one must be a single 4x4 matrix)."
)
out = torch.matmul(A, B) # broadcasts [1,4,4] against [K,4,4]
if both_single:
out = out.squeeze(0)
return (out,)
NODE_CLASS_MAPPINGS = {
"VideoPoseEstimator": VideoPoseEstimator,
"TrajectoryInvert": TrajectoryInvert,
"TrajectoryCompose": TrajectoryCompose,
}
+23
View File
@@ -0,0 +1,23 @@
[project]
name = "camera-comfyui"
description = "Custom ComfyUI nodes for camera projections (pinhole/fisheye/equirectangular), depth, point clouds, camera trajectories, and 3D/4D Gaussian splatting — including video-to-4D-world workflows."
version = "1.0.0"
license = { file = "LICENSE" }
dependencies = [
"transformers==4.50.0",
"diffusers==0.33.1",
"open3d==0.19.0",
"protobuf",
]
[project.urls]
Repository = "https://github.com/Alexankharin/camera-comfyUI"
[tool.comfy]
PublisherId = "alexk"
DisplayName = "camera-comfyUI"
# Force-include the SHARP submodule: its files are a gitlink in the parent repo
# (not git-tracked files), so without this the registry archive would ship
# without submodules/ml-sharpt and ImageToSplat/VideoToFusedSplats would be
# unavailable until users clone it manually.
includes = ["submodules/ml-sharpt/"]
+30 -23
View File
@@ -7,7 +7,7 @@ from typing import Dict, Any, Tuple
from tqdm import tqdm # Added tqdm import
# Import existing pointcloud nodes and projection definitions
from .pointcloud_nodes import DepthToPointCloud, TransformPointCloud, ProjectPointCloud, Projection, PointCloudCleaner
from .pointcloud_nodes import DepthToPointCloud, TransformPointCloud, ProjectPointCloud, Projection, PointCloudCleaner, interpolate_se3
import folder_paths
# Ensure video_depth_anything is on path
@@ -95,18 +95,8 @@ class VideoCameraMotionSequence:
# depth_seq: [T, H, W] or [T, H, W, 1]
T, H, W, _ = frames.shape
# Interpolate trajectory to match T
K = trajectory.shape[0]
if K < 2:
interp_traj = trajectory.expand(T, 4, 4).clone()
else:
idxs = torch.linspace(0, K - 1, T, device=trajectory.device)
lower = idxs.floor().long().clamp(max=K - 2)
upper = lower + 1
alpha = (idxs - lower.float()).unsqueeze(-1).unsqueeze(-1)
traj_lower = trajectory[lower]
traj_upper = trajectory[upper]
interp_traj = traj_lower * (1 - alpha) + traj_upper * alpha
# Interpolate trajectory to match T (SE(3): quaternion SLERP on R, lerp on t)
interp_traj = interpolate_se3(trajectory, T)
out_frames = []
out_masks = []
@@ -122,12 +112,12 @@ class VideoCameraMotionSequence:
for i, (frame, depth, pose) in enumerate(tqdm(zip(frames, depth_seq, interp_traj), total=T, desc="Processing video frames")):
if depth.dim() == 3 and depth.shape[-1] == 1:
depth = depth.squeeze(-1)
# Use mask if provided
mask = None
# Use mask if provided; must be (re)initialized every iteration
mask = None
if mask_seq is not None:
mask = mask_seq[i]
if mask.dim() == 3 and mask.shape[-1] == 1:
mask = mask.squeeze(-1)
if mask.dim() == 3 and mask.shape[-1] == 1:
mask = mask.squeeze(-1)
# to pointcloud
pc, = DepthToPointCloud().depth_to_pointcloud(
image=frame.permute(2, 0, 1),
@@ -237,6 +227,10 @@ class DepthFramesToVideo:
raw_color = raw_u8.unsqueeze(1).repeat(1, 3, 1, 1).permute(0, 2, 3, 1)
return raw_color, ds_color # [T, 3, H, W] -> [T, H, W, 3]
# Cache for loaded VideoDepthAnything models, keyed by (checkpoint, device)
_VIDEO_DEPTH_MODEL_CACHE: Dict[Tuple[str, str], Any] = {}
class VideoMetricDepthEstimate:
"""
Estimates metric depth for a sequence of frames using VideoDepthAnything.
@@ -267,16 +261,29 @@ class VideoMetricDepthEstimate:
input_size: int,
max_fps: int,
) -> Tuple[torch.Tensor, float]:
if VideoDepthAnything is None:
raise ImportError("VideoDepthAnything library not found")
if NO_VIDEO_DEPTH_ANYTHING:
raise ImportError(
f"VideoDepthAnything library not found. Clone "
f"https://github.com/DepthAnything/Video-Depth-Anything into {COMFYUI_ROOT!r} "
f"(expected module path: {video_depth_path!r})."
)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# if max input<1.5 normalize to 0-255
if frames.max() < 1.5:
frames = (frames * 255)
model = VideoDepthAnything(**{"encoder": "vitl", "features": 256, "out_channels": [256,512,1024,1024]})
state = torch.load("/root/ComfyUI/models/checkpoints/{}".format(model_checkpoint), map_location='cpu')
model.load_state_dict(state, strict=True)
model = model.to(device).eval()
cache_key = (model_checkpoint, str(device))
model = _VIDEO_DEPTH_MODEL_CACHE.get(cache_key)
if model is None:
# Same checkpoint directory as computed in INPUT_TYPES
model_dir = os.path.join(os.getcwd(), "models", "checkpoints")
checkpoint_path = os.path.join(model_dir, model_checkpoint)
if not os.path.isfile(checkpoint_path):
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
model = VideoDepthAnything(**{"encoder": "vitl", "features": 256, "out_channels": [256,512,1024,1024]})
state = torch.load(checkpoint_path, map_location='cpu')
model.load_state_dict(state, strict=True)
model = model.to(device).eval()
_VIDEO_DEPTH_MODEL_CACHE[cache_key] = model
np_frames = frames.cpu().numpy().astype(np.uint8)
metric_depths, fps = model.infer_video_depth(np_frames, max_fps, input_size=input_size, device=device.type, fp32=False)
return (torch.from_numpy(metric_depths), float(fps))
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
+632
View File
@@ -0,0 +1,632 @@
"""World-building nodes: depth-scale anchoring, splat world enrichment along a
trajectory (render -> outpaint -> SHARP -> align -> fuse) and panorama sphere seeding.
Contracts implemented here (see SPEC_4D.md):
C4: align_depth_scale(new_depth, ref_depth, valid_mask, mode) -> (aligned, scale, shift)
Heavy dependencies (Flux inpainting / diffusers via OutpaintAnyProjection, SHARP)
are only imported/loaded inside methods at call time.
"""
import math
from typing import Any, Dict, Optional, Tuple
import torch
from tqdm import tqdm
try:
import folder_paths
except ImportError: # Allow notebook usage outside ComfyUI
class _FolderPathsStub:
def __getattr__(self, name):
raise ModuleNotFoundError(
"folder_paths is unavailable; this node requires the ComfyUI runtime."
)
folder_paths = _FolderPathsStub()
try:
from . import GS_nodes as _gs
except Exception:
import GS_nodes as _gs
GaussianSplats = _gs.GaussianSplats
Projection = _gs.Projection
DEVICE_CHOICES = _gs.DEVICE_CHOICES
_resolve_device_choice = _gs._resolve_device_choice
splat_cloud_rotation = _gs.splat_cloud_rotation
_stitch_splats = _gs._stitch_splats
# Zeroth-order real SH constant; rendering with add_sh_bias=True computes
# rgb = C0 * f_dc + 0.5, so seeding uses f_dc = (rgb - 0.5) / C0.
SH_C0 = 0.28209479177387814
# ---------------------------------------------------------------------------
# Lazy accessors for symbols provided by sibling modules / heavy dependencies
# ---------------------------------------------------------------------------
def _get_render_gaussians():
"""Fetch GS_nodes.render_gaussians (contract C2) with an actionable error."""
fn = getattr(_gs, "render_gaussians", None)
if fn is None:
raise RuntimeError(
"GS_nodes.render_gaussians is unavailable. Update GS_nodes.py to a version "
"that provides the module-level render_gaussians function (contract C2)."
)
return fn
def _load_outpaint_node_class():
"""Lazy-import OutpaintAnyProjection (pulls in Flux/diffusers machinery)."""
try:
from .flux_fisheye_filling_nodes import OutpaintAnyProjection
return OutpaintAnyProjection
except Exception:
pass
try:
from flux_fisheye_filling_nodes import OutpaintAnyProjection
return OutpaintAnyProjection
except Exception as exc:
raise RuntimeError(
"OutpaintAnyProjection could not be imported from flux_fisheye_filling_nodes. "
"It requires the inpainting_flux custom node package (Flux NF4 inpainting, "
"diffusers). Install/fix custom_nodes/inpainting_flux and its dependencies. "
f"Import error: {exc}"
) from exc
# ---------------------------------------------------------------------------
# C4: robust depth-scale alignment in the disparity domain
# ---------------------------------------------------------------------------
def align_depth_scale(
new_depth: torch.Tensor,
ref_depth: torch.Tensor,
valid_mask: torch.Tensor,
mode: str = "scale_shift",
) -> Tuple[torch.Tensor, float, float]:
"""Least-squares scale(+shift) in DISPARITY (1/d) domain on valid_mask pixels,
robust (clip residual outliers, 2 IRLS rounds). Returns (aligned_depth, scale, shift).
Fits 1/ref_depth ~= scale * (1/new_depth) + shift over valid pixels and returns
new_depth remapped through the fitted disparity transform. If the fit is
degenerate (too few valid pixels, non-positive/non-finite scale), returns the
input depth unchanged with (scale=1.0, shift=0.0).
"""
if mode not in ("scale", "scale_shift"):
raise ValueError(f"Unknown align mode: {mode}")
nd = torch.as_tensor(new_depth).float()
# Harmonize devices: the inputs may arrive on different devices (e.g. a
# CUDA motion mask from MotionMaskFromDepth combined with CPU depth
# estimates); compute everything on new_depth's device.
rd = torch.as_tensor(ref_depth).float().to(nd.device)
vm = torch.as_tensor(valid_mask).float().to(nd.device)
nd_flat = nd.reshape(-1)
rd_flat = rd.reshape(-1)
if vm.numel() == nd_flat.numel():
vm_flat = vm.reshape(-1)
else:
try:
vm_flat = vm.expand_as(nd).reshape(-1)
except RuntimeError as exc:
raise ValueError(
f"valid_mask shape {tuple(vm.shape)} is not broadcastable to depth shape {tuple(nd.shape)}"
) from exc
eps = 1e-8
valid = (
(vm_flat > 0.5)
& (nd_flat > eps)
& (rd_flat > eps)
& torch.isfinite(nd_flat)
& torch.isfinite(rd_flat)
)
if int(valid.sum().item()) < 10:
return nd.clone(), 1.0, 0.0
x = 1.0 / nd_flat[valid] # new disparity
y = 1.0 / rd_flat[valid] # reference disparity
w = torch.ones_like(x)
scale, shift = 1.0, 0.0
# Initial weighted LSQ fit + 2 IRLS re-weighting rounds (outlier clipping).
for _ in range(3):
sw = w.sum().clamp(min=eps)
sx = (w * x).sum()
sy = (w * y).sum()
if mode == "scale_shift":
sxx = (w * x * x).sum()
sxy = (w * x * y).sum()
denom = sw * sxx - sx * sx
if float(denom.abs().item()) < eps:
s = (sxy / sxx.clamp(min=eps)).item()
b = 0.0
else:
s = float(((sw * sxy - sx * sy) / denom).item())
b = float(((sy - s * sx) / sw).item())
else:
sxx = (w * x * x).sum()
sxy = (w * x * y).sum()
s = float((sxy / sxx.clamp(min=eps)).item())
b = 0.0
scale, shift = s, b
resid = y - (scale * x + shift)
sigma = 1.4826 * resid.abs().median()
sigma = sigma.clamp(min=eps)
w = (resid.abs() <= 2.5 * sigma).float()
if float(w.sum().item()) < 10:
break
if not math.isfinite(scale) or scale <= 0.0 or not math.isfinite(shift):
return nd.clone(), 1.0, 0.0
disp = scale / nd.clamp(min=eps) + shift
aligned = 1.0 / disp.clamp(min=eps)
return aligned, float(scale), float(shift)
# ---------------------------------------------------------------------------
# Internal helpers
# ---------------------------------------------------------------------------
def _coerce_trajectory(trajectory: Any, device: torch.device) -> torch.Tensor:
"""Coerce trajectory input to a [K,4,4] float tensor on device."""
if isinstance(trajectory, torch.Tensor):
traj = trajectory
else:
traj = torch.as_tensor(trajectory)
traj = traj.to(device=device, dtype=torch.float32)
if traj.dim() == 2:
traj = traj.unsqueeze(0)
if traj.dim() != 3 or traj.shape[-2:] != (4, 4):
raise ValueError(f"trajectory must be [K,4,4], got shape {tuple(traj.shape)}")
return traj
def _project_to_pixels(
xyz: torch.Tensor,
projection: str,
horizontal_fov: float,
width: int,
height: int,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
"""Project camera-frame points to integer pixel indices.
Returns (ix [N], iy [N], ray_depth [N], valid [N]) where valid means the point
is in front of the camera (pinhole) and lands inside the image bounds. Uses the
same projection math as GS_nodes rendering so pixels line up with renders.
"""
X, Y, Z = xyz.unbind(-1)
if projection == "PINHOLE":
u, v, depth = _gs._xyz_to_pinhole(X, Y, Z, horizontal_fov)
front = Z > 1e-6
elif projection == "FISHEYE":
u, v, depth = _gs._xyz_to_fisheye(X, Y, Z, horizontal_fov)
front = depth > 1e-6
else:
u, v, depth = _gs._xyz_to_equirect(X, Y, Z, horizontal_fov)
front = depth > 1e-6
ix = torch.round((u * 0.5 + 0.5) * (width - 1)).long()
iy = torch.round((v * 0.5 + 0.5) * (height - 1)).long()
inside = (u >= -1.0) & (u <= 1.0) & (v >= -1.0) & (v <= 1.0)
valid = front & inside & torch.isfinite(u) & torch.isfinite(v)
ix = ix.clamp(0, width - 1)
iy = iy.clamp(0, height - 1)
return ix, iy, depth, valid
def _pad_f_rest_to_order(splats: GaussianSplats, sh_order: int) -> GaussianSplats:
"""Zero-pad SH coefficients so splats match the requested (higher) SH order.
Delegates to GS_nodes._pad_sh_order, which handles the renderer's
channel-major SH layout (cat([f_dc, f_rest]).view(-1, 3, total)) correctly.
Naively appending zeros to f_rest would shift the green/blue DC terms into
the red channel's l>=1 slots and corrupt colors.
"""
return _gs._pad_sh_order(splats, sh_order)
def _match_sh_orders(a: GaussianSplats, b: GaussianSplats) -> Tuple[GaussianSplats, GaussianSplats]:
"""Bring two splat sets to a common (max) SH order via zero padding."""
return _gs._match_sh_orders(a, b)
def _scale_splats_metric(splats: GaussianSplats, factor: float) -> GaussianSplats:
"""Uniformly rescale splat positions and sizes by a metric factor."""
out = splats.clone()
out.xyz = out.xyz * factor
out.scale = out.scale + math.log(max(factor, 1e-12))
return out
# ---------------------------------------------------------------------------
# Nodes
# ---------------------------------------------------------------------------
class DepthScaleAnchor:
"""Aligns a depth map's scale (and optionally shift) to a reference depth map
using a robust least-squares fit in the disparity domain (contract C4)."""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
"new_depth": ("TENSOR", {"tooltip": "Depth map to be aligned (any shape)."}),
"ref_depth": ("TENSOR", {"tooltip": "Reference metric depth map (same shape)."}),
"valid_mask": ("MASK", {"tooltip": "1.0 where both depths are trustworthy."}),
"mode": (
["scale", "scale_shift"],
{"default": "scale_shift", "tooltip": "Fit scale only, or scale + shift, in disparity (1/d) domain."},
),
},
}
RETURN_TYPES = ("TENSOR", "FLOAT", "FLOAT")
RETURN_NAMES = ("aligned_depth", "scale", "shift")
FUNCTION = "anchor"
CATEGORY = "Camera/World"
DESCRIPTION = "Robustly aligns a depth map to a reference depth via disparity-domain scale(+shift)."
def anchor(
self,
new_depth: torch.Tensor,
ref_depth: torch.Tensor,
valid_mask: torch.Tensor,
mode: str = "scale_shift",
):
aligned, scale, shift = align_depth_scale(new_depth, ref_depth, valid_mask, mode=mode)
return (aligned, scale, shift)
class SplatTrajectoryEnricher:
"""World-expansion loop for Gaussian splats.
For each pose along a trajectory: render the current splats, detect uncovered
(hole) regions, fill them with Flux outpainting, lift the filled view to new
splats with SHARP, align the SHARP metric scale to the rendered reference
depth, keep only the splats that cover holes, transform them to world space
and fuse them into the running splat set.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
choices = _gs._list_sharp_checkpoint_choices()
return {
"required": {
"splats": ("GSPLAT",),
"trajectory": ("TENSOR", {"tooltip": "[K,4,4] world-to-camera matrices of poses to visit."}),
"camera_projection": (Projection.PROJECTIONS, {}),
"horizontal_fov": ("FLOAT", {"default": 90.0, "min": 1.0, "max": 360.0}),
"width": ("INT", {"default": 512, "min": 8, "max": 8192}),
"height": ("INT", {"default": 512, "min": 8, "max": 8192}),
"checkpoint": (
choices,
{
"default": _gs._SHARP_DEFAULT_CHECKPOINT_LABEL,
"file_chooser": True,
"tooltip": "SHARP .pt checkpoint from the input folder, or download the default model.",
},
),
"prompt": ("STRING", {"default": "", "multiline": True}),
"num_inference_steps": ("INT", {"default": 28, "min": 10, "max": 60}),
"guidance_scale": ("FLOAT", {"default": 5.0, "min": 0.1, "max": 30.0}),
"mask_blur": ("INT", {"default": 5, "min": 0, "max": 512}),
"hole_min_frac": (
"FLOAT",
{"default": 0.02, "min": 0.0, "max": 1.0, "step": 0.001,
"tooltip": "Skip a view if the uncovered area is below this fraction of pixels."},
),
"stitch_voxel_size": ("FLOAT", {"default": 0.01, "min": 0.0, "max": 10.0}),
"max_views": ("INT", {"default": 10, "min": 1, "max": 1000}),
},
"optional": {
"device": (DEVICE_CHOICES, {"default": "auto"}),
"cache_flux": (
"BOOLEAN",
{"default": True,
"tooltip": "Keep the Flux inpainting pipeline loaded between views (avoids a multi-GB "
"model reload per view). Disable to free VRAM after each outpaint on "
"low-memory GPUs."},
),
"patch_projection": (Projection.PROJECTIONS, {"default": "PINHOLE", "tooltip": "Projection used for the outpaint patch."}),
"patch_horiz_fov": ("FLOAT", {"default": 90.0, "min": 1.0, "max": 180.0}),
"patch_res": ("INT", {"default": 1024, "min": 64, "max": 8192}),
"patch_phi": ("FLOAT", {"default": 0.0, "min": -180.0, "max": 180.0}),
"patch_theta": ("FLOAT", {"default": 0.0, "min": -90.0, "max": 90.0}),
},
}
RETURN_TYPES = ("GSPLAT", "IMAGE", "IMAGE")
RETURN_NAMES = ("enriched_splats", "last_render", "last_filled")
FUNCTION = "enrich"
CATEGORY = "Camera/World"
DESCRIPTION = (
"Expands a splat world along a camera trajectory: render, outpaint holes with Flux, "
"lift with SHARP, scale-align, and smart-stitch the new content."
)
@torch.no_grad()
def enrich(
self,
splats: GaussianSplats,
trajectory: torch.Tensor,
camera_projection: str,
horizontal_fov: float,
width: int,
height: int,
checkpoint: str,
prompt: str,
num_inference_steps: int,
guidance_scale: float,
mask_blur: int,
hole_min_frac: float,
stitch_voxel_size: float,
max_views: int,
device: str = "auto",
cache_flux: bool = True,
patch_projection: str = "PINHOLE",
patch_horiz_fov: float = 90.0,
patch_res: int = 1024,
patch_phi: float = 0.0,
patch_theta: float = 0.0,
) -> Tuple[GaussianSplats, torch.Tensor, torch.Tensor]:
# Fail fast: the SHARP lift (ImageToSplat) is pinhole-only and requires
# horizontal_fov < 179 degrees. Validating here avoids crashing in the
# lift step AFTER minutes of rendering + Flux outpainting work.
if not (0.0 < float(horizontal_fov) < 179.0):
raise ValueError(
"SplatTrajectoryEnricher lifts filled views with SHARP (pinhole), which requires "
f"0 < horizontal_fov < 179 degrees (got {horizontal_fov}). For panoramic worlds "
"(EQUIRECTANGULAR/FISHEYE with fov >= 179), visit several narrower pinhole poses "
"along the trajectory instead (e.g. 90-120 degree views after SphereSplatSeed)."
)
render_gaussians = _get_render_gaussians()
outpaint_cls = _load_outpaint_node_class()
outpaint_node = outpaint_cls()
image_to_splat = _gs.ImageToSplat()
target_device = _resolve_device_choice(device)
current = splats.to(target_device) if splats.xyz.device != target_device else splats
traj = _coerce_trajectory(trajectory, target_device)
if camera_projection != "PINHOLE":
print(
"[SplatTrajectoryEnricher] Warning: SHARP assumes pinhole geometry; "
f"lifting filled {camera_projection} views may distort new splats."
)
last_render = torch.zeros((1, height, width, 3), device=target_device)
last_filled = torch.zeros((1, height, width, 3), device=target_device)
added_views = 0
for pose in tqdm(traj[: max(1, int(max_views))], desc="Enriching splat world"):
# 1) Render the current world from this pose.
image, alpha, disparity = render_gaussians(
current,
pose,
camera_projection,
horizontal_fov,
width,
height,
max_splats=0,
opacity_is_logit=True,
add_sh_bias=True,
render_mode="auto",
device=str(target_device).split(":")[0],
)
last_render = image
alpha_map = alpha.view(height, width).to(target_device)
disp_map = disparity.view(height, width).to(target_device)
hole_mask = (alpha_map < 0.5).float()
hole_frac = float(hole_mask.mean().item())
if hole_frac < hole_min_frac:
continue
# 2) Outpaint the uncovered region.
filled_img, _ = outpaint_node.outpaint_any(
image,
input_projection=camera_projection,
input_horiz_fov=horizontal_fov,
output_projection=camera_projection,
output_horiz_fov=horizontal_fov,
output_width=width,
output_height=height,
patch_projection=patch_projection,
patch_horiz_fov=patch_horiz_fov,
patch_res=patch_res,
patch_phi=patch_phi,
patch_theta=patch_theta,
prompt=prompt,
num_inference_steps=num_inference_steps,
# cached=True keeps the Flux NF4 pipeline resident between views
# (cached=False forced a full multi-GB pipeline reload per view).
cached=bool(cache_flux),
guidance_scale=guidance_scale,
mask_blur=mask_blur,
mask=hole_mask.unsqueeze(0),
debug=False,
)
last_filled = filled_img
# 3) Lift the filled view to splats in this camera frame (SHARP, metric).
new_splats, = image_to_splat.image_to_splat(
filled_img,
horizontal_fov,
checkpoint,
device,
)
new_splats = new_splats.to(target_device)
if len(new_splats) == 0:
continue
# 4) Robust metric-scale alignment against the rendered reference depth.
# Reference ray depth from the renderer: disparity = alpha / depth.
ix, iy, sharp_depth, proj_valid = _project_to_pixels(
new_splats.xyz, camera_projection, horizontal_fov, width, height
)
samp_alpha = alpha_map[iy, ix]
samp_disp = disp_map[iy, ix]
overlap = proj_valid & (samp_alpha >= 0.5) & (samp_disp > 1e-6) & (sharp_depth > 1e-6)
if int(overlap.sum().item()) >= 10:
d_ref = (samp_alpha[overlap] / samp_disp[overlap]).clamp(min=1e-6)
ratio = d_ref / sharp_depth[overlap]
scale_factor = float(ratio.median().item())
if math.isfinite(scale_factor) and scale_factor > 0.0:
new_splats = _scale_splats_metric(new_splats, scale_factor)
# 5) Keep only NEW content: splats whose projected pixel lies in a hole.
samp_hole = hole_mask[iy, ix]
keep = proj_valid & (samp_hole > 0.5)
if not bool(keep.any().item()):
continue
new_splats = new_splats[keep]
# 6) Camera frame -> world frame (pose is world-to-camera).
new_world = splat_cloud_rotation(new_splats, torch.inverse(pose))
# 7) Fuse into the running world. Concatenation is cheap; the full
# smart voxel reduce is deferred to a single pass after the loop,
# so each view does not re-copy and re-unique-sort the entire
# accumulated cloud (O(views x N) work/memory otherwise).
cur_m, new_m = _match_sh_orders(current, new_world)
current = _gs._concat_splats([cur_m, new_m])
added_views += 1
if added_views > 0 and stitch_voxel_size > 0.0:
current = _stitch_splats([current], "smart", stitch_voxel_size, 5.0)
return (current, last_render, last_filled)
class SphereSplatSeed:
"""Seeds a 360-degree splat world from an equirectangular panorama: one Gaussian
per (subsampled) pixel, placed on a depth sphere around the origin."""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
"image": ("IMAGE", {"tooltip": "Equirectangular panorama [1,H,W,3]."}),
"horizontal_fov": ("FLOAT", {"default": 360.0, "min": 1.0, "max": 360.0}),
"radius": ("FLOAT", {"default": 5.0, "min": 0.01, "max": 10000.0, "tooltip": "Sphere radius used when no depth map is provided."}),
"splat_scale_frac": (
"FLOAT",
{"default": 1.5, "min": 0.1, "max": 10.0,
"tooltip": "Splat sigma as a fraction of the local point spacing (larger = smoother, fewer holes)."},
),
"stride": ("INT", {"default": 2, "min": 1, "max": 64, "tooltip": "Pixel subsampling stride (1 Gaussian per stride x stride block)."}),
},
"optional": {
"depth": ("TENSOR", {"tooltip": "Optional ray-depth map [H,W] (or [1,H,W]/[H,W,1]) matching the panorama."}),
"opacity_logit": ("FLOAT", {"default": 6.0, "min": -10.0, "max": 20.0}),
"device": (DEVICE_CHOICES, {"default": "auto"}),
},
}
RETURN_TYPES = ("GSPLAT",)
RETURN_NAMES = ("splats",)
FUNCTION = "seed_sphere"
CATEGORY = "Camera/World"
DESCRIPTION = "Converts an equirectangular panorama into a Gaussian sphere seeding a 360-degree world."
@torch.no_grad()
def seed_sphere(
self,
image: torch.Tensor,
horizontal_fov: float = 360.0,
radius: float = 5.0,
splat_scale_frac: float = 1.5,
stride: int = 2,
depth: Optional[torch.Tensor] = None,
opacity_logit: float = 6.0,
device: str = "auto",
) -> Tuple[GaussianSplats]:
target_device = _resolve_device_choice(device)
img = image
if img.dim() == 4:
img = img[0]
if img.dim() != 3 or img.shape[-1] < 3:
raise ValueError(f"Expected IMAGE [1,H,W,3], got shape {tuple(image.shape)}")
img = img[..., :3].to(device=target_device, dtype=torch.float32)
H, W = int(img.shape[0]), int(img.shape[1])
depth_map = None
if depth is not None:
d = torch.as_tensor(depth).to(device=target_device, dtype=torch.float32)
if d.dim() == 3:
# [1,H,W], [T,H,W] (take first) or [H,W,1]
d = d[..., 0] if d.shape[-1] == 1 else d[0]
if d.dim() != 2:
raise ValueError(f"depth must reduce to [H,W], got shape {tuple(depth.shape)}")
if d.shape != (H, W):
d = torch.nn.functional.interpolate(
d.unsqueeze(0).unsqueeze(0), size=(H, W), mode="bilinear", align_corners=True
)[0, 0]
depth_map = d.clamp(min=1e-6)
stride = max(1, int(stride))
ys = torch.arange(0, H, stride, device=target_device)
xs = torch.arange(0, W, stride, device=target_device)
yy, xx = torch.meshgrid(ys, xs, indexing="ij")
yy = yy.reshape(-1)
xx = xx.reshape(-1)
# Match the renderer's equirect mapping (GS_nodes._xyz_to_equirect):
# u = lon / (fov_rad/2), v = lat / (pi/2), px = (u*0.5+0.5)*(W-1)
fov_rad = math.radians(horizontal_fov)
u = xx.float() / max(W - 1, 1) * 2.0 - 1.0
v = yy.float() / max(H - 1, 1) * 2.0 - 1.0
lon = u * (fov_rad / 2.0)
lat = v * (math.pi / 2.0)
if depth_map is not None:
d = depth_map[yy, xx]
else:
d = torch.full_like(lon, float(radius))
cos_lat = torch.cos(lat)
X = d * cos_lat * torch.sin(lon)
Y = d * torch.sin(lat)
Z = d * cos_lat * torch.cos(lon)
xyz = torch.stack([X, Y, Z], dim=-1)
rgb = img[yy, xx, :]
# Rendering with add_sh_bias=True evaluates rgb = C0 * f_dc + 0.5.
f_dc = (rgb - 0.5) / SH_C0
# Isotropic sigma from local angular spacing (radians per sample) times depth.
ang_spacing = float(stride) * max(fov_rad / max(W, 1), math.pi / max(H, 1))
sigma = (splat_scale_frac * ang_spacing * d).clamp(min=1e-6)
scale = torch.log(sigma).unsqueeze(-1).expand(-1, 3).contiguous()
n = xyz.shape[0]
rotation = torch.zeros((n, 4), device=target_device, dtype=torch.float32)
rotation[:, 0] = 1.0 # identity wxyz quaternion
opacity = torch.full((n, 1), float(opacity_logit), device=target_device, dtype=torch.float32)
f_rest = torch.zeros((n, 0), device=target_device, dtype=torch.float32)
splats = GaussianSplats(
xyz=xyz,
scale=scale,
rotation=rotation,
opacity=opacity,
f_dc=f_dc,
f_rest=f_rest,
sh_order=0,
)
return (splats,)
NODE_CLASS_MAPPINGS = {
"DepthScaleAnchor": DepthScaleAnchor,
"SplatTrajectoryEnricher": SplatTrajectoryEnricher,
"SphereSplatSeed": SphereSplatSeed,
}