Compare commits

...
Author SHA1 Message Date
Alexander KharinandClaude Fable 5 358b2b55ae Add LingBot-World 2.0 video-to-4D research report
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-13 12:07:21 +03:00
Alexander KharinandClaude Fable 5 8598a7b51c Add ComfyUI Registry publishing setup (MIT license, pyproject, publish workflow)
- pyproject.toml with [tool.comfy] metadata (publisher alexk); force-includes
  the ml-sharpt submodule since gitlinks are not packed into registry archives
- GitHub Action publishing on pyproject.toml version bumps on main
- .comfyignore to keep demos/notebooks/docs out of the download archive
- MIT LICENSE (ml-sharpt submodule keeps its own Apple license)
- README: install via ComfyUI Manager section

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-13 12:07:09 +03:00
Alexander KharinandClaude Fable 5 90df387c2a Add 4D Gaussian splatting, pose, and world nodes with video-to-4D workflows
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-07-11 09:35:40 +03:00
Alexander Kharin 74d71a91f2 update splats stitching 2025-12-26 19:27:09 +03:00
Alexander Kharin 55820b53e8 add image to splat submodule 2025-12-26 15:01:42 +03:00
Alexander Kharin f6b6707dd8 add GS nodes 2025-12-25 01:32:59 +03:00
Alexander Kharin 49f9700880 Merge pull request #19 from gitcapoom/claude/camera-comf-modified-011CUt8BMsMPugiCMpz1j9ah
Fix equirectangular projection aspect ratio to 2:1
2025-12-24 23:34:29 +03:00
Claude b961c2988d Revert input vertical FOV to use grid aspect ratio
- Reverts input_vertical_fov calculation to use output grid aspect ratio
- Keeps output_vertical_fov fix for equirectangular (2:1 aspect)
- Fixes vertical squeezing when converting FROM equirectangular
- Input FOV now adapts to output dimensions as originally intended
2025-11-08 19:54:09 +00:00
Claude 2b1b716495 Fix vertical squeezing in equirectangular to pinhole conversion
- Remove output grid aspect ratio from input vertical FOV calculation
- For non-equirectangular inputs, use square aspect (horizontal FOV)
- The normalized_grid already handles output projection aspect ratio
- Fixes 2x vertical compression when converting from equirectangular to pinhole
2025-11-08 18:48:00 +00:00
Claude 839fa5883a Fix width/height swap in ReprojectImage meshgrid
- Corrected meshgrid arguments to use height for y-axis and width for x-axis
- With indexing="ij", first arg should be height, second should be width
- ReprojectDepth was already correct and didn't need changes
2025-11-08 08:43:20 +00:00
Claude 74f98be5c2 Fix equirectangular projection aspect ratio to 2:1
- Set vertical FOV to horizontal FOV / 2 for equirectangular projections
- Applies to both input and output projections
- Eliminates wasted computation on top/bottom quarters
- Reduces resource usage by ~50% for equirectangular outputs
- Affects both ReprojectImage and ReprojectDepth nodes
2025-11-08 07:46:45 +00:00
Alexander Kharin c1e1b55464 replace video with gif for readme 2025-07-13 20:05:49 +02:00
Alexander Kharin 08b48b2fcb Update camera movement on video 2025-07-13 20:03:09 +02:00
Alexander Kharin 4812973e0f Fix background-foregraound. Update video_camera workflow 2025-07-13 19:39:30 +02:00
Alexander Kharin 57fc1e8593 do not import videonodes if video-depth anything is not installed, update projection logics to fill the holes in foreground 2025-07-13 12:27:30 +02:00
Alexander Kharin abdfdc188e fix mask format 2025-07-07 22:59:41 +02:00
Alexander Kharin ca6f904fc3 update projection node to accept masks for processing non square videos. Add video workflow for camera movement in videos 2025-07-07 22:22:42 +02:00
Alexander Kharin 696a4ad763 refactor pointcloud projector for better hole-filling mechanism 2025-07-07 20:15:10 +02:00
Alexander Kharin d67cd61c3c update installation script (add video depth anything). Add some ugly path management for video-depth-anything 2025-06-28 18:30:39 +02:00
Alexander Kharin 17a0edb3c1 fix bug with wrong optional type 2025-06-28 17:29:11 +02:00
Alexander Kharin 5592382b44 Merge pull request #15 from Alexankharin/5-loadpointcloud-node-errors
add video nodes
2025-06-28 17:23:51 +02:00
Alexander Kharin de679db043 add vidoe nodes 2025-06-21 23:59:09 +02:00
Alexander Kharin bc358f8263 Merge pull request #13 from Alexankharin/5-loadpointcloud-node-errors
update test workflow and readme, add trajectory example to load
2025-06-12 21:44:03 +03:00
Alexander Kharin 1d0ae6dc61 Merge branch 'main' into 5-loadpointcloud-node-errors 2025-06-12 21:43:47 +03:00
Alexander Kharin 3bedb49949 update test workflow and readme, add trajectory example to load 2025-06-10 22:51:36 +02:00
Alexander Kharin b67ae9a0a2 Merge pull request #11 from Alexankharin/codex/refactor-node-structuring-by-categories
Improve node category organization
2025-06-10 23:35:34 +03:00
26 changed files with 12415 additions and 138 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.
+137 -8
View File
@@ -1,6 +1,7 @@
# camera-comfyUI
[![Ask DeepWiki](https://deepwiki.com/badge.svg)](https://deepwiki.com/Alexankharin/camera-comfyUI)
![ComfyUI Custom Nodes](demo_images/Camera_interpolation_pointcloud.gif)
![Camera Movement Demo](demo_images/camera_movement.gif)
> Custom ComfyUI nodes for advanced reprojections, point cloud processing, and camera-driven workflows.
@@ -13,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)
@@ -34,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
@@ -54,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:
@@ -93,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
@@ -115,12 +152,65 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
| `DepthToPointCloud` | Converts Depth and image to → 3D point cloud tensor (N×7). |
| `DepthToImageNode` | Converts depth to image (N×3) using a color map. |
| `ZDepthToRayDepthNode` | Converts Z-depth (output of metric-depth-anything) to ray depth to compensate lens curvature. |
| `TransformPointCloud` | Applies 4×4 rotation matrix to point cloud |
| `TransformPointCloud` | Applies 4×4 rotation matrix to point cloud. |
| `ProjectPointCloud` | Z-buffer–based projection of point cloud into image + mask. |
| `PointCloudCleaner` | Removes isolated points via voxel filtering. |
| `PointCloudUnion` | Combines multiple point clouds into one. |
| `LoadPointCloud` | Loads a point cloud from `.npy` or `.ply` format. |
| `SavePointCloud` | Saves a point cloud to `.npy` or `.ply` format. |
| `CameraMotionNode` | Generates image and mask sequences along a camera trajectory with optional mask dilation/inversion. |
| `CameraInterpolationNode` | Builds a trajectory tensor from two poses. |
| `CameraTrajectoryNode` | Interactive Open3D GUI for recording camera waypoints. |
| `PointCloudCleaner` | Removes isolated points via voxel filtering. |
| `SaveTrajectory` | Saves a trajectory tensor to a file. |
| `LoadTrajectory` | Loads a trajectory tensor from a file. |
| `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`.
---
@@ -140,6 +230,9 @@ A set of JSON workflows illustrating typical use cases. Each workflow lives in `
| **pointcloud\_inpaint.json** | Inpaint + backproject to 3D for dynamic camera motion videos |
| **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. |
---
@@ -208,11 +301,44 @@ Take a wide-angle (fisheye or equirectangular) high-resolution (e.g., 4096×4096
Interactive Open3D-based GUI for walking and setting camera trajectory inside pointcloud.
### 11. `wan-vace_ref_to_video.json`
### 11. `video_camera.json`
Integrate the [wan2.1-vace] video generation model to inpaint empty or newly revealed regions during camera movement or view synthesis. This workflow demonstrates how to use the camera-comfyUI nodes to generate camera trajectories and masks, then fill missing areas with the video inpainting model for smooth, high-quality results.
This workflow demonstrates camera trajectory movement using the `wan-vace` video inpainting model. It generates smooth camera movements along a trajectory while filling missing regions with high-quality inpainting.
<img src="demo_images/wan-vace-camera.gif" alt="wan2.1-vace Camera Inpainting Demo" width="80%" />
<div style="display:flex; gap:10px;">
<img src="demo_images/camera_movement.gif" alt="Camera Movement Demo" width="80%" />
</div>
---
## Trajectory Concept
A **trajectory** in camera-comfyUI is a sequence of camera poses, each represented as a 4×4 transformation matrix. This set of matrices defines the path and orientation of the camera through 3D space, enabling smooth and complex camera movements for view synthesis, point cloud rendering, and video generation.
### Creating Trajectories
There are two main ways to create a trajectory:
- **Camera Matrices Interpolation:**
Define two or more camera poses (as matrices), and interpolate between them to generate a smooth path. The `CameraInterpolationNode` automates this process, producing a trajectory tensor for use in camera motion nodes.
- **Walking in Open3D Environment:**
Use the interactive Open3D GUI (`CameraTrajectoryNode`) to "walk" through the point cloud. As you move the camera, waypoints (poses) are recorded, forming a trajectory that can be exported and reused.
### Using Trajectories
The `CameraMotionNode` takes a trajectory (set of matrices) and interpolates camera positions and orientations along it, producing smooth camera movements for rendering sequences or videos.
---
## Point Cloud Formats
Point clouds can be saved and loaded in two formats:
- **.npy**: Numpy array format (fast, preserves all tensor data, recommended for internal pipelines).
- **.ply**: Polygon File Format (widely supported, viewable in external 3D tools).
Use the `SavePointCloud` and `LoadPointCloud` nodes to handle I/O operations in either format.
---
@@ -225,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
* [ ] Fix imports for renamed folders (e.g., inpainting_flux)
* [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.
+24 -2
View File
@@ -3,6 +3,28 @@ from .reprojection_nodes import NODE_CLASS_MAPPINGS as NCM2
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
NODE_CLASS_MAPPINGS = {**NCM1, **NCM2, **NCM3, **NCM4, **NCM5}
from .video_nodes import NODE_CLASS_MAPPINGS as 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"]
Binary file not shown.
Binary file not shown.
Binary file not shown.

After

Width:  |  Height:  |  Size: 12 MiB

+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.*
+64 -27
View File
@@ -5,76 +5,104 @@ set -euo pipefail
# Functions
# ----------------------------------------
install_pytorch() {
echo "Installing PyTorch, TorchVision, TorchAudio..."
echo "==> Installing PyTorch, TorchVision, TorchAudio, bitsandbytes, accelerate…"
pip3 install -U torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128
pip3 install bitsandbytes
pip3 install accelerate
pip3 install -U bitsandbytes accelerate
}
install_system_deps() {
echo "Updating apt and installing system dependencies..."
echo "==> Updating apt and installing system packages…"
sudo apt-get update
sudo apt-get install -y build-essential ffmpeg libsm6 libxext6
sudo apt-get install -y build-essential ffmpeg libsm6 libxext6 python3.10-dev
}
clone_and_install_comfyui() {
echo "Cloning ComfyUI..."
echo "==> Cloning ComfyUI…"
git clone https://github.com/comfyanonymous/ComfyUI.git
echo "Installing ComfyUI requirements..."
echo "==> Installing ComfyUI Python requirements…"
pip3 install -r ComfyUI/requirements.txt
}
install_camera_node() {
echo "Cloning camera‑ComfyUI..."
echo "==> Installing camera‑ComfyUI node…"
mkdir -p ComfyUI/custom_nodes
git clone https://github.com/Alexankharin/camera-comfyUI.git \
ComfyUI/custom_nodes/camera-comfyUI
echo "Installing camera‑ComfyUI requirements..."
pip3 install -r ComfyUI/custom_nodes/camera-comfyUI/requirements.txt
}
install_image_filters() {
echo "Cloning Image‑Filters node..."
echo "==> Installing Image‑Filters node…"
mkdir -p ComfyUI/custom_nodes
git clone https://github.com/spacepxl/ComfyUI-Image-Filters.git \
ComfyUI/custom_nodes/ComfyUI-Image-Filters
echo "Installing Image‑Filters requirements..."
pip3 install -r ComfyUI/custom_nodes/ComfyUI-Image-Filters/requirements.txt
}
clone_flux_inpainting() {
echo "Cloning ComfyUI‑Flux‑Inpainting..."
echo "==> Installing Flux‑Inpainting node…"
mkdir -p ComfyUI/custom_nodes
git clone https://github.com/rubi-du/ComfyUI-Flux-Inpainting.git \
ComfyUI/custom_nodes/Flux-Inpainting
echo "Renaming Flux‑Inpainting folder..."
mv ComfyUI/custom_nodes/Flux-Inpainting \
ComfyUI/custom_nodes/inpainting_flux
ComfyUI/custom_nodes/inpainting_flux
}
install_metric_video_depth_anything() {
echo "==> Installing Metric Video Depth Anything…"
# Clone into ComfyUI root, not custom_nodes
git clone https://github.com/DepthAnything/Video-Depth-Anything.git \
ComfyUI/Video-Depth-Anything
echo " • Installing easydict…"
pip3 install -U easydict
echo " • Copying util.py…"
mkdir -p ComfyUI/utils
cp ComfyUI/Video-Depth-Anything/metric_depth/utils/util.py \
ComfyUI/utils/util.py
echo " • Downloading Metric Video Depth checkpoint…"
mkdir -p ComfyUI/models/checkpoints
wget -q -O ComfyUI/models/checkpoints/metric_video_depth_anything_vitl.pth \
"https://huggingface.co/depth-anything/Metric-Video-Depth-Anything-Large/resolve/main/metric_video_depth_anything_vitl.pth"
}
install_comfyui_manager() {
echo "==> Installing ComfyUI-Manager extension…"
mkdir -p ComfyUI/custom_nodes
git clone https://github.com/Comfy-Org/ComfyUI-Manager.git \
ComfyUI/custom_nodes/ComfyUI-Manager
pip3 install -r ComfyUI/custom_nodes/ComfyUI-Manager/requirements.txt
}
install_hf_hub() {
echo "Installing huggingface_hub..."
pip3 install huggingface_hub
echo "==> Installing huggingface_hub…"
pip3 install -U huggingface_hub
}
download_vae_models() {
echo "Downloading WAN‑VACE models via wget..."
wget -O ComfyUI/models/vae/wan_2.1_vae.safetensors \
echo "==> Downloading WAN‑VACE models…"
mkdir -p ComfyUI/models/vae
wget -q -O ComfyUI/models/vae/wan_2.1_vae.safetensors \
"https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/resolve/main/split_files/vae/wan_2.1_vae.safetensors?download=true"
wget -O ComfyUI/models/text_encoders/umt5_xxl_fp16.safetensors \
mkdir -p ComfyUI/models/text_encoders
wget -q -O ComfyUI/models/text_encoders/umt5_xxl_fp16.safetensors \
"https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/resolve/main/split_files/text_encoders/umt5_xxl_fp16.safetensors?download=true"
wget -O ComfyUI/models/diffusion_models/wan2.1_vace_14B_fp16.safetensors \
mkdir -p ComfyUI/models/diffusion_models
wget -q -O ComfyUI/models/diffusion_models/wan2.1_vace_14B_fp16.safetensors \
"https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/resolve/main/split_files/diffusion_models/wan2.1_vace_14B_fp16.safetensors"
}
login_hf_hub() {
echo "Logging in to Hugging Face Hub..."
echo "==> Hugging Face login…"
huggingface-cli login
}
# ----------------------------------------
# Main CLI
# Main
# ----------------------------------------
# default to "install" if no arg given
MODE="${1:-install}"
case "$MODE" in
@@ -84,6 +112,7 @@ case "$MODE" in
clone_and_install_comfyui
install_camera_node
install_image_filters
install_comfyui_manager
install_hf_hub
;;
@@ -95,6 +124,10 @@ case "$MODE" in
download_vae_models
;;
depth)
install_metric_video_depth_anything
;;
install)
install_pytorch
install_system_deps
@@ -102,7 +135,9 @@ case "$MODE" in
install_camera_node
clone_flux_inpainting
install_image_filters
install_comfyui_manager
install_hf_hub
install_metric_video_depth_anything
;;
all)
@@ -112,14 +147,16 @@ case "$MODE" in
install_camera_node
clone_flux_inpainting
install_image_filters
install_comfyui_manager
install_hf_hub
download_vae_models
install_metric_video_depth_anything
login_hf_hub
echo "All done! 🎉"
echo "✅ All done!"
;;
*)
echo "Usage: $0 {install|modules|flux|vae|all}"
echo "Usage: $0 {install|modules|flux|vae|depth|all}"
exit 1
;;
esac
+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
+351 -98
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:
"""
@@ -325,120 +438,165 @@ class ProjectPointCloud:
def project_pointcloud(
self,
pointcloud: torch.Tensor,
pointcloud: torch.Tensor,
output_projection: str,
output_horizontal_fov: float,
output_width: int,
output_width: int,
output_height: int,
point_size: int = 1,
point_size: int = 1,
return_inverse_depth: bool = False,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Projects an (N×6) XYZRGB point cloud into an image,
fills occlusion holes robustly, and returns:
• img: [1,H,W,3] RGB image
• mask: [H,W] foreground mask
• depth: [1,H,W,1] depth (or inverse depth)
"""
device = pointcloud.device
coords = pointcloud[:, :3]
colors = pointcloud[:, 3:].float()
xyz, rgb_raw = pointcloud[:, :3], pointcloud[:, 3:6].float()
# 1) Filter points in front of the camera
mask_front = coords[:, 2] > 0
coords = coords[mask_front]
colors = colors[mask_front]
# 1) Keep only points in front of camera
in_front = xyz[:, 2] > 0
xyz, rgb_raw = xyz[in_front], rgb_raw[in_front]
# 2) Project to normalized UV + depth
X, Y, Z = coords.unbind(1)
X, Y, Z = xyz.unbind(1)
if output_projection == "PINHOLE":
u, v, depth = XYZ_to_pinhole(X, Y, Z, output_horizontal_fov)
u, v, d = XYZ_to_pinhole(X, Y, Z, output_horizontal_fov)
elif output_projection == "FISHEYE":
u, v, depth = XYZ_to_fisheye(X, Y, Z, output_horizontal_fov)
u, v, d = XYZ_to_fisheye(X, Y, Z, output_horizontal_fov)
else:
u, v, depth = XYZ_to_equirect(X, Y, Z, output_horizontal_fov)
u, v, d = XYZ_to_equirect(X, Y, Z, output_horizontal_fov)
# 3) Rasterize to pixel indices
px = (u * (output_width - 1) / 2) + (output_width - 1) / 2
py = (v * (output_height - 1) / 2) + (output_height - 1) / 2
ix = px.round().clamp(0, output_width - 1).long()
iy = py.round().clamp(0, output_height - 1).long()
pix = iy * output_width + ix
M = output_width * output_height
W, H = output_width, output_height
ix = ((u * 0.5 + 0.5) * (W - 1)).round().clamp(0, W - 1).long()
iy = ((v * 0.5 + 0.5) * (H - 1)).round().clamp(0, H - 1).long()
pix = iy * W + ix
valid = (pix >= 0) & (pix < W * H)
pix, d, rgb_raw = pix[valid], d[valid], rgb_raw[valid]
# —— NEW: drop any invalid / NaN→int_min projections ——
valid = (pix >= 0) & (pix < M)
depth = depth[valid]
colors = colors[valid]
pix = pix[valid]
order = torch.arange(depth.size(0), device=device)
# rebuild your "order" to match
M = W * H
# 3a) Front‐layer (nearest) depth
z1 = torch.full((M,), float('inf'), device=device)
z1.scatter_reduce_(0, pix, d, reduce='amin', include_self=True)
# 4) Allocate or reuse buffers
if not hasattr(self, '_z_front') or self._z_front.numel() != M:
self._z_front = torch.empty((M,), device=device)
self._z_back = torch.empty((M,), device=device)
self._idx = torch.full((M,), -1, dtype=torch.long, device=device)
self._flat = torch.zeros((M, 4), device=device)
z_front = self._z_front
z_back = self._z_back
idxbuf = self._idx
flat = self._flat
# 3b) Second‐layer (background) depth
farther = d > z1[pix]
pix2, d2 = pix[farther], d[farther]
z2 = torch.full((M,), float('inf'), device=device)
z2.scatter_reduce_(0, pix2, d2, reduce='amin', include_self=True)
# 5) Front z-buffer pass (nearest)
z_front.fill_(float('inf'))
z_front.scatter_reduce_(0, pix, depth, reduce='amin', include_self=True)
sel_front = depth == z_front[pix]
order = torch.arange(depth.size(0), device=device)
order_m = torch.where(sel_front, order, depth.size(0))
idxbuf.fill_(depth.size(0))
idxbuf.scatter_reduce_(0, pix, order_m, reduce='amin', include_self=True)
win_front = order == idxbuf[pix]
# 3c) Foreground colour (from z1)
keep = d == z1[pix]
rgb = torch.zeros((M, 3), device=device)
rgb[pix[keep]] = rgb_raw[keep].clamp(0, 255)
flat.fill_(0)
flat[pix[win_front]] = colors[win_front]
img4 = flat.view(output_height, output_width, 4)
rgb = img4[..., :3].clamp(0, 255)
alpha = (img4[..., 3] > 0).float()
rgb *= alpha.unsqueeze(-1)
depth_img = z_front.view(output_height, output_width)
rgb_HR = rgb
# reshape to image
rgb = rgb.view(H, W, 3)
z1 = z1.view(H, W)
z2 = z2.view(H, W)
fg_mask = z1 < float('inf') # has front hit
occl = (z2 < float('inf')) # has any back hit
rear_only = occl & ~fg_mask # true holes
# 6) Back z-buffer pass (farthest) for hole-filling
# ── Iterative ring‐based in‐painting of rear‐only pixels ─────────────────────
ker3 = torch.ones((1,1,3,3), device=device)
ker3c = ker3.repeat(3,1,1,1)
for _ in range(max(W, H)):
# find rear_only pixels adjacent to current FG
neigh = (
F.max_pool2d(fg_mask.float()[None,None], 3, 1, 1).bool()[0,0]
& ~fg_mask
)
to_fill = rear_only & neigh
if not to_fill.any():
break
# average depth + colour from current FG frontier
d_t = z1.masked_fill(~fg_mask, 0)[None,None]
c_t = rgb.permute(2,0,1)[None] # [1,3,H,W]
m_t = fg_mask.float()[None,None]
sum_d = F.conv2d(d_t * m_t, ker3, padding=1)
cnt_d = F.conv2d(m_t, ker3, padding=1).clamp(min=1)
sum_c = F.conv2d(c_t * m_t, ker3c, padding=1, groups=3)
cnt_c = cnt_d.repeat(1,3,1,1)
avg_d = (sum_d / cnt_d).squeeze()
avg_c = (sum_c / cnt_c).squeeze().permute(1,2,0)
z1[to_fill] = avg_d[to_fill]
rgb[to_fill] = avg_c[to_fill]
fg_mask[to_fill] = True
rear_only[to_fill] = False
# ── Depth‐aware generic hole closure ─────────────────────────────────────────
# close_rad: radius of hole to close; depth_eps: depth jump tolerance
close_rad = max(1, point_size // 2)
depth_eps = 0.015
pad = close_rad
k = 2 * close_rad + 1
ker = torch.ones((1,1,k,k), device=device)
kerc = ker.repeat(3,1,1,1)
front_t = fg_mask.float()[None,None]
# binary closing: dilate then erode
D = F.max_pool2d(front_t, k, 1, pad)
E = 1 - F.max_pool2d(1 - D, k, 1, pad)
small_hole = E[0,0].bool() & ~fg_mask
if small_hole.any():
# compute local mean depth of FG
z_t = z1.masked_fill(~fg_mask, 0)[None,None]
cnt = F.conv2d(front_t, ker, padding=pad).clamp(min=1)
z_avg = (F.conv2d(z_t, ker, padding=pad) / cnt)[0,0]
# depth‐range test
z_near = F.max_pool2d(z1[None,None], 3,1,1)[0,0]
z_far = -F.max_pool2d(-z1[None,None],3,1,1)[0,0]
flat = (z_far - z_near) / z_avg.clamp(min=1e-6) < depth_eps
final = small_hole & flat
if final.any():
sum_d = F.conv2d(z_t, ker, padding=pad)
sum_c = F.conv2d(rgb.permute(2,0,1)[None] * front_t, kerc,
padding=pad, groups=3)
avg_d = (sum_d / cnt)[0,0]
avg_c = (sum_c / cnt.repeat(1,3,1,1))[0].permute(1,2,0)
z1[final] = avg_d[final]
rgb[final] = avg_c[final]
fg_mask[final] = True
# ── Optional morphological blur for larger point_size ───────────────────────
if point_size > 1:
z_back.fill_(-float('inf'))
z_back.scatter_reduce_(0, pix, depth, reduce='amax', include_self=True)
sel_back = depth == z_back[pix]
order_m = torch.where(sel_back, order, -1)
idxbuf.fill_(-1)
idxbuf.scatter_reduce_(0, pix, order_m, reduce='amax', include_self=True)
win_back = idxbuf[pix] >= 0
r = point_size // 2
k = 2 * r + 1
pad = r
ker = torch.ones((1,1,k,k), device=device)
kerc = ker.repeat(3,1,1,1)
d_t = z1[None,None]
c_t = rgb.permute(2,0,1)[None]
m_t = fg_mask.float()[None,None]
flat.fill_(0)
flat[pix[win_back]] = colors[win_back]
back4 = flat.view(output_height, output_width, 4)
rgb_back = back4[..., :3].clamp(0,255)
alpha_back = (back4[..., 3] > 0).float()
z1 = (F.conv2d(d_t * m_t, ker, padding=pad) /
F.conv2d(m_t, ker, padding=pad).clamp(min=1)).squeeze()
rgb = (F.conv2d(c_t * m_t, kerc, padding=pad, groups=3) /
F.conv2d(m_t, ker, padding=pad).repeat(1,3,1,1).clamp(min=1)
).squeeze().permute(1,2,0)
# fill holes where front missed
hole = (alpha == 0) & (alpha_back > 0)
rgb[hole] = rgb_back[hole]
alpha[hole] = 1.0
depth_img[hole] = z_back.view(output_height, output_width)[hole]
# 7) Median-filter _only_ in hole regions
if hole.any():
# prepare for kornia median_blur: [B,C,H,W]
rgb_t = rgb.permute(2,0,1).unsqueeze(0) # [1,3,H,W]
# apply median filter
rgb_med = median_blur(rgb_t, (point_size, point_size))
# back to HWC
rgb_med = rgb_med.squeeze(0).permute(1,2,0)
# merge only at hole locations
rgb[hole] = rgb_med[hole]
# alpha already set to 1.0 for holes
# 8) Pack and return with original script shapes
img = rgb.unsqueeze(0) # [1,H,W,3]
mask_out = alpha # [H,W]
depth4 = depth_img.unsqueeze(0).unsqueeze(-1) # [1,H,W,1]
# 10) Pack outputs
img = rgb.unsqueeze(0) # [1,H,W,3]
mask = fg_mask.float() # [H,W]
depth = z1.unsqueeze(0).unsqueeze(-1) # [1,H,W,1]
if return_inverse_depth:
depth4 = 1.0 / depth4.clamp(min=1e-6)
depth4 = depth4 * mask_out.unsqueeze(0).unsqueeze(-1)
return img, mask_out, depth4
depth = 1.0 / depth.clamp(min=1e-6)
depth *= mask.unsqueeze(0).unsqueeze(-1)
return img, mask, depth
class PointCloudUnion:
"""
@@ -749,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
@@ -759,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",)
@@ -770,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,)
@@ -792,7 +955,7 @@ class CameraTrajectoryNode:
"pointcloud": ("TENSOR",),
},
"optional": {
"initial_matrix": ("MAT_4X4"),
"initial_matrix": ("MAT_4X4",),
}
}
@@ -1205,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,
@@ -1219,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/"]
+9 -2
View File
@@ -42,7 +42,14 @@ def map_grid(
output_horizontal_fov = torch.tensor(output_horizontal_fov, device=grid_torch.device).float()
# Calculate vertical field of view for input and output projections
output_vertical_fov = output_horizontal_fov # Assuming square aspect ratio
# For equirectangular, use 2:1 aspect ratio (vertical FOV = horizontal FOV / 2)
if output_projection == "EQUIRECTANGULAR":
output_vertical_fov = output_horizontal_fov / 2.0
else:
output_vertical_fov = output_horizontal_fov # Assuming square aspect ratio for other projections
# Calculate input vertical FOV based on output grid aspect ratio
# This allows the input's vertical range to adapt to the output dimensions
input_vertical_fov = input_horizontal_fov * (grid_torch.shape[0] / grid_torch.shape[1])
# Normalize the grid for vertical FOV adjustment
@@ -222,8 +229,8 @@ class ReprojectImage:
)
grid_y, grid_x = torch.meshgrid(
torch.linspace(-1, 1, output_width, device=image_tensor.device),
torch.linspace(-1, 1, output_height, device=image_tensor.device),
torch.linspace(-1, 1, output_width, device=image_tensor.device),
indexing="ij"
)
grid_init = torch.stack((grid_x, grid_y), dim=-1)
+299
View File
@@ -0,0 +1,299 @@
import torch
import torch.nn.functional as F
import numpy as np
import os
import sys
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, interpolate_se3
import folder_paths
# Ensure video_depth_anything is on path
_here = os.path.dirname(os.path.abspath(__file__))
# climb up 3 levels: camera-comfyUI → custom_nodes → ComfyUI
COMFYUI_ROOT = os.path.abspath(os.path.join(_here, os.pardir, os.pardir))
# point at metric_depth inside the Video-Depth-Anything clone at the ComfyUI root
video_depth_path = os.path.join(COMFYUI_ROOT, "Video-Depth-Anything", "metric_depth")
# insert at front so it always wins
if video_depth_path not in sys.path:
sys.path.insert(0, video_depth_path)
NO_VIDEO_DEPTH_ANYTHING= False
try:
from video_depth_anything.video_depth import VideoDepthAnything
print("✅ video_depth_anything module loaded successfully.")
except ImportError as e:
NO_VIDEO_DEPTH_ANYTHING = True
print(
f"❌ Could not load video_depth_anything from {video_depth_path!r}: {e}"
)
class VideoCameraMotionSequence:
"""
Takes a sequence of RGB frames and corresponding depth maps,
converts each frame+depth to a pointcloud, interpolates a camera
trajectory to match video length, cleans the pointcloud if needed,
and outputs reprojected images, masks, and depth maps per frame.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
# Sequence of frames: Tensor [T, H, W, 3]
"frames": ("IMAGE", {"shape_hint": [None, None, None, 3]}),
# Sequence of depth maps: Tensor [T, H, W] or [T, H, W, 1]
"depth_seq": ("TENSOR", {"shape_hint": [None, None, None]}),
# Camera trajectory waypoints: Tensor [K, 4, 4]
"trajectory": ("TENSOR", {"shape_hint": [None, 4, 4]}),
# Optional mask sequence: Tensor [T, H, W] or [T, H, W, 1]
"mask_seq": ("MASK", {"shape_hint": [None, None, None], "optional": True}),
# Input projection parameters
"input_projection": (Projection.PROJECTIONS, {}),
"input_horizontal_fov": ("FLOAT", {"default": 90.0}),
"depth_scale": ("FLOAT", {"default": 1.0}),
"invert_depth": ("BOOLEAN", {"default": False}),
# Output projection parameters
"output_projection": (Projection.PROJECTIONS, {}),
"output_horizontal_fov": ("FLOAT", {"default": 90.0}),
"output_width": ("INT", {"default": 512, "min": 1}),
"output_height": ("INT", {"default": 512, "min": 1}),
"point_size": ("INT", {"default": 1, "min": 1}),
# Cleaning parameters
"voxel_size": ("FLOAT", {"default": 1.0, "min": 1e-3}),
"min_points_per_voxel": ("INT", {"default": 3, "min": 1}),
}
}
RETURN_TYPES = ("IMAGE", "MASK", "TENSOR")
RETURN_NAMES = ("video_frames", "mask_frames", "depths")
FUNCTION = "process_sequence"
CATEGORY = "Camera/Video"
def process_sequence(
self,
frames: torch.Tensor,
depth_seq: torch.Tensor,
trajectory: torch.Tensor,
input_projection: str,
input_horizontal_fov: float,
depth_scale: float,
invert_depth: bool,
output_projection: str,
output_horizontal_fov: float,
output_width: int,
output_height: int,
point_size: int,
voxel_size: float,
min_points_per_voxel: int,
mask_seq: torch.Tensor = None,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
# frames: [T, H, W, 3]
# depth_seq: [T, H, W] or [T, H, W, 1]
T, H, W, _ = frames.shape
# Interpolate trajectory to match T (SE(3): quaternion SLERP on R, lerp on t)
interp_traj = interpolate_se3(trajectory, T)
out_frames = []
out_masks = []
out_depths = []
# Add tqdm progress bar for the sequence
# If mask_seq is a single mask [H, W] or [H, W, 1], repeat it for all frames
if mask_seq is not None:
if mask_seq.dim() == 2 or (mask_seq.dim() == 3 and mask_seq.shape[0] == 1):
mask_seq = mask_seq.unsqueeze(0) if mask_seq.dim() == 2 else mask_seq
mask_seq = mask_seq.repeat(T, 1, 1, 1) if mask_seq.dim() == 4 else mask_seq.repeat(T, 1, 1)
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; 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)
# to pointcloud
pc, = DepthToPointCloud().depth_to_pointcloud(
image=frame.permute(2, 0, 1),
input_projection=input_projection,
input_horizontal_fov=input_horizontal_fov,
depth_scale=depth_scale,
invert_depth=invert_depth,
depthmap=depth,
mask=mask,
)
# optional cleaning
if min_points_per_voxel > 1:
pc, = PointCloudCleaner().clean_pointcloud(
pointcloud=pc,
width=output_width,
height=output_height,
voxel_size=voxel_size,
min_points_per_voxel=min_points_per_voxel,
)
# transform and project
pc_t, = TransformPointCloud().transform_pointcloud(pc, pose)
img_t, mask_t, depth_t = ProjectPointCloud().project_pointcloud(
pointcloud=pc_t,
output_projection=output_projection,
output_horizontal_fov=output_horizontal_fov,
output_width=output_width,
output_height=output_height,
point_size=point_size,
)
out_frames.append(img_t[0])
out_masks.append(mask_t)
out_depths.append(depth_t)
return (
torch.stack(out_frames, dim=0), # [T, 3, H, W]
torch.stack(out_masks, dim=0), # [T, H, W]
torch.stack(out_depths, dim=0), # [T, H, W]
)
class DepthFramesToVideo:
"""
Converts a sequence of depth maps into video frame tensors for saving.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
return {
"required": {
"depth_seq": ("TENSOR", {"shape_hint": [None, None, None]}),
"mask_seq": ("MASK", {"shape_hint": [None, None, None]}),
"normalize": ("BOOLEAN", {"default": True}),
"invert_depth": ("BOOLEAN", {"default": False}),
}
}
RETURN_TYPES = ("TENSOR", "IMAGE")
RETURN_NAMES = ("video_frames", "depth_video")
FUNCTION = "depth_to_video_frames"
CATEGORY = "Camera/Video"
def depth_to_video_frames(
self,
depth_seq: torch.Tensor,
normalize: bool,
invert_depth: bool,
mask_seq: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
ds = depth_seq.clone().squeeze()
if ds.dim() == 2:
ds = ds.unsqueeze(0) # [H, W] -> [1, H, W]
if ds.dim() != 3:
raise ValueError(f"Expected ds to be 3D [T, H, W], got shape {ds.shape}")
if invert_depth:
ds= 1.0 / (ds + 1e-8) # Avoid division by zero
if normalize:
# Mask: only normalize where depth > 0
mask = mask_seq>0.5
if mask.any():
#percentile first 10 percent min
# sample
minv = ds[mask]
# sample 10000 and find 10% quantile
if minv.numel() > 10000:
minv = minv[torch.randperm(minv.numel())[:10000]]
minv = minv.quantile(0.2)
minv = minv if minv > 0.1 else 0.1 # Avoid division by zero
#percentile last 10 percent max
maxv = ds[mask]
if maxv.numel() > 10000:
maxv = maxv[torch.randperm(maxv.numel())[:10000]]
maxv = maxv.quantile(0.98)
maxv = maxv if maxv < 100 else 100
print(f"Normalizing depth: min={minv}, max={maxv}")
ds_norm = (ds - minv) / (maxv - minv + 1e-8)
ds = ds_norm.clamp(0, 1) # torch.where(mask, ds_norm, ds) # Only normalize valid values
else:
print("Warning: No valid depth values for normalization.")
# expand to 3 channels: [T, H, W] -> [T, 3, H, W]
raw = depth_seq.clone().squeeze()
ds_u8 = (ds * 255.0).round().to(torch.uint8)
raw_u8 = (raw.clamp(0, 255)).to(torch.uint8) # if raw is already in a displayable range
# expand to 3 channels and permute to HWC
ds_color = ds_u8.unsqueeze(1).repeat(1, 3, 1, 1).permute(0, 2, 3, 1)
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.
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
# model files (.pth) in input directory
model_dir = os.path.join(os.getcwd(), "models", "checkpoints")
os.makedirs(model_dir, exist_ok=True)
files = [f for f in os.listdir(model_dir) if f.lower().endswith(('.pth', '.ckpt', '.safetensors'))]
return {
"required": {
"frames": ("IMAGE", {"shape_hint": [None, None, None, 3]}),
"model_checkpoint": (files, {"file_chooser": True}),
"input_size": ("INT", {"default": 518, "min": 64, "max": 2048}),
"max_fps": ("INT", {"default": 60, "min": 1}),
}
}
RETURN_TYPES = ("TENSOR", "FLOAT")
RETURN_NAMES = ("metric_depths", "fps")
FUNCTION = "estimate_metric_depth"
CATEGORY = "Camera/Video"
def estimate_metric_depth(
self,
frames: torch.Tensor,
model_checkpoint: str,
input_size: int,
max_fps: int,
) -> Tuple[torch.Tensor, float]:
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)
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))
# Register nodes
if NO_VIDEO_DEPTH_ANYTHING:
NODE_CLASS_MAPPINGS = {}
else:
NODE_CLASS_MAPPINGS = {
"VideoCameraMotionSequence": VideoCameraMotionSequence,
"VideoMetricDepthEstimate": VideoMetricDepthEstimate,
"DepthFramesToVideo": DepthFramesToVideo,
}
+289 -1
View File
@@ -1 +1,289 @@
{"id":"dd56c0bf-7405-406e-924f-42b2feacb73f","revision":0,"last_node_id":6,"last_link_id":4,"nodes":[{"id":1,"type":"TransformToMatrix","pos":[-337.9580993652344,1500.5211181640625],"size":[315,154],"flags":{},"order":0,"mode":0,"inputs":[],"outputs":[{"localized_name":"MAT_4X4","name":"MAT_4X4","type":"MAT_4X4","links":[1]}],"properties":{"Node name for S&R":"TransformToMatrix"},"widgets_values":[0,0,0,0,0]},{"id":2,"type":"CameraMotion","pos":[130.0218505859375,1516.7567138671875],"size":[367.79998779296875,218],"flags":{},"order":3,"mode":0,"inputs":[{"localized_name":"pointcloud","name":"pointcloud","type":"TENSOR","link":4},{"localized_name":"initial_matrix","name":"initial_matrix","type":"MAT_4X4","link":1},{"localized_name":"final_matrix","name":"final_matrix","type":"MAT_4X4","link":2}],"outputs":[{"localized_name":"IMAGE","name":"IMAGE","type":"IMAGE","links":[3]}],"properties":{"Node name for S&R":"CameraMotion"},"widgets_values":[24,"PINHOLE",90,1024,1024,2]},{"id":3,"type":"SaveWEBM","pos":[606.5880737304688,1515.5966796875],"size":[315,437],"flags":{},"order":4,"mode":0,"inputs":[{"localized_name":"images","name":"images","type":"IMAGE","link":3}],"outputs":[],"properties":{},"widgets_values":["ComfyUI","vp9",10.000000000000002,32]},{"id":5,"type":"TransformToMatrix","pos":[-262.9134521484375,1740.8876953125],"size":[315,154],"flags":{},"order":1,"mode":0,"inputs":[],"outputs":[{"localized_name":"MAT_4X4","name":"MAT_4X4","type":"MAT_4X4","links":[2]}],"properties":{"Node name for S&R":"TransformToMatrix"},"widgets_values":[0.10000000000000002,0,0,0,0]},{"id":6,"type":"LoadPointCloud","pos":[-357.24761962890625,1294.427734375],"size":[315,58],"flags":{},"order":2,"mode":0,"inputs":[],"outputs":[{"localized_name":"TENSOR","name":"TENSOR","type":"TENSOR","links":[4]}],"properties":{"Node name for S&R":"LoadPointCloud"},"widgets_values":["ComfyUIPointCloud_00001.ply"]}],"links":[[1,1,0,2,1,"MAT_4X4"],[2,5,0,2,2,"MAT_4X4"],[3,2,0,3,0,"IMAGE"],[4,6,0,2,0,"TENSOR"]],"groups":[],"config":{},"extra":{"ds":{"scale":1.351305709310409,"offset":[-187.6257577580669,-1567.2143321744395]}},"version":0.4}
{
"id": "dd56c0bf-7405-406e-924f-42b2feacb73f",
"revision": 0,
"last_node_id": 8,
"last_link_id": 9,
"nodes": [
{
"id": 3,
"type": "SaveWEBM",
"pos": [
606.5880737304688,
1515.5966796875
],
"size": [
315,
437
],
"flags": {},
"order": 5,
"mode": 0,
"inputs": [
{
"name": "images",
"type": "IMAGE",
"link": 5
}
],
"outputs": [],
"properties": {},
"widgets_values": [
"ComfyUI",
"vp9",
10.000000000000002,
32
]
},
{
"id": 7,
"type": "CameraMotionNode",
"pos": [
176.07933044433594,
1504.22705078125
],
"size": [
278.75,
270
],
"flags": {},
"order": 4,
"mode": 0,
"inputs": [
{
"name": "pointcloud",
"type": "TENSOR",
"link": 6
},
{
"name": "trajectory",
"type": "TENSOR",
"link": 9
}
],
"outputs": [
{
"name": "motion_frames",
"type": "IMAGE",
"links": [
5
]
},
{
"name": "mask_frames",
"type": "MASK",
"links": null
}
],
"properties": {
"Node name for S&R": "CameraMotionNode"
},
"widgets_values": [
10,
"PINHOLE",
90,
512,
512,
1,
0,
false,
false
]
},
{
"id": 6,
"type": "LoadPointCloud",
"pos": [
-357.24761962890625,
1294.427734375
],
"size": [
315,
58
],
"flags": {},
"order": 0,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "loaded pointcloud",
"type": "TENSOR",
"links": [
6
]
}
],
"properties": {
"Node name for S&R": "LoadPointCloud"
},
"widgets_values": [
"ComfyUIPointCloud_00001.ply"
]
},
{
"id": 1,
"type": "TransformToMatrix",
"pos": [
-537.2319946289062,
1491.40380859375
],
"size": [
315,
154
],
"flags": {},
"order": 1,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "transformation matrix",
"type": "MAT_4X4",
"links": [
7
]
}
],
"properties": {
"Node name for S&R": "TransformToMatrix"
},
"widgets_values": [
0,
0,
0,
0,
0
]
},
{
"id": 5,
"type": "TransformToMatrix",
"pos": [
-531.216796875,
1701.1632080078125
],
"size": [
315,
154
],
"flags": {},
"order": 2,
"mode": 0,
"inputs": [],
"outputs": [
{
"name": "transformation matrix",
"type": "MAT_4X4",
"links": [
8
]
}
],
"properties": {
"Node name for S&R": "TransformToMatrix"
},
"widgets_values": [
0.10000000000000002,
0,
0,
0,
0
]
},
{
"id": 8,
"type": "CameraInterpolationNode",
"pos": [
-121.91971588134766,
1597.3514404296875
],
"size": [
200.21640014648438,
46
],
"flags": {},
"order": 3,
"mode": 0,
"inputs": [
{
"name": "initial_matrix",
"type": "MAT_4X4",
"link": 7
},
{
"name": "final_matrix",
"type": "MAT_4X4",
"link": 8
}
],
"outputs": [
{
"name": "trajectory",
"type": "TENSOR",
"links": [
9
]
}
],
"properties": {
"Node name for S&R": "CameraInterpolationNode"
},
"widgets_values": []
}
],
"links": [
[
5,
7,
0,
3,
0,
"IMAGE"
],
[
6,
6,
0,
7,
0,
"TENSOR"
],
[
7,
1,
0,
8,
0,
"MAT_4X4"
],
[
8,
5,
0,
8,
1,
"MAT_4X4"
],
[
9,
8,
0,
7,
1,
"TENSOR"
]
],
"groups": [],
"config": {},
"extra": {
"ds": {
"scale": 1.015255979947716,
"offset": [
636.7531305750655,
-1186.3099424359816
]
},
"frontendVersion": "1.21.7"
},
"version": 0.4
}
File diff suppressed because it is too large Load Diff
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,
}