Merge pull request #21 from Alexankharin/video-world-experiments
Publish camera-comfyUI to the ComfyUI Registry (+ 4D world nodes)
This commit is contained in:
@@ -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
|
||||
@@ -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 }}
|
||||
@@ -0,0 +1,3 @@
|
||||
[submodule "submodules/ml-sharpt"]
|
||||
path = submodules/ml-sharpt
|
||||
url = https://github.com/apple/ml-sharp
|
||||
+1227
File diff suppressed because it is too large
Load Diff
+2499
File diff suppressed because it is too large
Load Diff
@@ -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.
|
||||
@@ -14,6 +14,7 @@
|
||||
* [Installation](#installation)
|
||||
* [Node Categories](#node-categories)
|
||||
* [Node Reference](#node-reference)
|
||||
* [Video → 4D World](#video--4d-world)
|
||||
* [Workflows](#workflows)
|
||||
* [Example Workflows](#example-workflows)
|
||||
* [Contributing](#contributing)
|
||||
@@ -35,6 +36,14 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
|
||||
|
||||
## Installation
|
||||
|
||||
### Option A — ComfyUI Manager (recommended)
|
||||
|
||||
The node pack is published to the [ComfyUI Registry](https://registry.comfy.org) as **`camera-comfyui`** (publisher `alexk`). In ComfyUI, open **Manager → Custom Nodes Manager**, search for **camera-comfyUI**, and click **Install**, then restart ComfyUI. The registry package bundles the SHARP submodule and installs the base Python requirements automatically; optional CUDA-specific extras (`gsplat`, `vggt`) still follow the manual steps below.
|
||||
|
||||
> **Maintainers:** releases are automated — bumping `version` in `pyproject.toml` on `main` triggers `.github/workflows/publish_action.yml`, which publishes the new version to the registry (requires the `REGISTRY_ACCESS_TOKEN` repo secret).
|
||||
|
||||
### Option B — Manual install (git)
|
||||
|
||||
1. **Clone** into your ComfyUI custom nodes folder:
|
||||
|
||||
```bash
|
||||
@@ -55,6 +64,13 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
|
||||
|
||||
* *Optional:* `open3d` for GUI point cloud tools.
|
||||
|
||||
**Optional dependencies** (only needed for specific nodes):
|
||||
|
||||
* **gsplat** — CUDA-accelerated Gaussian splat rasterizer. Required by `SplatPolish` and used as the fast render backend for `RenderSplat` / `RenderSplats4D*`. Needs a CUDA GPU and a matching PyTorch build: `pip install gsplat`.
|
||||
* **vggt** — camera pose + depth estimation (`VideoPoseEstimator`). Install with `pip install vggt` (or `pip install git+https://github.com/facebookresearch/vggt.git`), or clone [facebookresearch/vggt](https://github.com/facebookresearch/vggt) as a sibling folder in your ComfyUI root. The `facebook/VGGT-1B` weights (~5 GB) download via `huggingface_hub` on first use.
|
||||
* **CoTracker3** — point tracking for `EstimateTracks`. No manual install: it is fetched automatically via `torch.hub` on first use.
|
||||
* **SHARP** — image→splat prediction (`ImageToSplat`, `FisheyeToGaussian`, `VideoToFusedSplats`, `SplatTrajectoryEnricher`). Ships as the existing git submodule at `submodules/ml-sharpt` ([apple/ml-sharp](https://github.com/apple/ml-sharp)) — run `git submodule update --init` after cloning.
|
||||
|
||||
4. **Additional Nodes** (for certain workflows):
|
||||
|
||||
* Clone the following repositories directly into your `custom_nodes` folder:
|
||||
@@ -94,13 +110,33 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
|
||||
* ### Point Cloud Nodes
|
||||
|
||||
* `DepthToPointCloud`, `TransformPointCloud`, `ProjectPointCloud`, `PointCloudUnion`
|
||||
* `PointCloudCleaner`, `LoadPointCloud`, `SavePointCloud`, `ProjectAndClean`
|
||||
* `PointCloudCleaner`, `LoadPointCloud`, `SavePointCloud`, `ProjectAndClean`, `DepthEdgeFilter`
|
||||
|
||||
* ### Trajectory Nodes
|
||||
|
||||
* `CameraMotionNode`, `CameraInterpolationNode`, `CameraTrajectoryNode`
|
||||
* `SaveTrajectory`, `LoadTrajectory`, `PointcloudTrajectoryEnricher`
|
||||
|
||||
* ### Gaussian Splat Nodes
|
||||
|
||||
* `LoadPlySplat`, `SavePlySplat`, `ImageToSplat`, `FisheyeToGaussian`
|
||||
* `RotateSplats`, `MergeSplats`, `FuseSplats`, `RenderSplat`
|
||||
* `VideoToFusedSplats`, `SplatPolish`
|
||||
|
||||
* ### 4D Gaussian Splat Nodes
|
||||
|
||||
* `MotionMaskFromDepth`, `EstimateTracks`, `TracksToTrajectories`, `SplitSplatsByMask`
|
||||
* `BuildSplats4D`, `RenderSplats4DFrame`, `RenderSplats4DVideo`
|
||||
* `SaveSplats4D`, `LoadSplats4D`
|
||||
|
||||
* ### Pose Nodes
|
||||
|
||||
* `VideoPoseEstimator`, `TrajectoryInvert`, `TrajectoryCompose`
|
||||
|
||||
* ### World Nodes
|
||||
|
||||
* `DepthScaleAnchor`, `SplatTrajectoryEnricher`, `SphereSplatSeed`
|
||||
|
||||
---
|
||||
|
||||
## Node Reference
|
||||
@@ -130,6 +166,51 @@ A collection of ComfyUI custom nodes to handle diverse camera projections (pinho
|
||||
| `VideoCameraMotionSequence` | Processes video frames and depth maps along a camera trajectory, generating reprojected outputs. |
|
||||
| `DepthFramesToVideo` | Converts a sequence of depth maps into video frame tensors for saving. |
|
||||
| `VideoMetricDepthEstimate` | Estimates metric depth for a sequence of frames using VideoDepthAnything. |
|
||||
| `DepthEdgeFilter` | Detects "flying pixel" depth discontinuities and outputs a validity mask (1.0 = valid). |
|
||||
| `LoadPlySplat` | Loads a 3D Gaussian Splatting `.ply` file into a `GSPLAT` object. |
|
||||
| `SavePlySplat` | Saves a `GSPLAT` to the ComfyUI output directory as a `.ply` file. |
|
||||
| `ImageToSplat` | Predicts Gaussian splats from a single image using SHARP. |
|
||||
| `FisheyeToGaussian` | Reprojects a fisheye view to multiple pinhole angles, predicts splats, rotates and merges them. |
|
||||
| `RotateSplats` | Applies a 4×4 transform matrix to a splat cloud. |
|
||||
| `MergeSplats` | Concatenates two `GSPLAT` objects into one. |
|
||||
| `FuseSplats` | Fuses two splat clouds with weighted voxel merging (keep/discard/average/smart modes). |
|
||||
| `RenderSplat` | Renders a splat cloud from a camera pose into an image + mask. |
|
||||
| `VideoToFusedSplats` | Runs SHARP on video keyframes, scale-aligns to metric depth, filters dynamic pixels, and fuses all keyframes into one world-frame splat cloud. |
|
||||
| `SplatPolish` | Optimizes a world-frame splat cloud against posed video frames (L1 + D-SSIM) using gsplat's differentiable rasterizer. |
|
||||
| `MotionMaskFromDepth` | Detects dynamic pixels from a depth+pose sequence (1.0 = moving). |
|
||||
| `EstimateTracks` | Runs CoTracker3 on a video; returns tracks `[T,N,2]` (pixels) and visibility `[T,N]`. |
|
||||
| `TracksToTrajectories` | Unprojects 2D tracks with depth and camera poses into world-space 3D trajectories `[T,M,3]`. |
|
||||
| `SplitSplatsByMask` | Projects splat centers into a 2D mask and splits the cloud into inside/outside parts. |
|
||||
| `BuildSplats4D` | Builds a 4D splat scene: each canonical splat follows a kNN blend of track control-point motions. |
|
||||
| `RenderSplats4DFrame` | Evaluates the 4D scene at a single time value and renders it from a given camera. |
|
||||
| `RenderSplats4DVideo` | Interpolates the camera path, sweeps time from start to end, and renders each frame. |
|
||||
| `SaveSplats4D` | Saves a `GSPLAT4D` scene as an `.npz` archive (plus optional per-frame PLYs). |
|
||||
| `LoadSplats4D` | Loads a `GSPLAT4D` scene from an `.npz` archive. |
|
||||
| `VideoPoseEstimator` | VGGT-based per-frame camera poses `[T,4,4]`, depth maps, FOV and depth confidence from a video clip. |
|
||||
| `TrajectoryInvert` | Inverts each 4×4 pose (world-to-camera ↔ camera-to-world). |
|
||||
| `TrajectoryCompose` | Per-frame matrix product `A @ B`; a single 4×4 input broadcasts over the other. |
|
||||
| `DepthScaleAnchor` | Robustly aligns a depth map to a reference depth via disparity-domain scale(+shift). |
|
||||
| `SplatTrajectoryEnricher` | Expands a splat world along a trajectory: render, outpaint holes with Flux, lift with SHARP, scale-align, smart-stitch. |
|
||||
| `SphereSplatSeed` | Converts an equirectangular panorama into a Gaussian sphere seeding a 360° world. |
|
||||
|
||||
---
|
||||
|
||||
## Video → 4D World
|
||||
|
||||
Turn a monocular video into a navigable 4D (3D + time) Gaussian splat scene and re-render it from any novel camera trajectory. The reference workflow is **`workflows/video_to_4d_world.json`**; the stages are:
|
||||
|
||||
1. **Pose & depth (VGGT)** — `VideoPoseEstimator` estimates per-frame world-to-camera poses `[T,4,4]`, depth maps, FOV and depth confidence from the input frames. Since the depth maps are Z-depths, run `ZDepthToRayDepthNode` before any node that expects ray depth (see caveats below). `DepthEdgeFilter` can additionally mask out flying pixels at depth discontinuities.
|
||||
2. **Motion masking** — `MotionMaskFromDepth` warps depth between frames using the estimated poses and flags pixels whose residual is too large as dynamic (moving objects vs. static background).
|
||||
3. **Static splat fusion + polish** — `VideoToFusedSplats` runs SHARP on keyframes, keeps only static pixels (via the motion mask), scale-aligns each keyframe to metric depth, transforms splats into the world frame and fuses them incrementally. `SplatPolish` then fine-tunes the fused cloud photometrically against the posed video frames.
|
||||
4. **Tracked dynamic 4D Gaussians** — `EstimateTracks` (CoTracker3) tracks a dense point grid across the video; `TracksToTrajectories` lifts the tracks to world-space 3D using depth + poses; `SplitSplatsByMask` separates dynamic splats from the static background; `BuildSplats4D` binds the dynamic canonical splats to track control points via kNN blending, producing a `GSPLAT4D` scene.
|
||||
5. **Render a novel trajectory** — build any new camera path (e.g. `CameraInterpolationNode`, `TrajectoryCompose` to retarget relative to a source pose) and render with `RenderSplats4DVideo` (or single frames with `RenderSplats4DFrame`). Save/reload scenes with `SaveSplats4D` / `LoadSplats4D`.
|
||||
|
||||
### Caveats
|
||||
|
||||
* **Z-depth vs ray depth**: depth estimators (including `VideoPoseEstimator`) output Z-depth; point-cloud and splat lifting nodes expect ray depth. Insert `ZDepthToRayDepthNode` where needed, or geometry will bow at wide FOVs.
|
||||
* **`SplatPolish` requires gsplat + CUDA**: without them it can fall back to the differentiable torch renderer at reduced resolution, which is extremely slow (minutes per 100 iterations).
|
||||
* **`EstimateTracks` downloads CoTracker3 via `torch.hub` on first use** — expect a one-time download and allow network access.
|
||||
* **`VideoPoseEstimator` downloads `facebook/VGGT-1B` (~5 GB)** on first use via `huggingface_hub`.
|
||||
|
||||
---
|
||||
|
||||
@@ -150,6 +231,8 @@ A set of JSON workflows illustrating typical use cases. Each workflow lives in `
|
||||
| **Pointcloud\_walker.json** | GUI‐based camera control via Open3D |
|
||||
| **sbs180\_workflow.json** | Generate stereo (side-by-side) wide-angle/fisheye/equirectangular stereo pairs from a high-res input |
|
||||
| **video_camera.json** | Camera trajectory movement workflow using `wan-vace` for video inpainting. |
|
||||
| **video_to_4d_world\.json** | Video → 4D world: VGGT poses/depth → motion masking → fused static splats + polish → tracked dynamic 4D Gaussians → novel-trajectory render. |
|
||||
| **video_to_4d_walkable_world\.json** | Video → 4D WALKABLE world (test-friendly defaults): polished static splats enriched along a walk trajectory (`SplatTrajectoryEnricher`, Flux outpaint + SHARP) → 4D scene → walk-through render + `.ply`/`.npz` exports for free walking in external 3DGS viewers. |
|
||||
|
||||
---
|
||||
|
||||
@@ -268,9 +351,12 @@ Contributions welcome! Please open issues or PRs to add features, improve docs,
|
||||
* [x] Add processing to pointcloud or depthmap to remove outlier and lonely points at depth borders.
|
||||
* [x] Use built-in comfyUI mask type an image.
|
||||
* [x] Unite nodes into groups to simplify workflows.
|
||||
* [ ] Create a single workflow for view synthesis.
|
||||
* [x] Create a single workflow for view synthesis (`video_to_4d_world.json`).
|
||||
* [x] Implement easier and more flexible camera control - more complex camera movements with more than 2 points.
|
||||
* [x] Add more examples and documentation for each node.
|
||||
* [x] Add pointcloud union
|
||||
* [x] Fix imports for renamed folders (e.g., inpainting_flux)
|
||||
* [x] Integrate camera movement pipeline with video models (e.g., wan2.1) for smooth, high-quality inpainting along camera trajectories.
|
||||
* [ ] Compressed export format for 4D scenes (current `.npz` stores raw tensors).
|
||||
* [ ] SAM2-based refinement of motion masks (current masks come from depth-warp residuals only).
|
||||
* [ ] Fisheye/equirectangular rendering through gsplat (e.g., via cubemap render + reprojection); the fast CUDA path is currently pinhole-only.
|
||||
|
||||
+23
-2
@@ -4,6 +4,27 @@ from .metric_depth_nodes import NODE_CLASS_MAPPINGS as NCM3
|
||||
from .flux_fisheye_filling_nodes import NODE_CLASS_MAPPINGS as NCM4
|
||||
from .complex_nodes import NODE_CLASS_MAPPINGS as NCM5
|
||||
from .video_nodes import NODE_CLASS_MAPPINGS as NCM6
|
||||
NODE_CLASS_MAPPINGS = {**NCM1, **NCM2, **NCM3, **NCM4, **NCM5, **NCM6}
|
||||
from .GS_nodes import NODE_CLASS_MAPPINGS as NCM7
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS"]
|
||||
# Optional node packs: a missing/broken optional dependency must never kill the
|
||||
# whole extension (mirrors how video_nodes degrades when video_depth_anything
|
||||
# is unavailable).
|
||||
try:
|
||||
from .GS4D_nodes import NODE_CLASS_MAPPINGS as NCM8
|
||||
except Exception as _exc:
|
||||
print(f"[camera-comfyUI] Warning: GS4D_nodes could not be loaded, 4D splat nodes disabled: {_exc}")
|
||||
NCM8 = {}
|
||||
try:
|
||||
from .pose_nodes import NODE_CLASS_MAPPINGS as NCM9
|
||||
except Exception as _exc:
|
||||
print(f"[camera-comfyUI] Warning: pose_nodes could not be loaded, pose estimation nodes disabled: {_exc}")
|
||||
NCM9 = {}
|
||||
try:
|
||||
from .world_nodes import NODE_CLASS_MAPPINGS as NCM10
|
||||
except Exception as _exc:
|
||||
print(f"[camera-comfyUI] Warning: world_nodes could not be loaded, world-building nodes disabled: {_exc}")
|
||||
NCM10 = {}
|
||||
|
||||
NODE_CLASS_MAPPINGS = {**NCM1, **NCM2, **NCM3, **NCM4, **NCM5, **NCM6, **NCM7, **NCM8, **NCM9, **NCM10}
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS"]
|
||||
|
||||
@@ -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.*
|
||||
@@ -0,0 +1,484 @@
|
||||
"""Standalone CPU smoke test for the 4D-world node stack (no ComfyUI, no CUDA,
|
||||
no model downloads).
|
||||
|
||||
Run with:
|
||||
python notebooks/smoke_test_4d.py
|
||||
|
||||
Stubs `folder_paths` via sys.modules injection so the repo modules import
|
||||
outside the ComfyUI runtime, then functionally exercises the NEW code paths
|
||||
with small synthetic data:
|
||||
|
||||
1. interpolate_se3 (pointcloud_nodes, contract C1)
|
||||
2. render_gaussians (GS_nodes, contract C2) shapes + empty case
|
||||
3. render_gaussians fast anisotropic footprint
|
||||
4. GaussianSplats4D.at_time (GS4D_nodes, contract C3)
|
||||
5. BuildSplats4D kNN track binding
|
||||
6. SplitSplatsByMask
|
||||
7. MotionMaskFromDepth
|
||||
8. align_depth_scale (world_nodes, contract C4) + DepthEdgeFilter
|
||||
9. FuseSplats weighted voxel fusion
|
||||
10. SphereSplatSeed pano -> splat sphere -> render round-trip
|
||||
"""
|
||||
|
||||
import math
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import traceback
|
||||
import types
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Environment setup: repo on sys.path + folder_paths stub (before repo imports)
|
||||
# --------------------------------------------------------------------------- #
|
||||
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
if REPO_ROOT not in sys.path:
|
||||
sys.path.insert(0, REPO_ROOT)
|
||||
|
||||
_TMP_DIR = tempfile.mkdtemp(prefix="smoke_test_4d_")
|
||||
|
||||
|
||||
def _stub_get_save_image_path(filename_prefix, output_dir, *args, **kwargs):
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
return output_dir, filename_prefix, 0, "", filename_prefix
|
||||
|
||||
|
||||
_fp_stub = types.ModuleType("folder_paths")
|
||||
_fp_stub.get_input_directory = lambda: _TMP_DIR
|
||||
_fp_stub.get_output_directory = lambda: _TMP_DIR
|
||||
_fp_stub.get_temp_directory = lambda: _TMP_DIR
|
||||
_fp_stub.get_save_image_path = _stub_get_save_image_path
|
||||
_fp_stub.get_annotated_filepath = lambda name: os.path.join(_TMP_DIR, name)
|
||||
_fp_stub.exists_annotated_filepath = lambda name: os.path.exists(os.path.join(_TMP_DIR, name))
|
||||
_fp_stub.get_filename_list = lambda folder: []
|
||||
_fp_stub.models_dir = _TMP_DIR
|
||||
sys.modules["folder_paths"] = _fp_stub
|
||||
|
||||
import numpy as np # noqa: E402
|
||||
import torch # noqa: E402
|
||||
|
||||
import GS_nodes # noqa: E402
|
||||
import GS4D_nodes # noqa: E402
|
||||
import pointcloud_nodes # noqa: E402
|
||||
import world_nodes # noqa: E402
|
||||
|
||||
GaussianSplats = GS_nodes.GaussianSplats
|
||||
|
||||
torch.manual_seed(0)
|
||||
np.random.seed(0)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Helpers
|
||||
# --------------------------------------------------------------------------- #
|
||||
def make_splats(
|
||||
xyz: torch.Tensor,
|
||||
sigma: float = 0.05,
|
||||
color: tuple = None,
|
||||
opacity_logit: float = 4.0,
|
||||
) -> GaussianSplats:
|
||||
"""Isotropic sh_order-0 splats at the given positions."""
|
||||
n = xyz.shape[0]
|
||||
if color is None:
|
||||
rgb = torch.rand(n, 3)
|
||||
else:
|
||||
rgb = torch.tensor(color, dtype=torch.float32).view(1, 3).expand(n, 3)
|
||||
C0 = 0.28209479177387814
|
||||
return GaussianSplats(
|
||||
xyz=xyz.float(),
|
||||
scale=torch.full((n, 3), math.log(sigma)),
|
||||
rotation=torch.tensor([1.0, 0.0, 0.0, 0.0]).view(1, 4).expand(n, 4).contiguous(),
|
||||
opacity=torch.full((n, 1), float(opacity_logit)),
|
||||
f_dc=((rgb - 0.5) / C0).contiguous(),
|
||||
f_rest=torch.zeros(n, 0),
|
||||
sh_order=0,
|
||||
)
|
||||
|
||||
|
||||
def rot_x(deg: float) -> torch.Tensor:
|
||||
a = math.radians(deg)
|
||||
return torch.tensor(
|
||||
[[1, 0, 0], [0, math.cos(a), -math.sin(a)], [0, math.sin(a), math.cos(a)]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
|
||||
def rot_y(deg: float) -> torch.Tensor:
|
||||
a = math.radians(deg)
|
||||
return torch.tensor(
|
||||
[[math.cos(a), 0, math.sin(a)], [0, 1, 0], [-math.sin(a), 0, math.cos(a)]],
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
|
||||
def make_pose(R: torch.Tensor, t) -> torch.Tensor:
|
||||
M = torch.eye(4)
|
||||
M[:3, :3] = R
|
||||
M[:3, 3] = torch.tensor(t, dtype=torch.float32)
|
||||
return M
|
||||
|
||||
|
||||
IDENTITY_4X4 = torch.eye(4)
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Tests
|
||||
# --------------------------------------------------------------------------- #
|
||||
def test_01_interpolate_se3():
|
||||
poses = torch.stack(
|
||||
[
|
||||
make_pose(torch.eye(3), [0.0, 0.0, 0.0]),
|
||||
make_pose(rot_y(90.0), [1.0, 2.0, 3.0]),
|
||||
make_pose(rot_y(90.0) @ rot_x(45.0), [-1.0, 0.0, 2.0]),
|
||||
]
|
||||
)
|
||||
out = pointcloud_nodes.interpolate_se3(poses, 10)
|
||||
assert out.shape == (10, 4, 4), f"shape {tuple(out.shape)}"
|
||||
|
||||
eye = torch.eye(3)
|
||||
for i in range(10):
|
||||
R = out[i, :3, :3]
|
||||
ortho_err = (R @ R.T - eye).abs().max().item()
|
||||
det = torch.det(R).item()
|
||||
assert ortho_err < 1e-4, f"step {i}: R@R.T deviates from I by {ortho_err}"
|
||||
assert abs(det - 1.0) < 1e-4, f"step {i}: det(R)={det}"
|
||||
assert torch.allclose(out[i, 3], torch.tensor([0.0, 0.0, 0.0, 1.0]), atol=1e-6)
|
||||
|
||||
assert (out[0] - poses[0]).abs().max().item() < 1e-4, "start pose mismatch"
|
||||
assert (out[-1] - poses[-1]).abs().max().item() < 1e-4, "end pose mismatch"
|
||||
|
||||
# K == 1 repeats.
|
||||
rep = pointcloud_nodes.interpolate_se3(poses[:1], 5)
|
||||
assert rep.shape == (5, 4, 4)
|
||||
assert (rep - poses[0]).abs().max().item() < 1e-6
|
||||
|
||||
|
||||
def test_02_render_gaussians_shapes_and_empty():
|
||||
n, H, W = 200, 48, 64
|
||||
xyz = torch.stack(
|
||||
[
|
||||
torch.rand(n) * 2.0 - 1.0,
|
||||
torch.rand(n) * 2.0 - 1.0,
|
||||
torch.rand(n) * 3.0 + 2.0,
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
splats = make_splats(xyz, sigma=0.05)
|
||||
|
||||
for projection, fov in (("PINHOLE", 90.0), ("EQUIRECTANGULAR", 360.0)):
|
||||
image, mask, disparity = GS_nodes.render_gaussians(
|
||||
splats, IDENTITY_4X4, projection, fov, W, H,
|
||||
render_mode="fast", device="cpu",
|
||||
)
|
||||
assert image.shape == (1, H, W, 3), f"{projection} image {tuple(image.shape)}"
|
||||
assert mask.shape == (H, W), f"{projection} mask {tuple(mask.shape)}"
|
||||
assert disparity.shape == (1, H, W, 1), f"{projection} disparity {tuple(disparity.shape)}"
|
||||
assert torch.isfinite(image).all() and torch.isfinite(disparity).all()
|
||||
assert float(mask.min()) >= 0.0 and float(mask.max()) <= 1.0 + 1e-6
|
||||
assert float(mask.sum()) > 0.0, f"{projection}: nothing rendered"
|
||||
|
||||
# Empty case: every splat strictly behind a pinhole camera (known past bug:
|
||||
# early return used to yield only 2 outputs).
|
||||
behind = make_splats(xyz * torch.tensor([1.0, 1.0, -1.0]), sigma=0.05)
|
||||
result = GS_nodes.render_gaussians(
|
||||
behind, IDENTITY_4X4, "PINHOLE", 90.0, W, H,
|
||||
render_mode="fast", device="cpu",
|
||||
)
|
||||
assert isinstance(result, tuple) and len(result) == 3, f"empty render returned {len(result)} outputs"
|
||||
image, mask, disparity = result
|
||||
assert image.shape == (1, H, W, 3)
|
||||
assert mask.shape == (H, W)
|
||||
assert disparity.shape == (1, H, W, 1)
|
||||
assert float(mask.sum()) == 0.0
|
||||
|
||||
|
||||
def test_03_fast_mode_anisotropy():
|
||||
H = W = 128
|
||||
ang = math.radians(45.0) / 2.0
|
||||
splats = GaussianSplats(
|
||||
xyz=torch.tensor([[0.0, 0.0, 3.0]]),
|
||||
scale=torch.log(torch.tensor([[0.5, 0.01, 0.01]])),
|
||||
rotation=torch.tensor([[math.cos(ang), 0.0, 0.0, math.sin(ang)]]), # 45 deg about +z
|
||||
opacity=torch.tensor([[6.0]]),
|
||||
f_dc=torch.zeros(1, 3),
|
||||
f_rest=torch.zeros(1, 0),
|
||||
sh_order=0,
|
||||
)
|
||||
image, mask, disparity = GS_nodes.render_gaussians(
|
||||
splats, IDENTITY_4X4, "PINHOLE", 60.0, W, H,
|
||||
render_mode="fast", max_radius=64, device="cpu",
|
||||
)
|
||||
assert float(mask.sum()) > 0.0, "elongated splat rendered nothing"
|
||||
|
||||
# Alpha-weighted pixel covariance of the footprint.
|
||||
ys, xs = torch.meshgrid(
|
||||
torch.arange(H, dtype=torch.float32), torch.arange(W, dtype=torch.float32),
|
||||
indexing="ij",
|
||||
)
|
||||
w = mask.flatten()
|
||||
wsum = w.sum()
|
||||
mx = (w * xs.flatten()).sum() / wsum
|
||||
my = (w * ys.flatten()).sum() / wsum
|
||||
dx = xs.flatten() - mx
|
||||
dy = ys.flatten() - my
|
||||
cxx = (w * dx * dx).sum() / wsum
|
||||
cyy = (w * dy * dy).sum() / wsum
|
||||
cxy = (w * dx * dy).sum() / wsum
|
||||
cov = torch.tensor([[cxx, cxy], [cxy, cyy]])
|
||||
evals, evecs = torch.linalg.eigh(cov)
|
||||
ratio = float(evals[1] / evals[0].clamp(min=1e-8))
|
||||
assert ratio > 2.0, f"footprint not elongated: eigenvalue ratio {ratio:.2f}"
|
||||
|
||||
# Principal axis should be near 45 degrees (rotation honored).
|
||||
major = evecs[:, 1]
|
||||
angle = math.degrees(math.atan2(float(major[1]), float(major[0]))) % 180.0
|
||||
assert abs(angle - 45.0) < 15.0, f"major axis at {angle:.1f} deg, expected ~45"
|
||||
|
||||
|
||||
def test_04_at_time():
|
||||
T = 5
|
||||
canonical = make_splats(torch.tensor([[0.0, 0.0, 2.0], [0.0, 1.0, 3.0]]))
|
||||
static = make_splats(torch.tensor([[5.0, 5.0, 5.0]]))
|
||||
start = torch.tensor([[0.0, 0.0, 2.0], [0.0, 1.0, 3.0]])
|
||||
end = torch.tensor([[1.0, 0.0, 2.0], [0.0, -1.0, 3.0]])
|
||||
ts = torch.linspace(0.0, 1.0, T)
|
||||
trajectories = torch.stack([start + (end - start) * t for t in ts]) # [5,2,3]
|
||||
|
||||
s4d = GS4D_nodes.GaussianSplats4D(
|
||||
static=static, canonical=canonical, trajectories=trajectories, times=ts,
|
||||
)
|
||||
|
||||
mid = s4d.at_time(0.5)
|
||||
assert len(mid) == 3, f"count {len(mid)} != dynamic+static (3)"
|
||||
# Concat order is [static, dynamic].
|
||||
assert torch.allclose(mid.xyz[0], static.xyz[0], atol=1e-6)
|
||||
expected_mid = 0.5 * (start + end)
|
||||
assert torch.allclose(mid.xyz[1:], expected_mid, atol=1e-5), (
|
||||
f"midpoint mismatch: {mid.xyz[1:]} vs {expected_mid}"
|
||||
)
|
||||
|
||||
lo = s4d.at_time(-1.0)
|
||||
hi = s4d.at_time(2.0)
|
||||
assert torch.allclose(lo.xyz[1:], start, atol=1e-5), "t<range should clamp to first step"
|
||||
assert torch.allclose(hi.xyz[1:], end, atol=1e-5), "t>range should clamp to last step"
|
||||
|
||||
|
||||
def test_05_build_splats4d():
|
||||
T = 5
|
||||
ts = torch.linspace(0.0, 1.0, T)
|
||||
# Two control tracks moving apart along x.
|
||||
track_a = torch.stack([torch.tensor([-1.0 - 2.0 * t, 0.0, 2.0]) for t in ts])
|
||||
track_b = torch.stack([torch.tensor([1.0 + 2.0 * t, 0.0, 2.0]) for t in ts])
|
||||
trajectories3d = torch.stack([track_a, track_b], dim=1) # [T,2,3]
|
||||
|
||||
canonical = make_splats(torch.tensor([[-1.05, 0.0, 2.0], [1.05, 0.0, 2.0]]))
|
||||
node = GS4D_nodes.BuildSplats4D()
|
||||
(s4d,) = node.build_splats4d(
|
||||
canonical=canonical,
|
||||
trajectories3d=trajectories3d,
|
||||
reference_index=0,
|
||||
knn=1,
|
||||
rbf_gamma=0.0,
|
||||
device="cpu",
|
||||
)
|
||||
traj = s4d.trajectories
|
||||
assert traj.shape == (T, 2, 3), f"trajectories shape {tuple(traj.shape)}"
|
||||
# Reference timestep: splats stay at their canonical positions.
|
||||
assert torch.allclose(traj[0], canonical.xyz, atol=1e-5)
|
||||
# Each splat follows its nearest track's displacement direction.
|
||||
disp0 = traj[-1, 0] - traj[0, 0]
|
||||
disp1 = traj[-1, 1] - traj[0, 1]
|
||||
assert disp0[0] < -1.0, f"splat 0 should move -x with track A, moved {disp0.tolist()}"
|
||||
assert disp1[0] > 1.0, f"splat 1 should move +x with track B, moved {disp1.tolist()}"
|
||||
assert torch.allclose(traj[-1, 0], torch.tensor([-3.05, 0.0, 2.0]), atol=1e-4)
|
||||
assert torch.allclose(traj[-1, 1], torch.tensor([3.05, 0.0, 2.0]), atol=1e-4)
|
||||
|
||||
|
||||
def test_06_split_splats_by_mask():
|
||||
H = W = 32
|
||||
mask = torch.zeros(H, W)
|
||||
mask[:, : W // 2] = 1.0 # left half white
|
||||
|
||||
# 10 splats projecting into the left half (x<0), 10 into the right half,
|
||||
# 5 behind the camera.
|
||||
jitter = torch.linspace(-0.1, 0.1, 10)
|
||||
left = torch.stack([torch.full((10,), -0.5) + jitter * 0.1, jitter, torch.full((10,), 2.0)], dim=-1)
|
||||
right = torch.stack([torch.full((10,), 0.5) + jitter * 0.1, jitter, torch.full((10,), 2.0)], dim=-1)
|
||||
behind = torch.stack([jitter[:5], jitter[:5], torch.full((5,), -2.0)], dim=-1)
|
||||
splats = make_splats(torch.cat([left, right, behind], dim=0))
|
||||
|
||||
node = GS4D_nodes.SplitSplatsByMask()
|
||||
inside, outside = node.split_splats(
|
||||
splats=splats,
|
||||
mask=mask,
|
||||
projection="PINHOLE",
|
||||
horizontal_fov=90.0,
|
||||
threshold=0.5,
|
||||
camera_matrix=None,
|
||||
device="cpu",
|
||||
)
|
||||
assert len(inside) == 10, f"inside count {len(inside)} != 10"
|
||||
assert len(outside) == 15, f"outside count {len(outside)} != 15 (10 right + 5 behind)"
|
||||
assert (inside.xyz[:, 0] < 0).all(), "inside splats should be the x<0 group"
|
||||
|
||||
|
||||
def test_07_motion_mask_from_depth():
|
||||
T, H, W = 6, 32, 32
|
||||
depth = torch.full((T, H, W), 5.0)
|
||||
r0, r1 = 8, 16
|
||||
for t in range(T):
|
||||
depth[t, r0:r1, r0:r1] = 3.0 + 0.4 * t # depth-changing square patch
|
||||
|
||||
poses = torch.eye(4).unsqueeze(0).expand(T, 4, 4).contiguous()
|
||||
node = GS4D_nodes.MotionMaskFromDepth()
|
||||
(mask,) = node.motion_mask(
|
||||
depth_seq=depth,
|
||||
trajectory=poses,
|
||||
input_projection="PINHOLE",
|
||||
input_horizontal_fov=90.0,
|
||||
threshold=0.10,
|
||||
frame_gap=2,
|
||||
dilate=0,
|
||||
device="cpu",
|
||||
)
|
||||
assert mask.shape == (T, H, W), f"mask shape {tuple(mask.shape)}"
|
||||
|
||||
patch = mask[:, r0:r1, r0:r1]
|
||||
background = mask.clone()
|
||||
background[:, r0:r1, r0:r1] = 0.0
|
||||
patch_mean = float(patch.mean())
|
||||
bg_sum = float(background.sum())
|
||||
assert patch_mean > 0.9, f"moving square under-detected: mean {patch_mean:.3f}"
|
||||
assert bg_sum == 0.0, f"static plane falsely flagged: {bg_sum} pixels"
|
||||
|
||||
|
||||
def test_08_align_depth_scale_and_depth_edge_filter():
|
||||
H = W = 32
|
||||
new_depth = torch.rand(H, W) * 9.0 + 1.0
|
||||
# ref disparity = 0.5 * new disparity + 0.1 (i.e. ref = 2*new before shift).
|
||||
true_scale, true_shift = 0.5, 0.1
|
||||
ref_depth = 1.0 / (true_scale / new_depth + true_shift)
|
||||
valid = torch.ones(H, W)
|
||||
|
||||
aligned, scale, shift = world_nodes.align_depth_scale(
|
||||
new_depth, ref_depth, valid, mode="scale_shift"
|
||||
)
|
||||
assert abs(scale - true_scale) / true_scale < 0.05, f"scale {scale} vs {true_scale}"
|
||||
assert abs(shift - true_shift) / true_shift < 0.05, f"shift {shift} vs {true_shift}"
|
||||
rel_err = float(((aligned - ref_depth).abs() / ref_depth).max())
|
||||
assert rel_err < 0.01, f"aligned depth off by {rel_err:.4f} (rel)"
|
||||
|
||||
# DepthEdgeFilter: a vertical step edge must be masked out, flat kept.
|
||||
depth = torch.full((H, W), 1.0)
|
||||
depth[:, W // 2 :] = 5.0
|
||||
node = pointcloud_nodes.DepthEdgeFilter()
|
||||
(valid_mask,) = node.filter_edges(depth, relative_threshold=0.05, dilate=1)
|
||||
assert valid_mask.shape == (H, W)
|
||||
edge_cols = valid_mask[:, W // 2 - 1 : W // 2 + 1]
|
||||
assert float(edge_cols.max()) == 0.0, "step-edge pixels not masked out"
|
||||
assert float(valid_mask[:, : W // 2 - 3].min()) == 1.0, "flat left region wrongly masked"
|
||||
assert float(valid_mask[:, W // 2 + 3 :].min()) == 1.0, "flat right region wrongly masked"
|
||||
|
||||
|
||||
def test_09_fuse_splats():
|
||||
n = 20
|
||||
voxel = 0.5
|
||||
base = torch.stack(
|
||||
[
|
||||
torch.arange(n, dtype=torch.float32) * voxel + 0.15,
|
||||
torch.full((n,), 0.15),
|
||||
torch.full((n,), 0.15),
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
cloud_a = make_splats(base)
|
||||
cloud_b = make_splats(base + 0.2) # same voxels as A (0.15+0.2 < 0.5)
|
||||
|
||||
node = GS_nodes.FuseSplats()
|
||||
(fused,) = node.fuse_splats(cloud_a, cloud_b, voxel, "smart", 1.0, 1.0, device="cpu")
|
||||
assert len(fused) < len(cloud_a) + len(cloud_b), (
|
||||
f"voxel fuse did not reduce: {len(fused)} vs {len(cloud_a) + len(cloud_b)}"
|
||||
)
|
||||
assert len(fused) == n, f"expected one splat per voxel ({n}), got {len(fused)}"
|
||||
|
||||
# Strong weight_a pulls fused positions onto cloud A.
|
||||
(fused_w,) = node.fuse_splats(cloud_a, cloud_b, voxel, "average", 1000.0, 1.0, device="cpu")
|
||||
assert len(fused_w) == n
|
||||
d_a = torch.cdist(fused_w.xyz, cloud_a.xyz).min(dim=1).values
|
||||
d_b = torch.cdist(fused_w.xyz, cloud_b.xyz).min(dim=1).values
|
||||
assert float(d_a.max()) < 0.01, f"fused positions not near cloud A (max dist {float(d_a.max()):.4f})"
|
||||
assert (d_a < d_b).all(), "weight_a=1000 should pull fused splats toward cloud A"
|
||||
|
||||
|
||||
def test_10_sphere_splat_seed():
|
||||
H, W = 64, 128
|
||||
stride = 2
|
||||
color = (0.2, 0.6, 0.9)
|
||||
pano = torch.tensor(color).view(1, 1, 1, 3).expand(1, H, W, 3).contiguous()
|
||||
|
||||
node = world_nodes.SphereSplatSeed()
|
||||
(splats,) = node.seed_sphere(
|
||||
image=pano,
|
||||
horizontal_fov=360.0,
|
||||
radius=5.0,
|
||||
splat_scale_frac=1.5,
|
||||
stride=stride,
|
||||
device="cpu",
|
||||
)
|
||||
expected = (H // stride) * (W // stride)
|
||||
assert abs(len(splats) - expected) <= max(4, expected // 20), (
|
||||
f"splat count {len(splats)} far from expected ~{expected}"
|
||||
)
|
||||
|
||||
image, mask, disparity = GS_nodes.render_gaussians(
|
||||
splats, IDENTITY_4X4, "PINHOLE", 60.0, 64, 64,
|
||||
render_mode="fast", device="cpu",
|
||||
)
|
||||
assert float(mask.sum()) > 0.0, "pinhole render of the sphere seed is empty"
|
||||
solid = mask > 0.9
|
||||
assert bool(solid.any()), "no confidently covered pixels in the render"
|
||||
rendered = image[0][solid] # [K,3]
|
||||
target = torch.tensor(color)
|
||||
err = (rendered.mean(dim=0) - target).abs().max().item()
|
||||
assert err < 0.05, f"color round-trip failed: rendered mean {rendered.mean(dim=0).tolist()} vs {color}"
|
||||
|
||||
|
||||
# --------------------------------------------------------------------------- #
|
||||
# Runner
|
||||
# --------------------------------------------------------------------------- #
|
||||
TESTS = [
|
||||
test_01_interpolate_se3,
|
||||
test_02_render_gaussians_shapes_and_empty,
|
||||
test_03_fast_mode_anisotropy,
|
||||
test_04_at_time,
|
||||
test_05_build_splats4d,
|
||||
test_06_split_splats_by_mask,
|
||||
test_07_motion_mask_from_depth,
|
||||
test_08_align_depth_scale_and_depth_edge_filter,
|
||||
test_09_fuse_splats,
|
||||
test_10_sphere_splat_seed,
|
||||
]
|
||||
|
||||
|
||||
def main() -> int:
|
||||
passed = 0
|
||||
failed = []
|
||||
for test in TESTS:
|
||||
name = test.__name__
|
||||
try:
|
||||
test()
|
||||
except Exception:
|
||||
failed.append(name)
|
||||
print(f"[FAIL] {name}")
|
||||
traceback.print_exc()
|
||||
else:
|
||||
passed += 1
|
||||
print(f"[ ok ] {name}")
|
||||
print(f"\n{passed}/{len(TESTS)} tests passed")
|
||||
if failed:
|
||||
print("Failed:", ", ".join(failed))
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
File diff suppressed because one or more lines are too long
+214
-6
@@ -7,7 +7,10 @@ import os
|
||||
import folder_paths
|
||||
import logging
|
||||
import hashlib
|
||||
from kornia.filters import median_blur
|
||||
try:
|
||||
from kornia.filters import median_blur
|
||||
except ImportError: # kornia is optional; median_blur is not used in this module
|
||||
median_blur = None
|
||||
|
||||
from tqdm import tqdm
|
||||
# Try importing open3d and its visualization modules; log a warning if not found
|
||||
@@ -136,6 +139,116 @@ def project_first_hit(volume_sparse: torch.Tensor) -> Tuple[torch.Tensor, torch.
|
||||
|
||||
return rgba.permute(2, 0, 1), first_hit.any(dim=2)
|
||||
|
||||
# ==== SE(3) trajectory interpolation ==== #
|
||||
def _rotmat_to_quat_wxyz(R: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Convert a batch of rotation matrices [K,3,3] to unit quaternions [K,4] (wxyz).
|
||||
Uses Shepperd's method for numerical robustness. K is expected to be small
|
||||
(trajectory waypoints), so a Python loop is acceptable.
|
||||
"""
|
||||
quats = []
|
||||
for i in range(R.shape[0]):
|
||||
m = R[i]
|
||||
trace = m[0, 0] + m[1, 1] + m[2, 2]
|
||||
if trace > 0.0:
|
||||
s = torch.sqrt(trace + 1.0) * 2.0
|
||||
w = 0.25 * s
|
||||
x = (m[2, 1] - m[1, 2]) / s
|
||||
y = (m[0, 2] - m[2, 0]) / s
|
||||
z = (m[1, 0] - m[0, 1]) / s
|
||||
elif m[0, 0] > m[1, 1] and m[0, 0] > m[2, 2]:
|
||||
s = torch.sqrt(1.0 + m[0, 0] - m[1, 1] - m[2, 2]) * 2.0
|
||||
w = (m[2, 1] - m[1, 2]) / s
|
||||
x = 0.25 * s
|
||||
y = (m[0, 1] + m[1, 0]) / s
|
||||
z = (m[0, 2] + m[2, 0]) / s
|
||||
elif m[1, 1] > m[2, 2]:
|
||||
s = torch.sqrt(1.0 + m[1, 1] - m[0, 0] - m[2, 2]) * 2.0
|
||||
w = (m[0, 2] - m[2, 0]) / s
|
||||
x = (m[0, 1] + m[1, 0]) / s
|
||||
y = 0.25 * s
|
||||
z = (m[1, 2] + m[2, 1]) / s
|
||||
else:
|
||||
s = torch.sqrt(1.0 + m[2, 2] - m[0, 0] - m[1, 1]) * 2.0
|
||||
w = (m[1, 0] - m[0, 1]) / s
|
||||
x = (m[0, 2] + m[2, 0]) / s
|
||||
y = (m[1, 2] + m[2, 1]) / s
|
||||
z = 0.25 * s
|
||||
quats.append(torch.stack([w, x, y, z]))
|
||||
q = torch.stack(quats, dim=0)
|
||||
return q / q.norm(dim=-1, keepdim=True).clamp(min=1e-12)
|
||||
|
||||
|
||||
def _quat_wxyz_to_rotmat(q: torch.Tensor) -> torch.Tensor:
|
||||
"""Convert unit quaternions [N,4] (wxyz) to rotation matrices [N,3,3]."""
|
||||
q = q / q.norm(dim=-1, keepdim=True).clamp(min=1e-12)
|
||||
w, x, y, z = q.unbind(-1)
|
||||
R = torch.stack([
|
||||
1 - 2 * (y * y + z * z), 2 * (x * y - w * z), 2 * (x * z + w * y),
|
||||
2 * (x * y + w * z), 1 - 2 * (x * x + z * z), 2 * (y * z - w * x),
|
||||
2 * (x * z - w * y), 2 * (y * z + w * x), 1 - 2 * (x * x + y * y),
|
||||
], dim=-1).reshape(*q.shape[:-1], 3, 3)
|
||||
return R
|
||||
|
||||
|
||||
def _quat_slerp(q0: torch.Tensor, q1: torch.Tensor, alpha: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Spherical linear interpolation between quaternion batches q0, q1 [N,4] (wxyz)
|
||||
with per-element interpolation factors alpha [N]. Falls back to normalized
|
||||
lerp when the quaternions are nearly parallel.
|
||||
"""
|
||||
dot = (q0 * q1).sum(dim=-1, keepdim=True)
|
||||
q1 = torch.where(dot < 0.0, -q1, q1) # shortest arc
|
||||
dot = dot.abs().clamp(max=1.0)
|
||||
a = alpha.reshape(-1, 1).to(q0.dtype)
|
||||
theta = torch.acos(dot)
|
||||
sin_theta = torch.sin(theta)
|
||||
near_parallel = sin_theta < 1e-6
|
||||
denom = sin_theta.clamp(min=1e-12)
|
||||
w0 = torch.where(near_parallel, 1.0 - a, torch.sin((1.0 - a) * theta) / denom)
|
||||
w1 = torch.where(near_parallel, a, torch.sin(a * theta) / denom)
|
||||
q = w0 * q0 + w1 * q1
|
||||
return q / q.norm(dim=-1, keepdim=True).clamp(min=1e-12)
|
||||
|
||||
|
||||
def interpolate_se3(trajectory: torch.Tensor, num_steps: int) -> torch.Tensor:
|
||||
"""trajectory [K,4,4] -> [num_steps,4,4]. Piecewise: quaternion SLERP on R, lerp on t.
|
||||
K==1 -> repeat. Must return valid rotation matrices (orthonormal)."""
|
||||
if isinstance(trajectory, np.ndarray):
|
||||
trajectory = torch.from_numpy(trajectory)
|
||||
trajectory = trajectory.float()
|
||||
if trajectory.dim() == 2:
|
||||
trajectory = trajectory.unsqueeze(0)
|
||||
if trajectory.dim() != 3 or trajectory.shape[-2:] != (4, 4):
|
||||
raise ValueError(f"interpolate_se3 expects trajectory of shape [K,4,4], got {tuple(trajectory.shape)}")
|
||||
if num_steps < 1:
|
||||
raise ValueError(f"interpolate_se3 requires num_steps >= 1, got {num_steps}")
|
||||
K = trajectory.shape[0]
|
||||
if K == 1:
|
||||
return trajectory.expand(num_steps, 4, 4).clone()
|
||||
|
||||
R = trajectory[:, :3, :3]
|
||||
t = trajectory[:, :3, 3]
|
||||
q = _rotmat_to_quat_wxyz(R)
|
||||
# Enforce hemisphere continuity along the waypoint sequence so piecewise
|
||||
# SLERP always takes the shortest arc between consecutive poses.
|
||||
for k in range(1, K):
|
||||
if (q[k] * q[k - 1]).sum() < 0.0:
|
||||
q[k] = -q[k]
|
||||
|
||||
idxs = torch.linspace(0, K - 1, num_steps, device=trajectory.device)
|
||||
lower = idxs.floor().long().clamp(max=K - 2)
|
||||
upper = lower + 1
|
||||
alpha = (idxs - lower.float())
|
||||
|
||||
q_interp = _quat_slerp(q[lower], q[upper], alpha)
|
||||
t_interp = t[lower] * (1.0 - alpha).unsqueeze(-1) + t[upper] * alpha.unsqueeze(-1)
|
||||
|
||||
out = torch.eye(4, dtype=trajectory.dtype, device=trajectory.device).repeat(num_steps, 1, 1)
|
||||
out[:, :3, :3] = _quat_wxyz_to_rotmat(q_interp)
|
||||
out[:, :3, 3] = t_interp
|
||||
return out
|
||||
|
||||
# ==== Node Definitions ==== #
|
||||
class DepthToPointCloud:
|
||||
"""
|
||||
@@ -794,8 +907,9 @@ class CameraMotionNode:
|
||||
|
||||
class CameraInterpolationNode:
|
||||
"""
|
||||
Wrap two 4×4 poses into a trajectory tensor.
|
||||
Outputs only `trajectory` (shape 2×4×4).
|
||||
Interpolate between two 4×4 poses into a trajectory tensor using proper
|
||||
SE(3) interpolation (quaternion SLERP on rotation, lerp on translation).
|
||||
Outputs `trajectory` (shape num_steps×4×4, default 2×4×4).
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
@@ -804,7 +918,10 @@ class CameraInterpolationNode:
|
||||
"required": {
|
||||
"initial_matrix": ("MAT_4X4",),
|
||||
"final_matrix": ("MAT_4X4",),
|
||||
}
|
||||
},
|
||||
"optional": {
|
||||
"num_steps": ("INT", {"default": 2, "min": 2, "max": 4096, "tooltip": "Number of poses in the output trajectory, SE(3)-interpolated between the two matrices."}),
|
||||
},
|
||||
}
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("trajectory",)
|
||||
@@ -815,14 +932,15 @@ class CameraInterpolationNode:
|
||||
self,
|
||||
initial_matrix: torch.Tensor,
|
||||
final_matrix: torch.Tensor,
|
||||
num_steps: int = 2,
|
||||
) -> Tuple[torch.Tensor]:
|
||||
# stack into a (2,4,4) trajectory
|
||||
# convert to tensor if needed
|
||||
if isinstance(initial_matrix, np.ndarray):
|
||||
initial_matrix = torch.from_numpy(initial_matrix).float()
|
||||
if isinstance(final_matrix, np.ndarray):
|
||||
final_matrix = torch.from_numpy(final_matrix).float()
|
||||
traj = torch.stack([initial_matrix, final_matrix], dim=0)
|
||||
keyframes = torch.stack([initial_matrix.float(), final_matrix.float()], dim=0)
|
||||
traj = interpolate_se3(keyframes, num_steps)
|
||||
return (traj,)
|
||||
|
||||
|
||||
@@ -1250,6 +1368,95 @@ class LoadTrajectory:
|
||||
return f"Invalid trajectory file: {trajectory_file}"
|
||||
return True
|
||||
|
||||
class DepthEdgeFilter:
|
||||
"""
|
||||
Detect "flying pixel" depth discontinuities and output a validity mask.
|
||||
A pixel is flagged as an edge where |depth gradient| / depth exceeds
|
||||
`relative_threshold`; edges are optionally dilated. Returns a MASK with
|
||||
1.0 where the depth is valid (NOT a flying-pixel edge) and 0.0 on edges.
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
# Depth: [H,W] or [T,H,W], trailing channel dim of 1 accepted
|
||||
"depth": ("TENSOR", {"shape_hint": [None, None, None]}),
|
||||
"relative_threshold": ("FLOAT", {"default": 0.05, "min": 0.0, "max": 10.0, "step": 0.005, "tooltip": "Mark a pixel as edge where |depth gradient| / depth exceeds this value."}),
|
||||
"dilate": ("INT", {"default": 1, "min": 0, "max": 64, "tooltip": "Grow detected edges by this many pixels (max-pool dilation)."}),
|
||||
},
|
||||
"optional": {
|
||||
"mask": ("MASK", {"tooltip": "Optional validity mask ANDed with the edge-filter result."}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MASK",)
|
||||
RETURN_NAMES = ("valid_mask",)
|
||||
FUNCTION = "filter_edges"
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
|
||||
def filter_edges(
|
||||
self,
|
||||
depth: torch.Tensor,
|
||||
relative_threshold: float,
|
||||
dilate: int,
|
||||
mask: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor]:
|
||||
d = depth
|
||||
if isinstance(d, np.ndarray):
|
||||
d = torch.from_numpy(d)
|
||||
d = d.float()
|
||||
# Accept [H,W], [H,W,1], [T,H,W], [T,H,W,1]
|
||||
if d.dim() == 4 and d.shape[-1] == 1:
|
||||
d = d[..., 0]
|
||||
elif d.dim() == 3 and d.shape[-1] == 1:
|
||||
d = d[..., 0]
|
||||
squeeze_batch = False
|
||||
if d.dim() == 2:
|
||||
d = d.unsqueeze(0)
|
||||
squeeze_batch = True
|
||||
if d.dim() != 3:
|
||||
raise ValueError(f"DepthEdgeFilter expects depth of shape [H,W] or [T,H,W] (trailing 1 ok), got {tuple(depth.shape)}")
|
||||
|
||||
eps = 1e-8
|
||||
# Forward differences along x and y; propagate each difference to both
|
||||
# neighbouring pixels so both sides of a discontinuity are flagged.
|
||||
dx = (d[:, :, 1:] - d[:, :, :-1]).abs()
|
||||
dy = (d[:, 1:, :] - d[:, :-1, :]).abs()
|
||||
gx = torch.zeros_like(d)
|
||||
gx[:, :, :-1] = dx
|
||||
gx[:, :, 1:] = torch.maximum(gx[:, :, 1:], dx)
|
||||
gy = torch.zeros_like(d)
|
||||
gy[:, :-1, :] = dy
|
||||
gy[:, 1:, :] = torch.maximum(gy[:, 1:, :], dy)
|
||||
grad = torch.maximum(gx, gy)
|
||||
edge = (grad / d.abs().clamp(min=eps)) > relative_threshold
|
||||
|
||||
if dilate > 0:
|
||||
k = 2 * int(dilate) + 1
|
||||
edge = F.max_pool2d(edge.float().unsqueeze(1), kernel_size=k, stride=1, padding=int(dilate)).squeeze(1) > 0.5
|
||||
|
||||
valid = (~edge).float()
|
||||
|
||||
if mask is not None:
|
||||
m = mask
|
||||
if isinstance(m, np.ndarray):
|
||||
m = torch.from_numpy(m)
|
||||
m = m.float().to(valid.device)
|
||||
if m.dim() == 4 and m.shape[-1] == 1:
|
||||
m = m[..., 0]
|
||||
if m.dim() == 2:
|
||||
m = m.unsqueeze(0)
|
||||
if m.shape[0] == 1 and valid.shape[0] > 1:
|
||||
m = m.expand(valid.shape[0], -1, -1)
|
||||
if m.shape[-2:] != valid.shape[-2:]:
|
||||
m = F.interpolate(m.unsqueeze(1), size=valid.shape[-2:], mode="nearest").squeeze(1)
|
||||
valid = valid * (m > 0.5).float()
|
||||
|
||||
if squeeze_batch:
|
||||
valid = valid[0]
|
||||
return (valid,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"DepthToPointCloud": DepthToPointCloud,
|
||||
"TransformPointCloud": TransformPointCloud,
|
||||
@@ -1264,4 +1471,5 @@ NODE_CLASS_MAPPINGS = {
|
||||
"PointCloudCleaner": PointCloudCleaner,
|
||||
"SaveTrajectory": SaveTrajectory,
|
||||
"LoadTrajectory": LoadTrajectory,
|
||||
"DepthEdgeFilter": DepthEdgeFilter,
|
||||
}
|
||||
+461
@@ -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,
|
||||
}
|
||||
@@ -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/"]
|
||||
Submodule
+1
Submodule submodules/ml-sharpt added at 1eaa046834
+30
-23
@@ -7,7 +7,7 @@ from typing import Dict, Any, Tuple
|
||||
from tqdm import tqdm # Added tqdm import
|
||||
|
||||
# Import existing pointcloud nodes and projection definitions
|
||||
from .pointcloud_nodes import DepthToPointCloud, TransformPointCloud, ProjectPointCloud, Projection, PointCloudCleaner
|
||||
from .pointcloud_nodes import DepthToPointCloud, TransformPointCloud, ProjectPointCloud, Projection, PointCloudCleaner, interpolate_se3
|
||||
import folder_paths
|
||||
|
||||
# Ensure video_depth_anything is on path
|
||||
@@ -95,18 +95,8 @@ class VideoCameraMotionSequence:
|
||||
# depth_seq: [T, H, W] or [T, H, W, 1]
|
||||
T, H, W, _ = frames.shape
|
||||
|
||||
# Interpolate trajectory to match T
|
||||
K = trajectory.shape[0]
|
||||
if K < 2:
|
||||
interp_traj = trajectory.expand(T, 4, 4).clone()
|
||||
else:
|
||||
idxs = torch.linspace(0, K - 1, T, device=trajectory.device)
|
||||
lower = idxs.floor().long().clamp(max=K - 2)
|
||||
upper = lower + 1
|
||||
alpha = (idxs - lower.float()).unsqueeze(-1).unsqueeze(-1)
|
||||
traj_lower = trajectory[lower]
|
||||
traj_upper = trajectory[upper]
|
||||
interp_traj = traj_lower * (1 - alpha) + traj_upper * alpha
|
||||
# Interpolate trajectory to match T (SE(3): quaternion SLERP on R, lerp on t)
|
||||
interp_traj = interpolate_se3(trajectory, T)
|
||||
|
||||
out_frames = []
|
||||
out_masks = []
|
||||
@@ -122,12 +112,12 @@ class VideoCameraMotionSequence:
|
||||
for i, (frame, depth, pose) in enumerate(tqdm(zip(frames, depth_seq, interp_traj), total=T, desc="Processing video frames")):
|
||||
if depth.dim() == 3 and depth.shape[-1] == 1:
|
||||
depth = depth.squeeze(-1)
|
||||
# Use mask if provided
|
||||
mask = None
|
||||
# Use mask if provided; must be (re)initialized every iteration
|
||||
mask = None
|
||||
if mask_seq is not None:
|
||||
mask = mask_seq[i]
|
||||
if mask.dim() == 3 and mask.shape[-1] == 1:
|
||||
mask = mask.squeeze(-1)
|
||||
if mask.dim() == 3 and mask.shape[-1] == 1:
|
||||
mask = mask.squeeze(-1)
|
||||
# to pointcloud
|
||||
pc, = DepthToPointCloud().depth_to_pointcloud(
|
||||
image=frame.permute(2, 0, 1),
|
||||
@@ -237,6 +227,10 @@ class DepthFramesToVideo:
|
||||
raw_color = raw_u8.unsqueeze(1).repeat(1, 3, 1, 1).permute(0, 2, 3, 1)
|
||||
return raw_color, ds_color # [T, 3, H, W] -> [T, H, W, 3]
|
||||
|
||||
# Cache for loaded VideoDepthAnything models, keyed by (checkpoint, device)
|
||||
_VIDEO_DEPTH_MODEL_CACHE: Dict[Tuple[str, str], Any] = {}
|
||||
|
||||
|
||||
class VideoMetricDepthEstimate:
|
||||
"""
|
||||
Estimates metric depth for a sequence of frames using VideoDepthAnything.
|
||||
@@ -267,16 +261,29 @@ class VideoMetricDepthEstimate:
|
||||
input_size: int,
|
||||
max_fps: int,
|
||||
) -> Tuple[torch.Tensor, float]:
|
||||
if VideoDepthAnything is None:
|
||||
raise ImportError("VideoDepthAnything library not found")
|
||||
if NO_VIDEO_DEPTH_ANYTHING:
|
||||
raise ImportError(
|
||||
f"VideoDepthAnything library not found. Clone "
|
||||
f"https://github.com/DepthAnything/Video-Depth-Anything into {COMFYUI_ROOT!r} "
|
||||
f"(expected module path: {video_depth_path!r})."
|
||||
)
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
# if max input<1.5 normalize to 0-255
|
||||
if frames.max() < 1.5:
|
||||
frames = (frames * 255)
|
||||
model = VideoDepthAnything(**{"encoder": "vitl", "features": 256, "out_channels": [256,512,1024,1024]})
|
||||
state = torch.load("/root/ComfyUI/models/checkpoints/{}".format(model_checkpoint), map_location='cpu')
|
||||
model.load_state_dict(state, strict=True)
|
||||
model = model.to(device).eval()
|
||||
cache_key = (model_checkpoint, str(device))
|
||||
model = _VIDEO_DEPTH_MODEL_CACHE.get(cache_key)
|
||||
if model is None:
|
||||
# Same checkpoint directory as computed in INPUT_TYPES
|
||||
model_dir = os.path.join(os.getcwd(), "models", "checkpoints")
|
||||
checkpoint_path = os.path.join(model_dir, model_checkpoint)
|
||||
if not os.path.isfile(checkpoint_path):
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
model = VideoDepthAnything(**{"encoder": "vitl", "features": 256, "out_channels": [256,512,1024,1024]})
|
||||
state = torch.load(checkpoint_path, map_location='cpu')
|
||||
model.load_state_dict(state, strict=True)
|
||||
model = model.to(device).eval()
|
||||
_VIDEO_DEPTH_MODEL_CACHE[cache_key] = model
|
||||
np_frames = frames.cpu().numpy().astype(np.uint8)
|
||||
metric_depths, fps = model.infer_video_depth(np_frames, max_fps, input_size=input_size, device=device.type, fp32=False)
|
||||
return (torch.from_numpy(metric_depths), float(fps))
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+632
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user