Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d49f3f03ce |
@@ -1,10 +0,0 @@
|
||||
# Excluded from the ComfyUI Registry archive (not from git).
|
||||
demo_images/
|
||||
notebooks/
|
||||
docs/
|
||||
screenshot1.ply
|
||||
__pycache__/
|
||||
models/
|
||||
.github/
|
||||
Makefile
|
||||
install.sh
|
||||
@@ -1,28 +0,0 @@
|
||||
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 }}
|
||||
@@ -1,3 +0,0 @@
|
||||
[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
@@ -1,29 +0,0 @@
|
||||
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.
|
||||
@@ -1,19 +0,0 @@
|
||||
.PHONY: install install_all install_modules download_flux download_vae
|
||||
|
||||
# install everything except WAN‑VACE downloads
|
||||
install:
|
||||
./install.sh install
|
||||
|
||||
# install everything + WAN‑VACE + HF login
|
||||
install_all:
|
||||
./install.sh all
|
||||
|
||||
# lower‑level helpers
|
||||
install_modules:
|
||||
./install.sh modules
|
||||
|
||||
download_flux:
|
||||
./install.sh flux
|
||||
|
||||
download_vae:
|
||||
./install.sh vae
|
||||
@@ -1,7 +1,6 @@
|
||||
# camera-comfyUI
|
||||
[](https://deepwiki.com/Alexankharin/camera-comfyUI)
|
||||
|
||||

|
||||

|
||||
|
||||
> Custom ComfyUI nodes for advanced reprojections, point cloud processing, and camera-driven workflows.
|
||||
|
||||
@@ -14,7 +13,6 @@
|
||||
* [Installation](#installation)
|
||||
* [Node Categories](#node-categories)
|
||||
* [Node Reference](#node-reference)
|
||||
* [Video → 4D World](#video--4d-world)
|
||||
* [Workflows](#workflows)
|
||||
* [Example Workflows](#example-workflows)
|
||||
* [Contributing](#contributing)
|
||||
@@ -36,14 +34,6 @@ 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
|
||||
@@ -64,13 +54,6 @@ The node pack is published to the [ComfyUI Registry](https://registry.comfy.org)
|
||||
|
||||
* *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:
|
||||
@@ -97,45 +80,18 @@ The node pack is published to the [ComfyUI Registry](https://registry.comfy.org)
|
||||
* ### Reprojection Nodes
|
||||
|
||||
* `ReprojectImage`, `ReprojectDepth`, `OutpaintAnyProjection`
|
||||
|
||||
* ### Matrix Nodes
|
||||
|
||||
* `TransformToMatrix`, `TransformToMatrixManual`
|
||||
|
||||
* ### Depth Nodes
|
||||
|
||||
* `DepthEstimatorNode`, `DepthToImageNode`, `ZDepthToRayDepthNode`
|
||||
* `CombineDepthsNode`, `DepthRenormalizer`, `FisheyeDepthEstimator`
|
||||
* `CombineDepthsNode`, `DepthRenormalizer`
|
||||
|
||||
* ### Point Cloud Nodes
|
||||
|
||||
* `DepthToPointCloud`, `TransformPointCloud`, `ProjectPointCloud`, `PointCloudUnion`
|
||||
* `PointCloudCleaner`, `LoadPointCloud`, `SavePointCloud`, `ProjectAndClean`, `DepthEdgeFilter`
|
||||
|
||||
* ### Trajectory Nodes
|
||||
|
||||
* `DepthToPointCloud`, `TransformPointCloud`, `ProjectPointCloud`
|
||||
* `PointCloudUnion`, `PointCloudCleaner`, `LoadPointCloud`, `SavePointCloud`
|
||||
* `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`
|
||||
|
||||
---
|
||||
|
||||
@@ -152,65 +108,12 @@ The node pack is published to the [ComfyUI Registry](https://registry.comfy.org)
|
||||
| `DepthToPointCloud` | Converts Depth and image to → 3D point cloud tensor (N×7). |
|
||||
| `DepthToImageNode` | Converts depth to image (N×3) using a color map. |
|
||||
| `ZDepthToRayDepthNode` | Converts Z-depth (output of metric-depth-anything) to ray depth to compensate lens curvature. |
|
||||
| `TransformPointCloud` | Applies 4×4 rotation matrix to point cloud. |
|
||||
| `TransformPointCloud` | Applies 4×4 rotation matrix to point cloud |
|
||||
| `ProjectPointCloud` | Z-buffer–based projection of point cloud into image + mask. |
|
||||
| `PointCloudCleaner` | Removes isolated points via voxel filtering. |
|
||||
| `PointCloudUnion` | Combines multiple point clouds into one. |
|
||||
| `LoadPointCloud` | Loads a point cloud from `.npy` or `.ply` format. |
|
||||
| `SavePointCloud` | Saves a point cloud to `.npy` or `.ply` format. |
|
||||
| `CameraMotionNode` | Generates image and mask sequences along a camera trajectory with optional mask dilation/inversion. |
|
||||
| `CameraMotionNode` | Generates image sequences by moving camera along a trajectory. |
|
||||
| `CameraInterpolationNode` | Builds a trajectory tensor from two poses. |
|
||||
| `CameraTrajectoryNode` | Interactive Open3D GUI for recording camera waypoints. |
|
||||
| `SaveTrajectory` | Saves a trajectory tensor to a file. |
|
||||
| `LoadTrajectory` | Loads a trajectory tensor from a file. |
|
||||
| `VideoCameraMotionSequence` | Processes video frames and depth maps along a camera trajectory, generating reprojected outputs. |
|
||||
| `DepthFramesToVideo` | Converts a sequence of depth maps into video frame tensors for saving. |
|
||||
| `VideoMetricDepthEstimate` | Estimates metric depth for a sequence of frames using VideoDepthAnything. |
|
||||
| `DepthEdgeFilter` | Detects "flying pixel" depth discontinuities and outputs a validity mask (1.0 = valid). |
|
||||
| `LoadPlySplat` | Loads a 3D Gaussian Splatting `.ply` file into a `GSPLAT` object. |
|
||||
| `SavePlySplat` | Saves a `GSPLAT` to the ComfyUI output directory as a `.ply` file. |
|
||||
| `ImageToSplat` | Predicts Gaussian splats from a single image using SHARP. |
|
||||
| `FisheyeToGaussian` | Reprojects a fisheye view to multiple pinhole angles, predicts splats, rotates and merges them. |
|
||||
| `RotateSplats` | Applies a 4×4 transform matrix to a splat cloud. |
|
||||
| `MergeSplats` | Concatenates two `GSPLAT` objects into one. |
|
||||
| `FuseSplats` | Fuses two splat clouds with weighted voxel merging (keep/discard/average/smart modes). |
|
||||
| `RenderSplat` | Renders a splat cloud from a camera pose into an image + mask. |
|
||||
| `VideoToFusedSplats` | Runs SHARP on video keyframes, scale-aligns to metric depth, filters dynamic pixels, and fuses all keyframes into one world-frame splat cloud. |
|
||||
| `SplatPolish` | Optimizes a world-frame splat cloud against posed video frames (L1 + D-SSIM) using gsplat's differentiable rasterizer. |
|
||||
| `MotionMaskFromDepth` | Detects dynamic pixels from a depth+pose sequence (1.0 = moving). |
|
||||
| `EstimateTracks` | Runs CoTracker3 on a video; returns tracks `[T,N,2]` (pixels) and visibility `[T,N]`. |
|
||||
| `TracksToTrajectories` | Unprojects 2D tracks with depth and camera poses into world-space 3D trajectories `[T,M,3]`. |
|
||||
| `SplitSplatsByMask` | Projects splat centers into a 2D mask and splits the cloud into inside/outside parts. |
|
||||
| `BuildSplats4D` | Builds a 4D splat scene: each canonical splat follows a kNN blend of track control-point motions. |
|
||||
| `RenderSplats4DFrame` | Evaluates the 4D scene at a single time value and renders it from a given camera. |
|
||||
| `RenderSplats4DVideo` | Interpolates the camera path, sweeps time from start to end, and renders each frame. |
|
||||
| `SaveSplats4D` | Saves a `GSPLAT4D` scene as an `.npz` archive (plus optional per-frame PLYs). |
|
||||
| `LoadSplats4D` | Loads a `GSPLAT4D` scene from an `.npz` archive. |
|
||||
| `VideoPoseEstimator` | VGGT-based per-frame camera poses `[T,4,4]`, depth maps, FOV and depth confidence from a video clip. |
|
||||
| `TrajectoryInvert` | Inverts each 4×4 pose (world-to-camera ↔ camera-to-world). |
|
||||
| `TrajectoryCompose` | Per-frame matrix product `A @ B`; a single 4×4 input broadcasts over the other. |
|
||||
| `DepthScaleAnchor` | Robustly aligns a depth map to a reference depth via disparity-domain scale(+shift). |
|
||||
| `SplatTrajectoryEnricher` | Expands a splat world along a trajectory: render, outpaint holes with Flux, lift with SHARP, scale-align, smart-stitch. |
|
||||
| `SphereSplatSeed` | Converts an equirectangular panorama into a Gaussian sphere seeding a 360° world. |
|
||||
|
||||
---
|
||||
|
||||
## Video → 4D World
|
||||
|
||||
Turn a monocular video into a navigable 4D (3D + time) Gaussian splat scene and re-render it from any novel camera trajectory. The reference workflow is **`workflows/video_to_4d_world.json`**; the stages are:
|
||||
|
||||
1. **Pose & depth (VGGT)** — `VideoPoseEstimator` estimates per-frame world-to-camera poses `[T,4,4]`, depth maps, FOV and depth confidence from the input frames. Since the depth maps are Z-depths, run `ZDepthToRayDepthNode` before any node that expects ray depth (see caveats below). `DepthEdgeFilter` can additionally mask out flying pixels at depth discontinuities.
|
||||
2. **Motion masking** — `MotionMaskFromDepth` warps depth between frames using the estimated poses and flags pixels whose residual is too large as dynamic (moving objects vs. static background).
|
||||
3. **Static splat fusion + polish** — `VideoToFusedSplats` runs SHARP on keyframes, keeps only static pixels (via the motion mask), scale-aligns each keyframe to metric depth, transforms splats into the world frame and fuses them incrementally. `SplatPolish` then fine-tunes the fused cloud photometrically against the posed video frames.
|
||||
4. **Tracked dynamic 4D Gaussians** — `EstimateTracks` (CoTracker3) tracks a dense point grid across the video; `TracksToTrajectories` lifts the tracks to world-space 3D using depth + poses; `SplitSplatsByMask` separates dynamic splats from the static background; `BuildSplats4D` binds the dynamic canonical splats to track control points via kNN blending, producing a `GSPLAT4D` scene.
|
||||
5. **Render a novel trajectory** — build any new camera path (e.g. `CameraInterpolationNode`, `TrajectoryCompose` to retarget relative to a source pose) and render with `RenderSplats4DVideo` (or single frames with `RenderSplats4DFrame`). Save/reload scenes with `SaveSplats4D` / `LoadSplats4D`.
|
||||
|
||||
### Caveats
|
||||
|
||||
* **Z-depth vs ray depth**: depth estimators (including `VideoPoseEstimator`) output Z-depth; point-cloud and splat lifting nodes expect ray depth. Insert `ZDepthToRayDepthNode` where needed, or geometry will bow at wide FOVs.
|
||||
* **`SplatPolish` requires gsplat + CUDA**: without them it can fall back to the differentiable torch renderer at reduced resolution, which is extremely slow (minutes per 100 iterations).
|
||||
* **`EstimateTracks` downloads CoTracker3 via `torch.hub` on first use** — expect a one-time download and allow network access.
|
||||
* **`VideoPoseEstimator` downloads `facebook/VGGT-1B` (~5 GB)** on first use via `huggingface_hub`.
|
||||
| `PointCloudCleaner` | Removes isolated points via voxel filtering. |
|
||||
|
||||
---
|
||||
|
||||
@@ -229,10 +132,6 @@ A set of JSON workflows illustrating typical use cases. Each workflow lives in `
|
||||
| **Pointcloud.json** | Metric‐depth‐anything v2 → point cloud → camera view synthesis |
|
||||
| **pointcloud\_inpaint.json** | Inpaint + backproject to 3D for dynamic camera motion videos |
|
||||
| **Pointcloud\_walker.json** | GUI‐based camera control via Open3D |
|
||||
| **sbs180\_workflow.json** | Generate stereo (side-by-side) wide-angle/fisheye/equirectangular stereo pairs from a high-res input |
|
||||
| **video_camera.json** | Camera trajectory movement workflow using `wan-vace` for video inpainting. |
|
||||
| **video_to_4d_world\.json** | Video → 4D world: VGGT poses/depth → motion masking → fused static splats + polish → tracked dynamic 4D Gaussians → novel-trajectory render. |
|
||||
| **video_to_4d_walkable_world\.json** | Video → 4D WALKABLE world (test-friendly defaults): polished static splats enriched along a walk trajectory (`SplatTrajectoryEnricher`, Flux outpaint + SHARP) → 4D scene → walk-through render + `.ply`/`.npz` exports for free walking in external 3DGS viewers. |
|
||||
|
||||
---
|
||||
|
||||
@@ -291,55 +190,10 @@ Inpaint image with shifted camera and backproject for dynamic camera‐driven vi
|
||||
<img src="demo_images/Fisheye_camera_pointcloud_moved_outpainted.png" alt="PointCloud Inpaint" width="40%" />
|
||||
<img src="demo_images/Camera_interpolation_pointcloud.gif" alt="PointCloud Inpaint Video" width="40%" />
|
||||
|
||||
### 9. `sbs180_workflow.json`
|
||||
|
||||
Take a wide-angle (fisheye or equirectangular) high-resolution (e.g., 4096×4096) image and generate a stereo pair by moving the camera horizontally. The output is a wide-angle stereo pair (side-by-side), simulating a fisheye or equirectangular stereo camera.
|
||||
|
||||
<img src="demo_images/equirect_stereo.gif" alt="Equirectangular Stereo Demo" width="80%" />
|
||||
|
||||
### 10. `Pointcloud_walker.json`
|
||||
|
||||
Interactive Open3D-based GUI for walking and setting camera trajectory inside pointcloud.
|
||||
|
||||
### 11. `video_camera.json`
|
||||
|
||||
This workflow demonstrates camera trajectory movement using the `wan-vace` video inpainting model. It generates smooth camera movements along a trajectory while filling missing regions with high-quality inpainting.
|
||||
|
||||
<div style="display:flex; gap:10px;">
|
||||
<img src="demo_images/camera_movement.gif" alt="Camera Movement Demo" width="80%" />
|
||||
</div>
|
||||
|
||||
---
|
||||
|
||||
## Trajectory Concept
|
||||
|
||||
A **trajectory** in camera-comfyUI is a sequence of camera poses, each represented as a 4×4 transformation matrix. This set of matrices defines the path and orientation of the camera through 3D space, enabling smooth and complex camera movements for view synthesis, point cloud rendering, and video generation.
|
||||
|
||||
### Creating Trajectories
|
||||
|
||||
There are two main ways to create a trajectory:
|
||||
|
||||
- **Camera Matrices Interpolation:**
|
||||
Define two or more camera poses (as matrices), and interpolate between them to generate a smooth path. The `CameraInterpolationNode` automates this process, producing a trajectory tensor for use in camera motion nodes.
|
||||
|
||||
- **Walking in Open3D Environment:**
|
||||
Use the interactive Open3D GUI (`CameraTrajectoryNode`) to "walk" through the point cloud. As you move the camera, waypoints (poses) are recorded, forming a trajectory that can be exported and reused.
|
||||
|
||||
### Using Trajectories
|
||||
|
||||
The `CameraMotionNode` takes a trajectory (set of matrices) and interpolates camera positions and orientations along it, producing smooth camera movements for rendering sequences or videos.
|
||||
|
||||
---
|
||||
|
||||
## Point Cloud Formats
|
||||
|
||||
Point clouds can be saved and loaded in two formats:
|
||||
|
||||
- **.npy**: Numpy array format (fast, preserves all tensor data, recommended for internal pipelines).
|
||||
- **.ply**: Polygon File Format (widely supported, viewable in external 3D tools).
|
||||
|
||||
Use the `SavePointCloud` and `LoadPointCloud` nodes to handle I/O operations in either format.
|
||||
|
||||
---
|
||||
|
||||
## Contributing
|
||||
@@ -348,15 +202,11 @@ Contributions welcome! Please open issues or PRs to add features, improve docs,
|
||||
|
||||
## TODO List
|
||||
|
||||
* [x] Add processing to pointcloud or depthmap to remove outlier and lonely points at depth borders.
|
||||
* [ ] 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.
|
||||
* [x] Create a single workflow for view synthesis (`video_to_4d_world.json`).
|
||||
* [ ] Create a single workflow for view synthesis.
|
||||
* [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.
|
||||
* [ ] Fix imports for renamed folders (e.g., inpainting_flux)
|
||||
|
||||
+2
-24
@@ -3,28 +3,6 @@ from .reprojection_nodes import NODE_CLASS_MAPPINGS as NCM2
|
||||
from .metric_depth_nodes import NODE_CLASS_MAPPINGS as NCM3
|
||||
from .flux_fisheye_filling_nodes import NODE_CLASS_MAPPINGS as NCM4
|
||||
from .complex_nodes import NODE_CLASS_MAPPINGS as NCM5
|
||||
from .video_nodes import NODE_CLASS_MAPPINGS as NCM6
|
||||
from .GS_nodes import NODE_CLASS_MAPPINGS as NCM7
|
||||
NODE_CLASS_MAPPINGS = {**NCM1, **NCM2, **NCM3, **NCM4, **NCM5}
|
||||
|
||||
# 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"]
|
||||
__all__ = ["NODE_CLASS_MAPPINGS"]
|
||||
+376
-168
@@ -64,7 +64,7 @@ class FisheyeDepthEstimator:
|
||||
RETURN_TYPES = ("TENSOR","MASK")
|
||||
RETURN_NAMES = ("depthmap","mask")
|
||||
FUNCTION = "estimate_fisheye_depth"
|
||||
CATEGORY = "Camera/Depth"
|
||||
CATEGORY = "Camera/depth"
|
||||
|
||||
def estimate_fisheye_depth(
|
||||
self,
|
||||
@@ -95,76 +95,46 @@ class FisheyeDepthEstimator:
|
||||
depth_full, = de_node.estimate_depth(image, model_name, depth_scale)
|
||||
mask_full = (depth_full > 0).float()
|
||||
|
||||
# 2) Pinhole orientations (5 views)
|
||||
rotations = [
|
||||
(0, 0, 0), # front
|
||||
(0, 45, 0), # right
|
||||
(0, -45, 0), # left
|
||||
(45, 0, 0), # up
|
||||
(-45, 0, 0), # down
|
||||
]
|
||||
# 2) Generate Pinhole Views
|
||||
fisheye_depths, fisheye_masks = self._generate_pinhole_views(
|
||||
image,
|
||||
de_node, z2r_node, ri_node, rd_node,
|
||||
fisheye_fov, pinhole_fov,
|
||||
pin_w, pin_h, fish_w, fish_h,
|
||||
model_name, depth_scale, median_blur_kernel
|
||||
)
|
||||
|
||||
fisheye_depths = []
|
||||
fisheye_masks = []
|
||||
|
||||
# euler → matrix
|
||||
def euler_to_matrix(pitch, yaw, roll):
|
||||
p, y, r = map(math.radians, (pitch, yaw, roll))
|
||||
Rx = torch.tensor([[1,0,0],[0,math.cos(p),-math.sin(p)],[0,math.sin(p),math.cos(p)]], dtype=torch.float32)
|
||||
Ry = torch.tensor([[math.cos(y),0,math.sin(y)],[0,1,0],[-math.sin(y),0,math.cos(y)]], dtype=torch.float32)
|
||||
Rz = torch.tensor([[math.cos(r),-math.sin(r),0],[math.sin(r),math.cos(r),0],[0,0,1]], dtype=torch.float32)
|
||||
R = Rz @ Ry @ Rx
|
||||
M = torch.eye(4, dtype=torch.float32)
|
||||
M[:3, :3] = R
|
||||
return M
|
||||
|
||||
# 3) Process each orientation
|
||||
for pitch, yaw, roll in rotations:
|
||||
M = euler_to_matrix(pitch, yaw, roll)
|
||||
M_np = M.numpy()
|
||||
M_inv = torch.inverse(M).numpy()
|
||||
|
||||
# fisheye → pinhole
|
||||
img_pin, mask_pin = ri_node.reproject_image(
|
||||
image,
|
||||
input_horiszontal_fov = fisheye_fov,
|
||||
output_horiszontal_fov= pinhole_fov,
|
||||
input_projection = "FISHEYE",
|
||||
output_projection = "PINHOLE",
|
||||
output_width = pin_w,
|
||||
output_height = pin_h,
|
||||
transform_matrix = M_np,
|
||||
feathering = 0,
|
||||
)
|
||||
|
||||
# estimate pinhole depth
|
||||
depth_pin, = de_node.estimate_depth(img_pin, model_name, depth_scale, median_blur_kernel=median_blur_kernel)
|
||||
depth_pin, = z2r_node.depth_to_ray_depth(
|
||||
depth_pin,
|
||||
pinhole_fov,
|
||||
)
|
||||
# pinhole → fisheye
|
||||
fish_depth, fish_mask = rd_node.reproject_depth(
|
||||
depth_pin,
|
||||
input_horizontal_fov = pinhole_fov,
|
||||
output_horizontal_fov= fisheye_fov,
|
||||
input_projection = "PINHOLE",
|
||||
output_projection = "FISHEYE",
|
||||
output_width = fish_w,
|
||||
output_height = fish_h,
|
||||
transform_matrix = M_inv,
|
||||
)
|
||||
# squeeze mask to [B,H,W]
|
||||
fish_mask = fish_mask.squeeze(1)
|
||||
|
||||
fisheye_depths.append(fish_depth) # [B,H,W]
|
||||
fisheye_masks.append(fish_mask)
|
||||
fisheye_depths.append(depth_full) # [B,H,W]
|
||||
fisheye_masks.append(mask_full.squeeze(-1)) # [B,H,W 1]
|
||||
# merged mask
|
||||
merged_mask = torch.sum(torch.stack(fisheye_masks), dim=0) > 0.5
|
||||
# print(fisheye_depths[0].shape, fisheye_depths[-1].shape, merged_mask.shape)
|
||||
# 4) Merge in sequence
|
||||
d_acc, m_acc = self._merge_depths(
|
||||
fisheye_depths, fisheye_masks,
|
||||
ren_node, comb_node,
|
||||
mode, softmerge_radius
|
||||
)
|
||||
|
||||
# 5) Circular mask
|
||||
ys = torch.arange(fish_h, device=d_acc.device).view(1, fish_h, 1)
|
||||
xs = torch.arange(fish_w, device=d_acc.device).view(1, 1, fish_w)
|
||||
cy = (fish_h - 1) / 2.0
|
||||
cx = (fish_w - 1) / 2.0
|
||||
dist2 = (ys - cy)**2 + (xs - cx)**2
|
||||
radius2 = (min(fish_w, fish_h) / 2.0)**2
|
||||
circ_mask = (dist2 <= radius2).float()
|
||||
return d_acc, circ_mask
|
||||
|
||||
def _merge_depths(
|
||||
self,
|
||||
fisheye_depths: list,
|
||||
fisheye_masks: list,
|
||||
ren_node: DepthRenormalizer,
|
||||
comb_node: CombineDepthsNode,
|
||||
mode: str,
|
||||
softmerge_radius: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
d_acc = fisheye_depths[0]
|
||||
m_acc = fisheye_masks[0]
|
||||
for d_new, m_new in zip(fisheye_depths[1:-1], fisheye_masks[1:-1]):
|
||||
@@ -187,20 +157,120 @@ class FisheyeDepthEstimator:
|
||||
m_acc,
|
||||
d_norm,
|
||||
m_new_last,
|
||||
mode = "SRC",
|
||||
mode = "SRC", # Use SRC for the full fisheye to preserve its details
|
||||
invert_mask = False,
|
||||
softmerge_radius = softmerge_radius
|
||||
)
|
||||
return d_acc, m_acc
|
||||
|
||||
# 5) Circular mask
|
||||
ys = torch.arange(fish_h, device=d_acc.device).view(1, fish_h, 1)
|
||||
xs = torch.arange(fish_w, device=d_acc.device).view(1, 1, fish_w)
|
||||
cy = (fish_h - 1) / 2.0
|
||||
cx = (fish_w - 1) / 2.0
|
||||
dist2 = (ys - cy)**2 + (xs - cx)**2
|
||||
radius2 = (min(fish_w, fish_h) / 2.0)**2
|
||||
circ_mask = (dist2 <= radius2).float()
|
||||
return d_acc, circ_mask
|
||||
def _generate_pinhole_views(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
de_node: DepthEstimatorNode,
|
||||
z2r_node: ZDepthToRayDepthNode,
|
||||
ri_node: ReprojectImage,
|
||||
rd_node: ReprojectDepth,
|
||||
fisheye_fov: float,
|
||||
pinhole_fov: float,
|
||||
pin_w: int,
|
||||
pin_h: int,
|
||||
fish_w: int,
|
||||
fish_h: int,
|
||||
model_name: str,
|
||||
depth_scale: float,
|
||||
median_blur_kernel: int,
|
||||
) -> Tuple[list, list]:
|
||||
rotations = [
|
||||
(0, 0, 0), # front
|
||||
(0, 45, 0), # right
|
||||
(0, -45, 0), # left
|
||||
(45, 0, 0), # up
|
||||
(-45, 0, 0), # down
|
||||
]
|
||||
|
||||
fisheye_depths = []
|
||||
fisheye_masks = []
|
||||
|
||||
for pitch, yaw, roll in rotations:
|
||||
M = self._euler_to_matrix(pitch, yaw, roll)
|
||||
M_np = M.numpy()
|
||||
M_inv = torch.inverse(M).numpy()
|
||||
|
||||
fish_depth, fish_mask = self._process_view(
|
||||
image, M_np, M_inv,
|
||||
de_node, z2r_node, ri_node, rd_node,
|
||||
fisheye_fov, pinhole_fov,
|
||||
pin_w, pin_h, fish_w, fish_h,
|
||||
model_name, depth_scale, median_blur_kernel
|
||||
)
|
||||
|
||||
fisheye_depths.append(fish_depth)
|
||||
fisheye_masks.append(fish_mask)
|
||||
|
||||
return fisheye_depths, fisheye_masks
|
||||
|
||||
def _process_view(
|
||||
self,
|
||||
image: torch.Tensor,
|
||||
M_np: np.ndarray,
|
||||
M_inv: np.ndarray,
|
||||
de_node: DepthEstimatorNode,
|
||||
z2r_node: ZDepthToRayDepthNode,
|
||||
ri_node: ReprojectImage,
|
||||
rd_node: ReprojectDepth,
|
||||
fisheye_fov: float,
|
||||
pinhole_fov: float,
|
||||
pin_w: int,
|
||||
pin_h: int,
|
||||
fish_w: int,
|
||||
fish_h: int,
|
||||
model_name: str,
|
||||
depth_scale: float,
|
||||
median_blur_kernel: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# fisheye → pinhole
|
||||
img_pin, mask_pin = ri_node.reproject_image(
|
||||
image,
|
||||
input_horiszontal_fov = fisheye_fov,
|
||||
output_horiszontal_fov= pinhole_fov,
|
||||
input_projection = "FISHEYE",
|
||||
output_projection = "PINHOLE",
|
||||
output_width = pin_w,
|
||||
output_height = pin_h,
|
||||
transform_matrix = M_np,
|
||||
feathering = 0,
|
||||
)
|
||||
|
||||
# estimate pinhole depth
|
||||
depth_pin, = de_node.estimate_depth(img_pin, model_name, depth_scale, median_blur_kernel=median_blur_kernel)
|
||||
depth_pin, = z2r_node.depth_to_ray_depth(
|
||||
depth_pin,
|
||||
pinhole_fov,
|
||||
)
|
||||
# pinhole → fisheye
|
||||
fish_depth, fish_mask = rd_node.reproject_depth(
|
||||
depth_pin,
|
||||
input_horizontal_fov = pinhole_fov,
|
||||
output_horizontal_fov= fisheye_fov,
|
||||
input_projection = "PINHOLE",
|
||||
output_projection = "FISHEYE",
|
||||
output_width = fish_w,
|
||||
output_height = fish_h,
|
||||
transform_matrix = M_inv,
|
||||
)
|
||||
# squeeze mask to [B,H,W]
|
||||
fish_mask = fish_mask.squeeze(1)
|
||||
return fish_depth, fish_mask
|
||||
|
||||
def _euler_to_matrix(self, pitch, yaw, roll):
|
||||
p, y, r = map(math.radians, (pitch, yaw, roll))
|
||||
Rx = torch.tensor([[1,0,0],[0,math.cos(p),-math.sin(p)],[0,math.sin(p),math.cos(p)]], dtype=torch.float32)
|
||||
Ry = torch.tensor([[math.cos(y),0,math.sin(y)],[0,1,0],[-math.sin(y),0,math.cos(y)]], dtype=torch.float32)
|
||||
Rz = torch.tensor([[math.cos(r),-math.sin(r),0],[math.sin(r),math.cos(r),0],[0,0,1]], dtype=torch.float32)
|
||||
R = Rz @ Ry @ Rx
|
||||
M = torch.eye(4, dtype=torch.float32)
|
||||
M[:3, :3] = R
|
||||
return M
|
||||
|
||||
class PointcloudTrajectoryEnricher:
|
||||
"""
|
||||
@@ -241,7 +311,7 @@ class PointcloudTrajectoryEnricher:
|
||||
RETURN_TYPES = ("TENSOR","IMAGE","TENSOR")
|
||||
RETURN_NAMES = ("enriched_pointcloud","debug_image","debug_depth")
|
||||
FUNCTION = "enrich_trajectory"
|
||||
CATEGORY = "Camera/Trajectory"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def enrich_trajectory(
|
||||
self,
|
||||
@@ -288,106 +358,244 @@ class PointcloudTrajectoryEnricher:
|
||||
debug_img = torch.zeros((1, height, width, 3), device=device)
|
||||
debug_depth = torch.zeros((1, height, width, 1), device=device)
|
||||
enriched_pc = pointcloud
|
||||
# Initialize debug_img and debug_depth which will be updated in the loop
|
||||
# and will hold the values from the last processed view.
|
||||
debug_img = torch.zeros((1, height, width, 3), device=device)
|
||||
debug_depth = torch.zeros((1, height, width, 1), device=device)
|
||||
|
||||
# loop over trajectory (limit or full)
|
||||
for M in tqdm(trajectory[:15], desc="Enriching trajectory"):
|
||||
M_np = M.cpu().numpy()
|
||||
M_inv = np.linalg.inv(M_np)
|
||||
enriched_pc, view_debug_img, view_debug_depth = self._process_single_view(
|
||||
M, enriched_pc, device,
|
||||
proj_node, outpaint_node, depth_node, renorm_node,
|
||||
depth2pc_node, transform_node, clean_node, zdepth_node,
|
||||
camera_type, horizontal_fov, width, height,
|
||||
patch_projection, patch_horiz_fov, patch_res,
|
||||
patch_phi, patch_theta, prompt,
|
||||
num_inference_steps, guidance_scale, mask_blur,
|
||||
voxel_size, min_points_per_voxel, model_name
|
||||
)
|
||||
debug_img = view_debug_img
|
||||
debug_depth = view_debug_depth
|
||||
return enriched_pc, debug_img, debug_depth
|
||||
|
||||
# transform and select front points
|
||||
rotated, = transform_node.transform_pointcloud(enriched_pc, M_np)
|
||||
pc_front = rotated[rotated[:, 2] > 0]
|
||||
def _process_single_view(
|
||||
self,
|
||||
M_matrix: torch.Tensor,
|
||||
current_enriched_pc: torch.Tensor,
|
||||
device: torch.device,
|
||||
proj_node: ProjectPointCloud,
|
||||
outpaint_node: OutpaintAnyProjection,
|
||||
depth_node: DepthEstimatorNode,
|
||||
renorm_node: DepthRenormalizer,
|
||||
depth2pc_node: DepthToPointCloud,
|
||||
transform_node: TransformPointCloud,
|
||||
clean_node: PointCloudCleaner,
|
||||
zdepth_node: ZDepthToRayDepthNode,
|
||||
camera_type: str,
|
||||
horizontal_fov: float,
|
||||
width: int,
|
||||
height: int,
|
||||
patch_projection: str,
|
||||
patch_horiz_fov: float,
|
||||
patch_res: int,
|
||||
patch_phi: float,
|
||||
patch_theta: float,
|
||||
prompt: str,
|
||||
num_inference_steps: int,
|
||||
guidance_scale: float,
|
||||
mask_blur: int,
|
||||
voxel_size: float,
|
||||
min_points_per_voxel: int,
|
||||
model_name: str,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
M_np = M_matrix.cpu().numpy()
|
||||
M_inv = np.linalg.inv(M_np)
|
||||
|
||||
# clean front points
|
||||
pc_front, = clean_node.clean_pointcloud(
|
||||
pc_front,
|
||||
voxel_size=voxel_size,
|
||||
min_points_per_voxel=min_points_per_voxel,
|
||||
width=4096,
|
||||
height=4096,
|
||||
)
|
||||
img, mask, depth_map, pc_front = self._prepare_view_data(
|
||||
current_enriched_pc, M_np, transform_node, clean_node, proj_node,
|
||||
voxel_size, min_points_per_voxel, camera_type, horizontal_fov,
|
||||
width, height
|
||||
)
|
||||
|
||||
# project to image + depth
|
||||
img, mask, depth_map = proj_node.project_pointcloud(
|
||||
pc_front,
|
||||
camera_type,
|
||||
horizontal_fov,
|
||||
width,
|
||||
height,
|
||||
point_size=3,
|
||||
return_inverse_depth=False,
|
||||
)
|
||||
debug_img = img
|
||||
# fill nan in depthmap with (-1)
|
||||
# outpaint missing regions
|
||||
hole_mask = (mask < 0.5).float()
|
||||
out_img, out_mask = outpaint_node.outpaint_any(
|
||||
img,
|
||||
input_projection = camera_type,
|
||||
input_horiz_fov = horizontal_fov,
|
||||
output_projection = camera_type,
|
||||
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 = False,
|
||||
guidance_scale = guidance_scale,
|
||||
mask_blur = mask_blur,
|
||||
mask = hole_mask,
|
||||
debug = False,
|
||||
)
|
||||
debug_img = out_img
|
||||
# estimate and renormalize depth
|
||||
nan_mask = torch.isnan(depth_map)
|
||||
# …and replace them with –1.0 (in-place)
|
||||
depth_map[nan_mask] = 0
|
||||
# clip from -1 to 1000
|
||||
depth_map = torch.clamp(depth_map, 0, 1000.0)
|
||||
new_depth, = depth_node.estimate_depth(out_img, model_name, depth_scale=1.0)
|
||||
new_depth, = zdepth_node.depth_to_ray_depth(
|
||||
new_depth,
|
||||
horizontal_fov,
|
||||
)
|
||||
# renormalize depth
|
||||
norm_depth, = renorm_node.renormalize_depth(
|
||||
new_depth,
|
||||
depth_map,
|
||||
depth_mask=(mask>=0.5)*1,
|
||||
guidance_mask=(mask<0.5)*1,
|
||||
use_inverse=False,
|
||||
)
|
||||
# median blur on depth
|
||||
k = 5
|
||||
d = norm_depth.permute(0,3,1,2) # [B,1,H,W]
|
||||
pad = k//2
|
||||
pd = F.pad(d, (pad, pad, pad, pad), mode='reflect')
|
||||
patches = pd.unfold(2, k, 1).unfold(3, k, 1)
|
||||
patches = patches.contiguous().view(d.shape[0], d.shape[1], d.shape[2], d.shape[3], k*k)
|
||||
d, _ = patches.median(dim=-1)
|
||||
norm_depth = d.permute(0,2,3,1) # [B,H,W,1]
|
||||
debug_depth = norm_depth*hole_mask.unsqueeze(0).unsqueeze(-1)+depth_map*(1-hole_mask.unsqueeze(0).unsqueeze(-1))
|
||||
# fill nan in depthmap with (-1)
|
||||
# outpaint missing regions
|
||||
hole_mask = (mask < 0.5).float()
|
||||
out_img = self._outpaint_missing_regions(
|
||||
img, hole_mask, # Pass hole_mask instead of the full mask
|
||||
outpaint_node, camera_type, horizontal_fov, width, height,
|
||||
patch_projection, patch_horiz_fov, patch_res,
|
||||
patch_phi, patch_theta, prompt,
|
||||
num_inference_steps, guidance_scale, mask_blur
|
||||
)
|
||||
# estimate and renormalize depth
|
||||
norm_depth, debug_depth_view = self._estimate_and_refine_depth(
|
||||
out_img, depth_map, mask, hole_mask,
|
||||
depth_node, zdepth_node, renorm_node,
|
||||
model_name, horizontal_fov
|
||||
)
|
||||
|
||||
# back to pointcloud
|
||||
pc_new, = depth2pc_node.depth_to_pointcloud(
|
||||
out_img,
|
||||
camera_type,
|
||||
horizontal_fov,
|
||||
depth_scale=1.0,
|
||||
invert_depth=False,
|
||||
depthmap=norm_depth,
|
||||
mask=hole_mask,
|
||||
)
|
||||
# back to pointcloud
|
||||
pc_world = self._convert_depth_to_world_pointcloud(
|
||||
out_img, norm_depth, hole_mask, M_inv,
|
||||
depth2pc_node, transform_node,
|
||||
camera_type, horizontal_fov
|
||||
)
|
||||
# enriched_pc is not rotated
|
||||
current_enriched_pc = torch.cat([current_enriched_pc, pc_world.to(device)], dim=0)
|
||||
return current_enriched_pc, out_img, debug_depth_view # Return out_img and the depth for this view
|
||||
|
||||
pc_world, = transform_node.transform_pointcloud(pc_new, M_inv)
|
||||
# enriched_pc is not rotated
|
||||
enriched_pc = torch.cat([enriched_pc, pc_world.to(device)], dim=0)
|
||||
return enriched_pc, debug_img, norm_depth
|
||||
|
||||
def _convert_depth_to_world_pointcloud(
|
||||
self,
|
||||
out_img: torch.Tensor,
|
||||
norm_depth: torch.Tensor,
|
||||
hole_mask: torch.Tensor,
|
||||
M_inv: np.ndarray,
|
||||
depth2pc_node: DepthToPointCloud,
|
||||
transform_node: TransformPointCloud,
|
||||
camera_type: str,
|
||||
horizontal_fov: float,
|
||||
) -> torch.Tensor:
|
||||
pc_new, = depth2pc_node.depth_to_pointcloud(
|
||||
out_img,
|
||||
camera_type,
|
||||
horizontal_fov,
|
||||
depth_scale=1.0,
|
||||
invert_depth=False,
|
||||
depthmap=norm_depth,
|
||||
mask=hole_mask,
|
||||
)
|
||||
pc_world, = transform_node.transform_pointcloud(pc_new, M_inv)
|
||||
return pc_world
|
||||
|
||||
def _estimate_and_refine_depth(
|
||||
self,
|
||||
out_img: torch.Tensor,
|
||||
depth_map: torch.Tensor,
|
||||
original_mask: torch.Tensor, # Mask from projection
|
||||
hole_mask: torch.Tensor,
|
||||
depth_node: DepthEstimatorNode,
|
||||
zdepth_node: ZDepthToRayDepthNode,
|
||||
renorm_node: DepthRenormalizer,
|
||||
model_name: str,
|
||||
horizontal_fov: float,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
nan_mask = torch.isnan(depth_map)
|
||||
depth_map[nan_mask] = 0 # In-place modification
|
||||
depth_map = torch.clamp(depth_map, 0, 1000.0)
|
||||
|
||||
new_depth, = depth_node.estimate_depth(out_img, model_name, depth_scale=1.0)
|
||||
new_depth, = zdepth_node.depth_to_ray_depth(
|
||||
new_depth,
|
||||
horizontal_fov,
|
||||
)
|
||||
# renormalize depth
|
||||
# Use original_mask for depth_mask as it represents valid projected areas
|
||||
norm_depth, = renorm_node.renormalize_depth(
|
||||
new_depth,
|
||||
depth_map,
|
||||
depth_mask=(original_mask >= 0.5) * 1,
|
||||
guidance_mask=(hole_mask >= 0.5) * 1, # hole_mask is appropriate here
|
||||
use_inverse=False,
|
||||
)
|
||||
# median blur on depth
|
||||
k = 5
|
||||
d = norm_depth.permute(0,3,1,2) # [B,1,H,W]
|
||||
pad = k//2
|
||||
pd = F.pad(d, (pad, pad, pad, pad), mode='reflect')
|
||||
patches = pd.unfold(2, k, 1).unfold(3, k, 1)
|
||||
patches = patches.contiguous().view(d.shape[0], d.shape[1], d.shape[2], d.shape[3], k*k)
|
||||
d_median, _ = patches.median(dim=-1) # Renamed to avoid conflict
|
||||
norm_depth_blurred = d_median.permute(0,2,3,1) # [B,H,W,1]
|
||||
|
||||
# Create debug_depth_view using the blurred normalized depth for holes
|
||||
# and the original depth_map for non-holes.
|
||||
debug_depth_view = norm_depth_blurred * hole_mask.unsqueeze(0).unsqueeze(-1) + \
|
||||
depth_map * (1 - hole_mask.unsqueeze(0).unsqueeze(-1))
|
||||
|
||||
return norm_depth_blurred, debug_depth_view
|
||||
|
||||
|
||||
def _outpaint_missing_regions(
|
||||
self,
|
||||
img: torch.Tensor,
|
||||
hole_mask: torch.Tensor, # Expects the specific hole_mask
|
||||
outpaint_node: OutpaintAnyProjection,
|
||||
camera_type: str,
|
||||
horizontal_fov: float,
|
||||
width: int,
|
||||
height: int,
|
||||
patch_projection: str,
|
||||
patch_horiz_fov: float,
|
||||
patch_res: int,
|
||||
patch_phi: float,
|
||||
patch_theta: float,
|
||||
prompt: str,
|
||||
num_inference_steps: int,
|
||||
guidance_scale: float,
|
||||
mask_blur: int,
|
||||
) -> torch.Tensor: # Returns only out_img, out_mask is not used later
|
||||
out_img, _ = outpaint_node.outpaint_any( # Assign out_mask to _
|
||||
img,
|
||||
input_projection = camera_type,
|
||||
input_horiz_fov = horizontal_fov,
|
||||
output_projection = camera_type,
|
||||
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 = False,
|
||||
guidance_scale = guidance_scale,
|
||||
mask_blur = mask_blur,
|
||||
mask = hole_mask, # Use the passed hole_mask
|
||||
debug = False,
|
||||
)
|
||||
return out_img
|
||||
|
||||
def _prepare_view_data(
|
||||
self,
|
||||
current_enriched_pc: torch.Tensor,
|
||||
M_np: np.ndarray,
|
||||
transform_node: TransformPointCloud,
|
||||
clean_node: PointCloudCleaner,
|
||||
proj_node: ProjectPointCloud,
|
||||
voxel_size: float,
|
||||
min_points_per_voxel: int,
|
||||
camera_type: str,
|
||||
horizontal_fov: float,
|
||||
width: int,
|
||||
height: int,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
# transform and select front points
|
||||
rotated, = transform_node.transform_pointcloud(current_enriched_pc, M_np)
|
||||
pc_front = rotated[rotated[:, 2] > 0]
|
||||
|
||||
# clean front points
|
||||
pc_front, = clean_node.clean_pointcloud(
|
||||
pc_front,
|
||||
voxel_size=voxel_size,
|
||||
min_points_per_voxel=min_points_per_voxel,
|
||||
width=4096, # Consider passing these as params if they vary
|
||||
height=4096, # Consider passing these as params if they vary
|
||||
)
|
||||
|
||||
# project to image + depth
|
||||
img, mask, depth_map = proj_node.project_pointcloud(
|
||||
pc_front,
|
||||
camera_type,
|
||||
horizontal_fov,
|
||||
width,
|
||||
height,
|
||||
point_size=3,
|
||||
return_inverse_depth=False,
|
||||
)
|
||||
return img, mask, depth_map, pc_front
|
||||
|
||||
NODE_CLASS_MAPPINGS = {"FisheyeDepthEstimator": FisheyeDepthEstimator,
|
||||
"PointcloudTrajectoryEnricher": PointcloudTrajectoryEnricher}
|
||||
|
||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
|
Before Width: | Height: | Size: 12 MiB |
Binary file not shown.
|
Before Width: | Height: | Size: 224 KiB |
Binary file not shown.
|
Before Width: | Height: | Size: 13 MiB |
@@ -1,127 +0,0 @@
|
||||
# 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.*
|
||||
@@ -129,7 +129,7 @@ class OutpaintAnyProjection:
|
||||
patch_mask = torch.ones_like(patch_mask) * 1 if debug else patch_mask
|
||||
# 5) Reproject inpainted patch back
|
||||
|
||||
back_img, back_mask = reproj.reproject_image(
|
||||
back_img, inpainted_patch_reproj_mask_raw = reproj.reproject_image( # Store raw mask output
|
||||
inpainted_patch,
|
||||
patch_horiz_fov, output_horiz_fov,
|
||||
patch_projection, output_projection,
|
||||
@@ -138,19 +138,37 @@ class OutpaintAnyProjection:
|
||||
transform_matrix=rot_m,
|
||||
feathering=0,
|
||||
)
|
||||
back_mask = normalize_mask(back_mask).bool() # True where patch contributes
|
||||
# inpainted_patch_coverage_mask: 1.0 where the reprojected inpainted patch has content, 0.0 otherwise.
|
||||
inpainted_patch_coverage_mask = normalize_mask(back_mask) # Ensure it's float [0,1]
|
||||
|
||||
# original coverage: True = had data, False = hole
|
||||
orig_covered = ~base_mask.bool()
|
||||
base_img=base_img * orig_covered.unsqueeze(-1)
|
||||
# fill only holes where back_mask is False
|
||||
filled = back_img * (~back_mask.unsqueeze(-1))*base_mask.unsqueeze(-1)
|
||||
final_img = base_img+filled
|
||||
# --- Define Masks based on Conventions ---
|
||||
# base_mask: 1.0 where original reprojected image has content, 0.0 for holes.
|
||||
# initial_hole_mask: 1.0 where original reprojected image has holes (inverse of base_mask).
|
||||
initial_hole_mask = 1.0 - base_mask
|
||||
# inpainted_patch_coverage_mask: 1.0 where reprojected inpainted patch has content.
|
||||
|
||||
# anything that’s still a hole after back‐projection needs inpaint
|
||||
needs_inpaint = (~((orig_covered) | (~back_mask))).to(torch.float32)
|
||||
# --- Compositing Logic ---
|
||||
# Goal: Inpainted patch takes precedence in overlapping areas. Original content is used elsewhere.
|
||||
|
||||
return final_img, needs_inpaint
|
||||
# Contribution from the original image:
|
||||
# Valid original pixels, excluding areas covered by the inpainted patch.
|
||||
original_content_contribution = base_img * base_mask.unsqueeze(-1) * \
|
||||
(1.0 - inpainted_patch_coverage_mask.unsqueeze(-1))
|
||||
|
||||
# Contribution from the inpainted patch (reprojected as back_img):
|
||||
# Valid inpainted pixels, where the patch provides coverage.
|
||||
inpainted_patch_contribution = back_img * inpainted_patch_coverage_mask.unsqueeze(-1)
|
||||
|
||||
# Combine:
|
||||
final_img = original_content_contribution + inpainted_patch_contribution
|
||||
|
||||
# --- needs_inpaint_mask Derivation ---
|
||||
# Identifies areas that were initially holes AND remain un-filled by the reprojected inpainted patch.
|
||||
# These are areas that still require inpainting if a further pass was to be made.
|
||||
not_covered_by_inpainted_patch = 1.0 - inpainted_patch_coverage_mask
|
||||
needs_inpaint_mask = initial_hole_mask * not_covered_by_inpainted_patch # Element-wise multiplication (AND logic)
|
||||
|
||||
return final_img, needs_inpaint_mask
|
||||
|
||||
# register
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
|
||||
-162
@@ -1,162 +0,0 @@
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
|
||||
# ----------------------------------------
|
||||
# Functions
|
||||
# ----------------------------------------
|
||||
install_pytorch() {
|
||||
echo "==> Installing PyTorch, TorchVision, TorchAudio, bitsandbytes, accelerate…"
|
||||
pip3 install -U torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128
|
||||
pip3 install -U bitsandbytes accelerate
|
||||
}
|
||||
|
||||
install_system_deps() {
|
||||
echo "==> Updating apt and installing system packages…"
|
||||
sudo apt-get update
|
||||
sudo apt-get install -y build-essential ffmpeg libsm6 libxext6 python3.10-dev
|
||||
}
|
||||
|
||||
clone_and_install_comfyui() {
|
||||
echo "==> Cloning ComfyUI…"
|
||||
git clone https://github.com/comfyanonymous/ComfyUI.git
|
||||
echo "==> Installing ComfyUI Python requirements…"
|
||||
pip3 install -r ComfyUI/requirements.txt
|
||||
}
|
||||
|
||||
install_camera_node() {
|
||||
echo "==> Installing camera‑ComfyUI node…"
|
||||
mkdir -p ComfyUI/custom_nodes
|
||||
git clone https://github.com/Alexankharin/camera-comfyUI.git \
|
||||
ComfyUI/custom_nodes/camera-comfyUI
|
||||
pip3 install -r ComfyUI/custom_nodes/camera-comfyUI/requirements.txt
|
||||
}
|
||||
|
||||
install_image_filters() {
|
||||
echo "==> Installing Image‑Filters node…"
|
||||
mkdir -p ComfyUI/custom_nodes
|
||||
git clone https://github.com/spacepxl/ComfyUI-Image-Filters.git \
|
||||
ComfyUI/custom_nodes/ComfyUI-Image-Filters
|
||||
pip3 install -r ComfyUI/custom_nodes/ComfyUI-Image-Filters/requirements.txt
|
||||
}
|
||||
|
||||
clone_flux_inpainting() {
|
||||
echo "==> Installing Flux‑Inpainting node…"
|
||||
mkdir -p ComfyUI/custom_nodes
|
||||
git clone https://github.com/rubi-du/ComfyUI-Flux-Inpainting.git \
|
||||
ComfyUI/custom_nodes/inpainting_flux
|
||||
}
|
||||
|
||||
install_metric_video_depth_anything() {
|
||||
echo "==> Installing Metric Video Depth Anything…"
|
||||
# Clone into ComfyUI root, not custom_nodes
|
||||
git clone https://github.com/DepthAnything/Video-Depth-Anything.git \
|
||||
ComfyUI/Video-Depth-Anything
|
||||
|
||||
echo " • Installing easydict…"
|
||||
pip3 install -U easydict
|
||||
|
||||
echo " • Copying util.py…"
|
||||
mkdir -p ComfyUI/utils
|
||||
cp ComfyUI/Video-Depth-Anything/metric_depth/utils/util.py \
|
||||
ComfyUI/utils/util.py
|
||||
|
||||
echo " • Downloading Metric Video Depth checkpoint…"
|
||||
mkdir -p ComfyUI/models/checkpoints
|
||||
wget -q -O ComfyUI/models/checkpoints/metric_video_depth_anything_vitl.pth \
|
||||
"https://huggingface.co/depth-anything/Metric-Video-Depth-Anything-Large/resolve/main/metric_video_depth_anything_vitl.pth"
|
||||
}
|
||||
|
||||
|
||||
install_comfyui_manager() {
|
||||
echo "==> Installing ComfyUI-Manager extension…"
|
||||
mkdir -p ComfyUI/custom_nodes
|
||||
git clone https://github.com/Comfy-Org/ComfyUI-Manager.git \
|
||||
ComfyUI/custom_nodes/ComfyUI-Manager
|
||||
pip3 install -r ComfyUI/custom_nodes/ComfyUI-Manager/requirements.txt
|
||||
}
|
||||
|
||||
install_hf_hub() {
|
||||
echo "==> Installing huggingface_hub…"
|
||||
pip3 install -U huggingface_hub
|
||||
}
|
||||
|
||||
download_vae_models() {
|
||||
echo "==> Downloading WAN‑VACE models…"
|
||||
mkdir -p ComfyUI/models/vae
|
||||
wget -q -O ComfyUI/models/vae/wan_2.1_vae.safetensors \
|
||||
"https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/resolve/main/split_files/vae/wan_2.1_vae.safetensors?download=true"
|
||||
|
||||
mkdir -p ComfyUI/models/text_encoders
|
||||
wget -q -O ComfyUI/models/text_encoders/umt5_xxl_fp16.safetensors \
|
||||
"https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/resolve/main/split_files/text_encoders/umt5_xxl_fp16.safetensors?download=true"
|
||||
|
||||
mkdir -p ComfyUI/models/diffusion_models
|
||||
wget -q -O ComfyUI/models/diffusion_models/wan2.1_vace_14B_fp16.safetensors \
|
||||
"https://huggingface.co/Comfy-Org/Wan_2.1_ComfyUI_repackaged/resolve/main/split_files/diffusion_models/wan2.1_vace_14B_fp16.safetensors"
|
||||
}
|
||||
|
||||
login_hf_hub() {
|
||||
echo "==> Hugging Face login…"
|
||||
huggingface-cli login
|
||||
}
|
||||
|
||||
# ----------------------------------------
|
||||
# Main
|
||||
# ----------------------------------------
|
||||
MODE="${1:-install}"
|
||||
|
||||
case "$MODE" in
|
||||
modules)
|
||||
install_pytorch
|
||||
install_system_deps
|
||||
clone_and_install_comfyui
|
||||
install_camera_node
|
||||
install_image_filters
|
||||
install_comfyui_manager
|
||||
install_hf_hub
|
||||
;;
|
||||
|
||||
flux)
|
||||
clone_flux_inpainting
|
||||
;;
|
||||
|
||||
vae)
|
||||
download_vae_models
|
||||
;;
|
||||
|
||||
depth)
|
||||
install_metric_video_depth_anything
|
||||
;;
|
||||
|
||||
install)
|
||||
install_pytorch
|
||||
install_system_deps
|
||||
clone_and_install_comfyui
|
||||
install_camera_node
|
||||
clone_flux_inpainting
|
||||
install_image_filters
|
||||
install_comfyui_manager
|
||||
install_hf_hub
|
||||
install_metric_video_depth_anything
|
||||
;;
|
||||
|
||||
all)
|
||||
install_pytorch
|
||||
install_system_deps
|
||||
clone_and_install_comfyui
|
||||
install_camera_node
|
||||
clone_flux_inpainting
|
||||
install_image_filters
|
||||
install_comfyui_manager
|
||||
install_hf_hub
|
||||
download_vae_models
|
||||
install_metric_video_depth_anything
|
||||
login_hf_hub
|
||||
echo "✅ All done!"
|
||||
;;
|
||||
|
||||
*)
|
||||
echo "Usage: $0 {install|modules|flux|vae|depth|all}"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
@@ -41,7 +41,7 @@ class DepthEstimatorNode:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("depth tensor",)
|
||||
FUNCTION = "estimate_depth"
|
||||
CATEGORY = "Camera/Depth"
|
||||
CATEGORY = "Camera/depth"
|
||||
|
||||
def estimate_depth(
|
||||
self,
|
||||
@@ -110,7 +110,7 @@ class DepthToImageNode:
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("depth image",)
|
||||
FUNCTION = "depth_to_image"
|
||||
CATEGORY = "Camera/Depth"
|
||||
CATEGORY = "Camera/depth"
|
||||
|
||||
def depth_to_image(
|
||||
self,
|
||||
@@ -158,7 +158,7 @@ class ZDepthToRayDepthNode:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("ray depth",)
|
||||
FUNCTION = "depth_to_ray_depth"
|
||||
CATEGORY = "Camera/Depth"
|
||||
CATEGORY = "Camera/depth"
|
||||
|
||||
def depth_to_ray_depth(
|
||||
self,
|
||||
@@ -229,7 +229,7 @@ class CombineDepthsNode:
|
||||
RETURN_TYPES = ("TENSOR","MASK")
|
||||
RETURN_NAMES = ("combined_depth","combined_mask")
|
||||
FUNCTION = "combine_depths"
|
||||
CATEGORY = "Camera/Depth"
|
||||
CATEGORY = "Camera/depth"
|
||||
|
||||
def combine_depths(
|
||||
self,
|
||||
@@ -358,7 +358,7 @@ class DepthRenormalizer:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("depth tensor",)
|
||||
FUNCTION = "renormalize_depth"
|
||||
CATEGORY = "Camera/Depth"
|
||||
CATEGORY = "Camera/depth"
|
||||
|
||||
def renormalize_depth(
|
||||
self,
|
||||
|
||||
@@ -1,484 +0,0 @@
|
||||
"""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
+152
-454
@@ -7,10 +7,7 @@ import os
|
||||
import folder_paths
|
||||
import logging
|
||||
import hashlib
|
||||
try:
|
||||
from kornia.filters import median_blur
|
||||
except ImportError: # kornia is optional; median_blur is not used in this module
|
||||
median_blur = None
|
||||
from kornia.filters import median_blur
|
||||
|
||||
from tqdm import tqdm
|
||||
# Try importing open3d and its visualization modules; log a warning if not found
|
||||
@@ -116,7 +113,7 @@ def XYZ_to_equirect(X: torch.Tensor, Y: torch.Tensor, Z: torch.Tensor, fov: floa
|
||||
Convert XYZ coordinates to normalized UV and depth using equirectangular projection.
|
||||
"""
|
||||
# full 360°×180°
|
||||
fov_rad = math.radians(fov) / 2
|
||||
fov_rad = math.radians(fov)
|
||||
depth = torch.sqrt(X**2 + Y**2 + Z**2)
|
||||
lon = torch.atan2(X, Z) # –π → +π
|
||||
lat = torch.asin(Y / depth) # –π/2 → +π/2
|
||||
@@ -139,116 +136,6 @@ 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:
|
||||
"""
|
||||
@@ -280,7 +167,7 @@ class DepthToPointCloud:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("pointcloud",)
|
||||
FUNCTION = "depth_to_pointcloud"
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def depth_to_pointcloud(
|
||||
self,
|
||||
@@ -386,7 +273,7 @@ class TransformPointCloud:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("transformed pointcloud",)
|
||||
FUNCTION = "transform_pointcloud"
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def transform_pointcloud(
|
||||
self,
|
||||
@@ -434,169 +321,124 @@ class ProjectPointCloud:
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "TENSOR")
|
||||
RETURN_NAMES = ("image", "mask", "depth")
|
||||
FUNCTION = "project_pointcloud"
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def project_pointcloud(
|
||||
self,
|
||||
pointcloud: torch.Tensor,
|
||||
pointcloud: torch.Tensor,
|
||||
output_projection: str,
|
||||
output_horizontal_fov: float,
|
||||
output_width: int,
|
||||
output_width: int,
|
||||
output_height: int,
|
||||
point_size: int = 1,
|
||||
point_size: int = 1,
|
||||
return_inverse_depth: bool = False,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Projects an (N×6) XYZRGB point cloud into an image,
|
||||
fills occlusion holes robustly, and returns:
|
||||
• img: [1,H,W,3] RGB image
|
||||
• mask: [H,W] foreground mask
|
||||
• depth: [1,H,W,1] depth (or inverse depth)
|
||||
"""
|
||||
device = pointcloud.device
|
||||
xyz, rgb_raw = pointcloud[:, :3], pointcloud[:, 3:6].float()
|
||||
coords = pointcloud[:, :3]
|
||||
colors = pointcloud[:, 3:].float()
|
||||
|
||||
# 1) Keep only points in front of camera
|
||||
in_front = xyz[:, 2] > 0
|
||||
xyz, rgb_raw = xyz[in_front], rgb_raw[in_front]
|
||||
# 1) Filter points in front of the camera
|
||||
mask_front = coords[:, 2] > 0
|
||||
coords = coords[mask_front]
|
||||
colors = colors[mask_front]
|
||||
|
||||
# 2) Project to normalized UV + depth
|
||||
X, Y, Z = xyz.unbind(1)
|
||||
X, Y, Z = coords.unbind(1)
|
||||
if output_projection == "PINHOLE":
|
||||
u, v, d = XYZ_to_pinhole(X, Y, Z, output_horizontal_fov)
|
||||
u, v, depth = XYZ_to_pinhole(X, Y, Z, output_horizontal_fov)
|
||||
elif output_projection == "FISHEYE":
|
||||
u, v, d = XYZ_to_fisheye(X, Y, Z, output_horizontal_fov)
|
||||
u, v, depth = XYZ_to_fisheye(X, Y, Z, output_horizontal_fov)
|
||||
else:
|
||||
u, v, d = XYZ_to_equirect(X, Y, Z, output_horizontal_fov)
|
||||
u, v, depth = XYZ_to_equirect(X, Y, Z, output_horizontal_fov)
|
||||
|
||||
# 3) Rasterize to pixel indices
|
||||
W, H = output_width, output_height
|
||||
ix = ((u * 0.5 + 0.5) * (W - 1)).round().clamp(0, W - 1).long()
|
||||
iy = ((v * 0.5 + 0.5) * (H - 1)).round().clamp(0, H - 1).long()
|
||||
pix = iy * W + ix
|
||||
valid = (pix >= 0) & (pix < W * H)
|
||||
pix, d, rgb_raw = pix[valid], d[valid], rgb_raw[valid]
|
||||
px = (u * (output_width - 1) / 2) + (output_width - 1) / 2
|
||||
py = (v * (output_height - 1) / 2) + (output_height - 1) / 2
|
||||
ix = px.round().clamp(0, output_width - 1).long()
|
||||
iy = py.round().clamp(0, output_height - 1).long()
|
||||
pix = iy * output_width + ix
|
||||
M = output_width * output_height
|
||||
|
||||
M = W * H
|
||||
# 3a) Front‐layer (nearest) depth
|
||||
z1 = torch.full((M,), float('inf'), device=device)
|
||||
z1.scatter_reduce_(0, pix, d, reduce='amin', include_self=True)
|
||||
# —— NEW: drop any invalid / NaN→int_min projections ——
|
||||
valid = (pix >= 0) & (pix < M)
|
||||
depth = depth[valid]
|
||||
colors = colors[valid]
|
||||
pix = pix[valid]
|
||||
order = torch.arange(depth.size(0), device=device)
|
||||
# rebuild your "order" to match
|
||||
|
||||
# 3b) Second‐layer (background) depth
|
||||
farther = d > z1[pix]
|
||||
pix2, d2 = pix[farther], d[farther]
|
||||
z2 = torch.full((M,), float('inf'), device=device)
|
||||
z2.scatter_reduce_(0, pix2, d2, reduce='amin', include_self=True)
|
||||
# 4) Allocate or reuse buffers
|
||||
if not hasattr(self, '_z_front') or self._z_front.numel() != M:
|
||||
self._z_front = torch.empty((M,), device=device)
|
||||
self._z_back = torch.empty((M,), device=device)
|
||||
self._idx = torch.full((M,), -1, dtype=torch.long, device=device)
|
||||
self._flat = torch.zeros((M, 4), device=device)
|
||||
z_front = self._z_front
|
||||
z_back = self._z_back
|
||||
idxbuf = self._idx
|
||||
flat = self._flat
|
||||
|
||||
# 3c) Foreground colour (from z1)
|
||||
keep = d == z1[pix]
|
||||
rgb = torch.zeros((M, 3), device=device)
|
||||
rgb[pix[keep]] = rgb_raw[keep].clamp(0, 255)
|
||||
# 5) Front z-buffer pass (nearest)
|
||||
z_front.fill_(float('inf'))
|
||||
z_front.scatter_reduce_(0, pix, depth, reduce='amin', include_self=True)
|
||||
sel_front = depth == z_front[pix]
|
||||
order = torch.arange(depth.size(0), device=device)
|
||||
order_m = torch.where(sel_front, order, depth.size(0))
|
||||
idxbuf.fill_(depth.size(0))
|
||||
idxbuf.scatter_reduce_(0, pix, order_m, reduce='amin', include_self=True)
|
||||
win_front = order == idxbuf[pix]
|
||||
|
||||
# reshape to image
|
||||
rgb = rgb.view(H, W, 3)
|
||||
z1 = z1.view(H, W)
|
||||
z2 = z2.view(H, W)
|
||||
fg_mask = z1 < float('inf') # has front hit
|
||||
occl = (z2 < float('inf')) # has any back hit
|
||||
rear_only = occl & ~fg_mask # true holes
|
||||
flat.fill_(0)
|
||||
flat[pix[win_front]] = colors[win_front]
|
||||
img4 = flat.view(output_height, output_width, 4)
|
||||
rgb = img4[..., :3].clamp(0, 255)
|
||||
alpha = (img4[..., 3] > 0).float()
|
||||
rgb *= alpha.unsqueeze(-1)
|
||||
depth_img = z_front.view(output_height, output_width)
|
||||
rgb_HR = rgb
|
||||
|
||||
# ── Iterative ring‐based in‐painting of rear‐only pixels ─────────────────────
|
||||
ker3 = torch.ones((1,1,3,3), device=device)
|
||||
ker3c = ker3.repeat(3,1,1,1)
|
||||
for _ in range(max(W, H)):
|
||||
# find rear_only pixels adjacent to current FG
|
||||
neigh = (
|
||||
F.max_pool2d(fg_mask.float()[None,None], 3, 1, 1).bool()[0,0]
|
||||
& ~fg_mask
|
||||
)
|
||||
to_fill = rear_only & neigh
|
||||
if not to_fill.any():
|
||||
break
|
||||
|
||||
# average depth + colour from current FG frontier
|
||||
d_t = z1.masked_fill(~fg_mask, 0)[None,None]
|
||||
c_t = rgb.permute(2,0,1)[None] # [1,3,H,W]
|
||||
m_t = fg_mask.float()[None,None]
|
||||
|
||||
sum_d = F.conv2d(d_t * m_t, ker3, padding=1)
|
||||
cnt_d = F.conv2d(m_t, ker3, padding=1).clamp(min=1)
|
||||
sum_c = F.conv2d(c_t * m_t, ker3c, padding=1, groups=3)
|
||||
cnt_c = cnt_d.repeat(1,3,1,1)
|
||||
|
||||
avg_d = (sum_d / cnt_d).squeeze()
|
||||
avg_c = (sum_c / cnt_c).squeeze().permute(1,2,0)
|
||||
|
||||
z1[to_fill] = avg_d[to_fill]
|
||||
rgb[to_fill] = avg_c[to_fill]
|
||||
fg_mask[to_fill] = True
|
||||
rear_only[to_fill] = False
|
||||
|
||||
# ── Depth‐aware generic hole closure ─────────────────────────────────────────
|
||||
# close_rad: radius of hole to close; depth_eps: depth jump tolerance
|
||||
close_rad = max(1, point_size // 2)
|
||||
depth_eps = 0.015
|
||||
pad = close_rad
|
||||
k = 2 * close_rad + 1
|
||||
ker = torch.ones((1,1,k,k), device=device)
|
||||
kerc = ker.repeat(3,1,1,1)
|
||||
front_t = fg_mask.float()[None,None]
|
||||
|
||||
# binary closing: dilate then erode
|
||||
D = F.max_pool2d(front_t, k, 1, pad)
|
||||
E = 1 - F.max_pool2d(1 - D, k, 1, pad)
|
||||
small_hole = E[0,0].bool() & ~fg_mask
|
||||
if small_hole.any():
|
||||
# compute local mean depth of FG
|
||||
z_t = z1.masked_fill(~fg_mask, 0)[None,None]
|
||||
cnt = F.conv2d(front_t, ker, padding=pad).clamp(min=1)
|
||||
z_avg = (F.conv2d(z_t, ker, padding=pad) / cnt)[0,0]
|
||||
|
||||
# depth‐range test
|
||||
z_near = F.max_pool2d(z1[None,None], 3,1,1)[0,0]
|
||||
z_far = -F.max_pool2d(-z1[None,None],3,1,1)[0,0]
|
||||
flat = (z_far - z_near) / z_avg.clamp(min=1e-6) < depth_eps
|
||||
|
||||
final = small_hole & flat
|
||||
if final.any():
|
||||
sum_d = F.conv2d(z_t, ker, padding=pad)
|
||||
sum_c = F.conv2d(rgb.permute(2,0,1)[None] * front_t, kerc,
|
||||
padding=pad, groups=3)
|
||||
avg_d = (sum_d / cnt)[0,0]
|
||||
avg_c = (sum_c / cnt.repeat(1,3,1,1))[0].permute(1,2,0)
|
||||
|
||||
z1[final] = avg_d[final]
|
||||
rgb[final] = avg_c[final]
|
||||
fg_mask[final] = True
|
||||
|
||||
# ── Optional morphological blur for larger point_size ───────────────────────
|
||||
# 6) Back z-buffer pass (farthest) for hole-filling
|
||||
if point_size > 1:
|
||||
r = point_size // 2
|
||||
k = 2 * r + 1
|
||||
pad = r
|
||||
ker = torch.ones((1,1,k,k), device=device)
|
||||
kerc = ker.repeat(3,1,1,1)
|
||||
d_t = z1[None,None]
|
||||
c_t = rgb.permute(2,0,1)[None]
|
||||
m_t = fg_mask.float()[None,None]
|
||||
z_back.fill_(-float('inf'))
|
||||
z_back.scatter_reduce_(0, pix, depth, reduce='amax', include_self=True)
|
||||
sel_back = depth == z_back[pix]
|
||||
order_m = torch.where(sel_back, order, -1)
|
||||
idxbuf.fill_(-1)
|
||||
idxbuf.scatter_reduce_(0, pix, order_m, reduce='amax', include_self=True)
|
||||
win_back = idxbuf[pix] >= 0
|
||||
|
||||
z1 = (F.conv2d(d_t * m_t, ker, padding=pad) /
|
||||
F.conv2d(m_t, ker, padding=pad).clamp(min=1)).squeeze()
|
||||
rgb = (F.conv2d(c_t * m_t, kerc, padding=pad, groups=3) /
|
||||
F.conv2d(m_t, ker, padding=pad).repeat(1,3,1,1).clamp(min=1)
|
||||
).squeeze().permute(1,2,0)
|
||||
flat.fill_(0)
|
||||
flat[pix[win_back]] = colors[win_back]
|
||||
back4 = flat.view(output_height, output_width, 4)
|
||||
rgb_back = back4[..., :3].clamp(0,255)
|
||||
alpha_back = (back4[..., 3] > 0).float()
|
||||
|
||||
# 10) Pack outputs
|
||||
img = rgb.unsqueeze(0) # [1,H,W,3]
|
||||
mask = fg_mask.float() # [H,W]
|
||||
depth = z1.unsqueeze(0).unsqueeze(-1) # [1,H,W,1]
|
||||
# fill holes where front missed
|
||||
hole = (alpha == 0) & (alpha_back > 0)
|
||||
rgb[hole] = rgb_back[hole]
|
||||
alpha[hole] = 1.0
|
||||
depth_img[hole] = z_back.view(output_height, output_width)[hole]
|
||||
|
||||
# 7) Median-filter _only_ in hole regions
|
||||
if hole.any():
|
||||
# prepare for kornia median_blur: [B,C,H,W]
|
||||
rgb_t = rgb.permute(2,0,1).unsqueeze(0) # [1,3,H,W]
|
||||
# apply median filter
|
||||
rgb_med = median_blur(rgb_t, (point_size, point_size))
|
||||
# back to HWC
|
||||
rgb_med = rgb_med.squeeze(0).permute(1,2,0)
|
||||
# merge only at hole locations
|
||||
rgb[hole] = rgb_med[hole]
|
||||
# alpha already set to 1.0 for holes
|
||||
|
||||
# 8) Pack and return with original script shapes
|
||||
img = rgb.unsqueeze(0) # [1,H,W,3]
|
||||
mask_out = alpha # [H,W]
|
||||
depth4 = depth_img.unsqueeze(0).unsqueeze(-1) # [1,H,W,1]
|
||||
if return_inverse_depth:
|
||||
depth = 1.0 / depth.clamp(min=1e-6)
|
||||
depth *= mask.unsqueeze(0).unsqueeze(-1)
|
||||
return img, mask, depth
|
||||
|
||||
|
||||
|
||||
depth4 = 1.0 / depth4.clamp(min=1e-6)
|
||||
depth4 = depth4 * mask_out.unsqueeze(0).unsqueeze(-1)
|
||||
return img, mask_out, depth4
|
||||
|
||||
class PointCloudUnion:
|
||||
"""
|
||||
@@ -614,7 +456,7 @@ class PointCloudUnion:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES =("merged pointcloud",)
|
||||
FUNCTION = "union_pointclouds"
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def union_pointclouds(
|
||||
self,
|
||||
@@ -650,7 +492,7 @@ class LoadPointCloud:
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("loaded pointcloud",)
|
||||
FUNCTION = "load_pointcloud"
|
||||
@@ -662,45 +504,24 @@ class LoadPointCloud:
|
||||
arr = np.load(file_path)
|
||||
tensor_pc = torch.from_numpy(arr)
|
||||
return (tensor_pc,)
|
||||
|
||||
if o3d is None:
|
||||
logging.warning("[camera-comfyUI] open3d is not installed. Falling back to manual PLY parser.")
|
||||
coords = []
|
||||
colors = []
|
||||
with open(file_path, 'r') as f:
|
||||
coords = []
|
||||
colors = []
|
||||
with open(file_path, 'r') as f:
|
||||
line = f.readline().strip()
|
||||
while not line.startswith("end_header"):
|
||||
line = f.readline().strip()
|
||||
while not line.startswith("end_header"):
|
||||
line = f.readline().strip()
|
||||
for line in f:
|
||||
parts = line.strip().split()
|
||||
if len(parts) < 7:
|
||||
continue
|
||||
x, y, z = map(float, parts[0:3])
|
||||
r, g, b, a = map(float, parts[3:7])
|
||||
coords.append((x, y, z))
|
||||
colors.append((r, g, b, a))
|
||||
np_coords = np.array(coords, dtype=np.float32)
|
||||
np_colors = np.array(colors, dtype=np.float32)
|
||||
# if colors are > 1, normalize them to [0,1]
|
||||
if np_colors.max() > 1.0:
|
||||
np_colors = np_colors / 255.0
|
||||
else:
|
||||
pc = o3d.t.io.read_point_cloud(file_path)
|
||||
np_coords = pc.point["positions"].numpy().astype(np.float32)
|
||||
if "colors" in pc.point:
|
||||
cols = pc.point["colors"].numpy().astype(np.float32)
|
||||
else:
|
||||
cols = np.ones((np_coords.shape[0], 3), dtype=np.float32)
|
||||
if "alpha" in pc.point:
|
||||
alpha = pc.point["alpha"].numpy().astype(np.float32)
|
||||
else:
|
||||
alpha = np.ones((np_coords.shape[0], 1), dtype=np.float32)
|
||||
np_colors = np.concatenate([cols, alpha], axis=1)
|
||||
if np_colors.max() > 1.0:
|
||||
np_colors = np_colors / 255.0
|
||||
# combine coords and colors into a single tensor
|
||||
combined = np.concatenate([np_coords, np_colors], axis=1)
|
||||
tensor_pc = torch.from_numpy(combined)
|
||||
for line in f:
|
||||
parts = line.strip().split()
|
||||
if len(parts) < 7:
|
||||
continue
|
||||
x, y, z = map(float, parts[0:3])
|
||||
r, g, b, a = map(int, parts[3:7])
|
||||
coords.append((x, y, z))
|
||||
colors.append((r, g, b, a))
|
||||
np_coords = np.array(coords, dtype=np.float32)
|
||||
np_colors = np.array(colors, dtype=np.float32)/255.0
|
||||
combined = np.concatenate([np_coords, np_colors], axis=1)
|
||||
tensor_pc = torch.from_numpy(combined)
|
||||
return (tensor_pc,)
|
||||
|
||||
@classmethod
|
||||
@@ -748,7 +569,7 @@ class SavePointCloud:
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "save_pointcloud"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
DESCRIPTION = "Saves the input point cloud to your ComfyUI output directory as .ply or .npy."
|
||||
|
||||
def save_pointcloud(self, pointcloud: torch.Tensor, filename_prefix: str, save_as: str = "ply"):
|
||||
@@ -765,36 +586,24 @@ class SavePointCloud:
|
||||
os.makedirs(full_output_folder, exist_ok=True)
|
||||
base_name = filename.replace("%batch_num%", "0")
|
||||
if save_as == "ply":
|
||||
ply_name = f"{base_name}_{counter:05}.ply"
|
||||
ply_path = os.path.join(full_output_folder, ply_name)
|
||||
coords = pointcloud[:, :3].cpu().numpy().astype(np.float32)
|
||||
colors = pointcloud[:, 3:].cpu().numpy().clip(0, 1).astype(np.float32)
|
||||
|
||||
if o3d is None:
|
||||
logging.warning("[camera-comfyUI] open3d is not installed. Falling back to manual ASCII PLY writer.")
|
||||
with open(ply_path, 'w') as f:
|
||||
f.write("ply\n")
|
||||
f.write("format ascii 1.0\n")
|
||||
f.write(f"element vertex {coords.shape[0]}\n")
|
||||
f.write("property float x\n")
|
||||
f.write("property float y\n")
|
||||
f.write("property float z\n")
|
||||
f.write("property float red\n")
|
||||
f.write("property float green\n")
|
||||
f.write("property float blue\n")
|
||||
f.write("property float alpha\n")
|
||||
f.write("end_header\n")
|
||||
for (x, y, z), (r, g, b, a) in zip(coords, colors):
|
||||
f.write(f"{x} {y} {z} {r} {g} {b} {a}\n")
|
||||
else:
|
||||
pc = o3d.t.geometry.PointCloud()
|
||||
pc.point["positions"] = o3d.core.Tensor(coords, o3d.core.float32)
|
||||
pc.point["colors"] = o3d.core.Tensor(colors[:, :3], o3d.core.float32)
|
||||
if colors.shape[1] > 3:
|
||||
pc.point["alpha"] = o3d.core.Tensor(colors[:, 3:], o3d.core.float32)
|
||||
else:
|
||||
pc.point["alpha"] = o3d.core.Tensor(np.ones((coords.shape[0], 1), dtype=np.float32), o3d.core.float32)
|
||||
o3d.t.io.write_point_cloud(ply_path, pc)
|
||||
ply_name = f"{base_name}_{counter:05}.ply"
|
||||
ply_path = os.path.join(full_output_folder, ply_name)
|
||||
coords = pointcloud[:, :3].cpu().numpy()
|
||||
colors = pointcloud[:, 3:].cpu().numpy().clip(0,1)
|
||||
with open(ply_path, 'w') as f:
|
||||
f.write("ply\n")
|
||||
f.write("format ascii 1.0\n")
|
||||
f.write(f"element vertex {coords.shape[0]}\n")
|
||||
f.write("property float x\n")
|
||||
f.write("property float y\n")
|
||||
f.write("property float z\n")
|
||||
f.write("property uchar red\n")
|
||||
f.write("property uchar green\n")
|
||||
f.write("property uchar blue\n")
|
||||
f.write("property uchar alpha\n")
|
||||
f.write("end_header\n")
|
||||
for (x,y,z), (r,g,b,a) in zip(coords, colors):
|
||||
f.write(f"{x} {y} {z} {int(r*255)} {int(g*255)} {int(b*255)} {int(a*255)}\n")
|
||||
file_name = ply_name
|
||||
else:
|
||||
npy_name = f"{base_name}_{counter:05}.npy"
|
||||
@@ -831,15 +640,12 @@ class CameraMotionNode:
|
||||
"output_width": ("INT", {"default":512, "min":8, "max":16384}),
|
||||
"output_height": ("INT", {"default":512, "min":8, "max":16384}),
|
||||
"point_size": ("INT", {"default":1, "min":1}),
|
||||
"widen_mask": ("INT", {"default":0, "min":0, "max":64}),
|
||||
"invert_mask": ("BOOLEAN", {"default": False}),
|
||||
"points_to_mask": ("BOOLEAN", {"default": False, "tooltip": "Output mask frames of projected points"}),
|
||||
}}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("motion_frames", "mask_frames")
|
||||
RETURN_TYPES = ("IMAGE",)
|
||||
RETURN_NAMES = ("motion_frames",)
|
||||
FUNCTION = "generate_motion_frames"
|
||||
CATEGORY = "Camera/Trajectory"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def generate_motion_frames(
|
||||
self,
|
||||
@@ -850,10 +656,7 @@ class CameraMotionNode:
|
||||
output_horizontal_fov: float,
|
||||
output_width: int,
|
||||
output_height: int,
|
||||
point_size: int = 1,
|
||||
widen_mask: int = 0,
|
||||
invert_mask: bool = False,
|
||||
points_to_mask: bool = False
|
||||
point_size: int = 1
|
||||
) -> Tuple[torch.Tensor]:
|
||||
# validate trajectory shape
|
||||
if trajectory.dim() != 3 or trajectory.shape[1:] != (4,4):
|
||||
@@ -880,10 +683,9 @@ class CameraMotionNode:
|
||||
proj_node = ProjectPointCloud()
|
||||
transform_node = TransformPointCloud()
|
||||
frames = []
|
||||
masks = []
|
||||
for M in tqdm(full_traj):
|
||||
pc_t, = transform_node.transform_pointcloud(pointcloud, M)
|
||||
img, mask, _ = proj_node.project_pointcloud(
|
||||
img, _, _ = proj_node.project_pointcloud(
|
||||
pc_t,
|
||||
output_projection,
|
||||
output_horizontal_fov,
|
||||
@@ -891,25 +693,15 @@ class CameraMotionNode:
|
||||
output_height,
|
||||
point_size
|
||||
)
|
||||
if widen_mask > 0:
|
||||
k = 2 * widen_mask + 1
|
||||
pad = widen_mask
|
||||
mask = F.max_pool2d(mask.float().unsqueeze(0).unsqueeze(0), kernel_size=k, stride=1, padding=pad).squeeze(0).squeeze(0)
|
||||
if invert_mask:
|
||||
mask = 1.0 - mask
|
||||
masks.append(mask)
|
||||
if points_to_mask:
|
||||
img = mask.unsqueeze(-1).repeat(1,1,1,3)
|
||||
frames.append(img[0])
|
||||
|
||||
# output as (T,H,W,3)
|
||||
return (torch.stack(frames, dim=0), torch.stack(masks, dim=0))
|
||||
return (torch.stack(frames, dim=0),)
|
||||
|
||||
class CameraInterpolationNode:
|
||||
"""
|
||||
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).
|
||||
Wrap two 4×4 poses into a trajectory tensor.
|
||||
Outputs only `trajectory` (shape 2×4×4).
|
||||
"""
|
||||
|
||||
@classmethod
|
||||
@@ -918,29 +710,25 @@ 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",)
|
||||
FUNCTION = "interpolate"
|
||||
CATEGORY = "Camera/Trajectory"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def interpolate(
|
||||
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()
|
||||
keyframes = torch.stack([initial_matrix.float(), final_matrix.float()], dim=0)
|
||||
traj = interpolate_se3(keyframes, num_steps)
|
||||
traj = torch.stack([initial_matrix, final_matrix], dim=0)
|
||||
return (traj,)
|
||||
|
||||
|
||||
@@ -955,14 +743,14 @@ class CameraTrajectoryNode:
|
||||
"pointcloud": ("TENSOR",),
|
||||
},
|
||||
"optional": {
|
||||
"initial_matrix": ("MAT_4X4",),
|
||||
"initial_matrix": ("MAT_4X4"),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("trajectory",)
|
||||
FUNCTION = "build_trajectory"
|
||||
CATEGORY = "Camera/Trajectory"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def build_trajectory(
|
||||
self,
|
||||
@@ -1102,7 +890,7 @@ class PointCloudCleaner:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("cleaned_pointcloud",)
|
||||
FUNCTION = "clean_pointcloud"
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def clean_pointcloud(
|
||||
self,
|
||||
@@ -1178,7 +966,7 @@ class ProjectAndClean:
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("cleaned_pointcloud",)
|
||||
FUNCTION = "project_and_clean"
|
||||
CATEGORY = "Camera/PointCloud"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
|
||||
def project_and_clean(
|
||||
self,
|
||||
@@ -1292,7 +1080,7 @@ class SaveTrajectory:
|
||||
RETURN_TYPES = ()
|
||||
FUNCTION = "save_trajectory"
|
||||
OUTPUT_NODE = True
|
||||
CATEGORY = "Camera/Trajectory"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
DESCRIPTION = "Saves the input trajectory tensor (N,4,4) to your ComfyUI output directory as .npy."
|
||||
|
||||
def save_trajectory(self, trajectory: torch.Tensor, filename_prefix: str):
|
||||
@@ -1342,7 +1130,7 @@ class LoadTrajectory:
|
||||
}
|
||||
}
|
||||
|
||||
CATEGORY = "Camera/Trajectory"
|
||||
CATEGORY = "Camera/pointcloud"
|
||||
RETURN_TYPES = ("TENSOR",)
|
||||
RETURN_NAMES = ("loaded_trajectory",)
|
||||
FUNCTION = "load_trajectory"
|
||||
@@ -1368,95 +1156,6 @@ 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,
|
||||
@@ -1471,5 +1170,4 @@ NODE_CLASS_MAPPINGS = {
|
||||
"PointCloudCleaner": PointCloudCleaner,
|
||||
"SaveTrajectory": SaveTrajectory,
|
||||
"LoadTrajectory": LoadTrajectory,
|
||||
"DepthEdgeFilter": DepthEdgeFilter,
|
||||
}
|
||||
-461
@@ -1,461 +0,0 @@
|
||||
"""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,
|
||||
}
|
||||
@@ -1,23 +0,0 @@
|
||||
[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/"]
|
||||
+6
-13
@@ -42,14 +42,7 @@ def map_grid(
|
||||
output_horizontal_fov = torch.tensor(output_horizontal_fov, device=grid_torch.device).float()
|
||||
|
||||
# Calculate vertical field of view for input and output projections
|
||||
# For equirectangular, use 2:1 aspect ratio (vertical FOV = horizontal FOV / 2)
|
||||
if output_projection == "EQUIRECTANGULAR":
|
||||
output_vertical_fov = output_horizontal_fov / 2.0
|
||||
else:
|
||||
output_vertical_fov = output_horizontal_fov # Assuming square aspect ratio for other projections
|
||||
|
||||
# Calculate input vertical FOV based on output grid aspect ratio
|
||||
# This allows the input's vertical range to adapt to the output dimensions
|
||||
output_vertical_fov = output_horizontal_fov # Assuming square aspect ratio
|
||||
input_vertical_fov = input_horizontal_fov * (grid_torch.shape[0] / grid_torch.shape[1])
|
||||
|
||||
# Normalize the grid for vertical FOV adjustment
|
||||
@@ -155,7 +148,7 @@ class ReprojectImage:
|
||||
RETURN_TYPES: Tuple[str, str] = ("IMAGE", "MASK")
|
||||
RETURN_NAMES = ("reprojected image", "reprojected mask")
|
||||
FUNCTION: str = "reproject_image"
|
||||
CATEGORY: str = "Camera/Reprojection"
|
||||
CATEGORY: str = "Camera/reproject"
|
||||
|
||||
def reproject_image(
|
||||
self,
|
||||
@@ -229,8 +222,8 @@ class ReprojectImage:
|
||||
)
|
||||
|
||||
grid_y, grid_x = torch.meshgrid(
|
||||
torch.linspace(-1, 1, output_height, device=image_tensor.device),
|
||||
torch.linspace(-1, 1, output_width, device=image_tensor.device),
|
||||
torch.linspace(-1, 1, output_height, device=image_tensor.device),
|
||||
indexing="ij"
|
||||
)
|
||||
grid_init = torch.stack((grid_x, grid_y), dim=-1)
|
||||
@@ -310,7 +303,7 @@ class TransformToMatrix:
|
||||
RETURN_TYPES: Tuple[str] = ("MAT_4X4",)
|
||||
RETURN_NAMES = ("transformation matrix",)
|
||||
FUNCTION: str = "generate_matrix"
|
||||
CATEGORY: str = "Camera/Matrix"
|
||||
CATEGORY: str = "Camera/reproject"
|
||||
|
||||
def generate_matrix(
|
||||
self,
|
||||
@@ -399,7 +392,7 @@ class TransformToMatrixManual:
|
||||
RETURN_TYPES: Tuple[str] = ("MAT_4X4",)
|
||||
RETURN_NAMES = ("transformation matrix",)
|
||||
FUNCTION: str = "generate_matrix"
|
||||
CATEGORY: str = "Camera/Matrix"
|
||||
CATEGORY: str = "Camera/reproject"
|
||||
|
||||
def generate_matrix(
|
||||
self,
|
||||
@@ -451,7 +444,7 @@ class ReprojectDepth:
|
||||
RETURN_TYPES: Tuple[str, str] = ("TENSOR", "MASK")
|
||||
RETURN_NAMES = ("reprojected_depth", "reprojected_mask")
|
||||
FUNCTION: str = "reproject_depth"
|
||||
CATEGORY: str = "Camera/Reprojection"
|
||||
CATEGORY: str = "Camera/reproject"
|
||||
|
||||
def reproject_depth(
|
||||
self,
|
||||
|
||||
Submodule submodules/ml-sharpt deleted from 1eaa046834
-299
@@ -1,299 +0,0 @@
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
import os
|
||||
import sys
|
||||
from typing import Dict, Any, Tuple
|
||||
from tqdm import tqdm # Added tqdm import
|
||||
|
||||
# Import existing pointcloud nodes and projection definitions
|
||||
from .pointcloud_nodes import DepthToPointCloud, TransformPointCloud, ProjectPointCloud, Projection, PointCloudCleaner, interpolate_se3
|
||||
import folder_paths
|
||||
|
||||
# Ensure video_depth_anything is on path
|
||||
_here = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
# climb up 3 levels: camera-comfyUI → custom_nodes → ComfyUI
|
||||
COMFYUI_ROOT = os.path.abspath(os.path.join(_here, os.pardir, os.pardir))
|
||||
|
||||
# point at metric_depth inside the Video-Depth-Anything clone at the ComfyUI root
|
||||
video_depth_path = os.path.join(COMFYUI_ROOT, "Video-Depth-Anything", "metric_depth")
|
||||
|
||||
# insert at front so it always wins
|
||||
if video_depth_path not in sys.path:
|
||||
sys.path.insert(0, video_depth_path)
|
||||
NO_VIDEO_DEPTH_ANYTHING= False
|
||||
try:
|
||||
from video_depth_anything.video_depth import VideoDepthAnything
|
||||
print("✅ video_depth_anything module loaded successfully.")
|
||||
except ImportError as e:
|
||||
NO_VIDEO_DEPTH_ANYTHING = True
|
||||
print(
|
||||
f"❌ Could not load video_depth_anything from {video_depth_path!r}: {e}"
|
||||
)
|
||||
|
||||
class VideoCameraMotionSequence:
|
||||
"""
|
||||
Takes a sequence of RGB frames and corresponding depth maps,
|
||||
converts each frame+depth to a pointcloud, interpolates a camera
|
||||
trajectory to match video length, cleans the pointcloud if needed,
|
||||
and outputs reprojected images, masks, and depth maps per frame.
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
# Sequence of frames: Tensor [T, H, W, 3]
|
||||
"frames": ("IMAGE", {"shape_hint": [None, None, None, 3]}),
|
||||
# Sequence of depth maps: Tensor [T, H, W] or [T, H, W, 1]
|
||||
"depth_seq": ("TENSOR", {"shape_hint": [None, None, None]}),
|
||||
# Camera trajectory waypoints: Tensor [K, 4, 4]
|
||||
"trajectory": ("TENSOR", {"shape_hint": [None, 4, 4]}),
|
||||
# Optional mask sequence: Tensor [T, H, W] or [T, H, W, 1]
|
||||
"mask_seq": ("MASK", {"shape_hint": [None, None, None], "optional": True}),
|
||||
# Input projection parameters
|
||||
"input_projection": (Projection.PROJECTIONS, {}),
|
||||
"input_horizontal_fov": ("FLOAT", {"default": 90.0}),
|
||||
"depth_scale": ("FLOAT", {"default": 1.0}),
|
||||
"invert_depth": ("BOOLEAN", {"default": False}),
|
||||
# Output projection parameters
|
||||
"output_projection": (Projection.PROJECTIONS, {}),
|
||||
"output_horizontal_fov": ("FLOAT", {"default": 90.0}),
|
||||
"output_width": ("INT", {"default": 512, "min": 1}),
|
||||
"output_height": ("INT", {"default": 512, "min": 1}),
|
||||
"point_size": ("INT", {"default": 1, "min": 1}),
|
||||
# Cleaning parameters
|
||||
"voxel_size": ("FLOAT", {"default": 1.0, "min": 1e-3}),
|
||||
"min_points_per_voxel": ("INT", {"default": 3, "min": 1}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("IMAGE", "MASK", "TENSOR")
|
||||
RETURN_NAMES = ("video_frames", "mask_frames", "depths")
|
||||
FUNCTION = "process_sequence"
|
||||
CATEGORY = "Camera/Video"
|
||||
|
||||
def process_sequence(
|
||||
self,
|
||||
frames: torch.Tensor,
|
||||
depth_seq: torch.Tensor,
|
||||
trajectory: torch.Tensor,
|
||||
input_projection: str,
|
||||
input_horizontal_fov: float,
|
||||
depth_scale: float,
|
||||
invert_depth: bool,
|
||||
output_projection: str,
|
||||
output_horizontal_fov: float,
|
||||
output_width: int,
|
||||
output_height: int,
|
||||
point_size: int,
|
||||
voxel_size: float,
|
||||
min_points_per_voxel: int,
|
||||
mask_seq: torch.Tensor = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
# frames: [T, H, W, 3]
|
||||
# depth_seq: [T, H, W] or [T, H, W, 1]
|
||||
T, H, W, _ = frames.shape
|
||||
|
||||
# Interpolate trajectory to match T (SE(3): quaternion SLERP on R, lerp on t)
|
||||
interp_traj = interpolate_se3(trajectory, T)
|
||||
|
||||
out_frames = []
|
||||
out_masks = []
|
||||
out_depths = []
|
||||
|
||||
# Add tqdm progress bar for the sequence
|
||||
# If mask_seq is a single mask [H, W] or [H, W, 1], repeat it for all frames
|
||||
if mask_seq is not None:
|
||||
if mask_seq.dim() == 2 or (mask_seq.dim() == 3 and mask_seq.shape[0] == 1):
|
||||
mask_seq = mask_seq.unsqueeze(0) if mask_seq.dim() == 2 else mask_seq
|
||||
mask_seq = mask_seq.repeat(T, 1, 1, 1) if mask_seq.dim() == 4 else mask_seq.repeat(T, 1, 1)
|
||||
|
||||
for i, (frame, depth, pose) in enumerate(tqdm(zip(frames, depth_seq, interp_traj), total=T, desc="Processing video frames")):
|
||||
if depth.dim() == 3 and depth.shape[-1] == 1:
|
||||
depth = depth.squeeze(-1)
|
||||
# Use mask if provided; must be (re)initialized every iteration
|
||||
mask = None
|
||||
if mask_seq is not None:
|
||||
mask = mask_seq[i]
|
||||
if mask.dim() == 3 and mask.shape[-1] == 1:
|
||||
mask = mask.squeeze(-1)
|
||||
# to pointcloud
|
||||
pc, = DepthToPointCloud().depth_to_pointcloud(
|
||||
image=frame.permute(2, 0, 1),
|
||||
input_projection=input_projection,
|
||||
input_horizontal_fov=input_horizontal_fov,
|
||||
depth_scale=depth_scale,
|
||||
invert_depth=invert_depth,
|
||||
depthmap=depth,
|
||||
mask=mask,
|
||||
)
|
||||
# optional cleaning
|
||||
if min_points_per_voxel > 1:
|
||||
pc, = PointCloudCleaner().clean_pointcloud(
|
||||
pointcloud=pc,
|
||||
width=output_width,
|
||||
height=output_height,
|
||||
voxel_size=voxel_size,
|
||||
min_points_per_voxel=min_points_per_voxel,
|
||||
)
|
||||
# transform and project
|
||||
pc_t, = TransformPointCloud().transform_pointcloud(pc, pose)
|
||||
img_t, mask_t, depth_t = ProjectPointCloud().project_pointcloud(
|
||||
pointcloud=pc_t,
|
||||
output_projection=output_projection,
|
||||
output_horizontal_fov=output_horizontal_fov,
|
||||
output_width=output_width,
|
||||
output_height=output_height,
|
||||
point_size=point_size,
|
||||
)
|
||||
|
||||
out_frames.append(img_t[0])
|
||||
out_masks.append(mask_t)
|
||||
out_depths.append(depth_t)
|
||||
|
||||
return (
|
||||
torch.stack(out_frames, dim=0), # [T, 3, H, W]
|
||||
torch.stack(out_masks, dim=0), # [T, H, W]
|
||||
torch.stack(out_depths, dim=0), # [T, H, W]
|
||||
)
|
||||
|
||||
|
||||
class DepthFramesToVideo:
|
||||
"""
|
||||
Converts a sequence of depth maps into video frame tensors for saving.
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
return {
|
||||
"required": {
|
||||
"depth_seq": ("TENSOR", {"shape_hint": [None, None, None]}),
|
||||
"mask_seq": ("MASK", {"shape_hint": [None, None, None]}),
|
||||
"normalize": ("BOOLEAN", {"default": True}),
|
||||
"invert_depth": ("BOOLEAN", {"default": False}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("TENSOR", "IMAGE")
|
||||
RETURN_NAMES = ("video_frames", "depth_video")
|
||||
FUNCTION = "depth_to_video_frames"
|
||||
CATEGORY = "Camera/Video"
|
||||
|
||||
def depth_to_video_frames(
|
||||
self,
|
||||
depth_seq: torch.Tensor,
|
||||
normalize: bool,
|
||||
invert_depth: bool,
|
||||
mask_seq: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
ds = depth_seq.clone().squeeze()
|
||||
if ds.dim() == 2:
|
||||
ds = ds.unsqueeze(0) # [H, W] -> [1, H, W]
|
||||
if ds.dim() != 3:
|
||||
raise ValueError(f"Expected ds to be 3D [T, H, W], got shape {ds.shape}")
|
||||
if invert_depth:
|
||||
ds= 1.0 / (ds + 1e-8) # Avoid division by zero
|
||||
if normalize:
|
||||
# Mask: only normalize where depth > 0
|
||||
mask = mask_seq>0.5
|
||||
if mask.any():
|
||||
#percentile first 10 percent min
|
||||
# sample
|
||||
minv = ds[mask]
|
||||
# sample 10000 and find 10% quantile
|
||||
if minv.numel() > 10000:
|
||||
minv = minv[torch.randperm(minv.numel())[:10000]]
|
||||
|
||||
minv = minv.quantile(0.2)
|
||||
minv = minv if minv > 0.1 else 0.1 # Avoid division by zero
|
||||
#percentile last 10 percent max
|
||||
maxv = ds[mask]
|
||||
if maxv.numel() > 10000:
|
||||
maxv = maxv[torch.randperm(maxv.numel())[:10000]]
|
||||
maxv = maxv.quantile(0.98)
|
||||
maxv = maxv if maxv < 100 else 100
|
||||
print(f"Normalizing depth: min={minv}, max={maxv}")
|
||||
ds_norm = (ds - minv) / (maxv - minv + 1e-8)
|
||||
ds = ds_norm.clamp(0, 1) # torch.where(mask, ds_norm, ds) # Only normalize valid values
|
||||
else:
|
||||
print("Warning: No valid depth values for normalization.")
|
||||
# expand to 3 channels: [T, H, W] -> [T, 3, H, W]
|
||||
raw = depth_seq.clone().squeeze()
|
||||
ds_u8 = (ds * 255.0).round().to(torch.uint8)
|
||||
raw_u8 = (raw.clamp(0, 255)).to(torch.uint8) # if raw is already in a displayable range
|
||||
|
||||
# expand to 3 channels and permute to HWC
|
||||
ds_color = ds_u8.unsqueeze(1).repeat(1, 3, 1, 1).permute(0, 2, 3, 1)
|
||||
raw_color = raw_u8.unsqueeze(1).repeat(1, 3, 1, 1).permute(0, 2, 3, 1)
|
||||
return raw_color, ds_color # [T, 3, H, W] -> [T, H, W, 3]
|
||||
|
||||
# Cache for loaded VideoDepthAnything models, keyed by (checkpoint, device)
|
||||
_VIDEO_DEPTH_MODEL_CACHE: Dict[Tuple[str, str], Any] = {}
|
||||
|
||||
|
||||
class VideoMetricDepthEstimate:
|
||||
"""
|
||||
Estimates metric depth for a sequence of frames using VideoDepthAnything.
|
||||
"""
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls) -> Dict[str, Any]:
|
||||
# model files (.pth) in input directory
|
||||
model_dir = os.path.join(os.getcwd(), "models", "checkpoints")
|
||||
os.makedirs(model_dir, exist_ok=True)
|
||||
files = [f for f in os.listdir(model_dir) if f.lower().endswith(('.pth', '.ckpt', '.safetensors'))]
|
||||
return {
|
||||
"required": {
|
||||
"frames": ("IMAGE", {"shape_hint": [None, None, None, 3]}),
|
||||
"model_checkpoint": (files, {"file_chooser": True}),
|
||||
"input_size": ("INT", {"default": 518, "min": 64, "max": 2048}),
|
||||
"max_fps": ("INT", {"default": 60, "min": 1}),
|
||||
}
|
||||
}
|
||||
RETURN_TYPES = ("TENSOR", "FLOAT")
|
||||
RETURN_NAMES = ("metric_depths", "fps")
|
||||
FUNCTION = "estimate_metric_depth"
|
||||
CATEGORY = "Camera/Video"
|
||||
|
||||
def estimate_metric_depth(
|
||||
self,
|
||||
frames: torch.Tensor,
|
||||
model_checkpoint: str,
|
||||
input_size: int,
|
||||
max_fps: int,
|
||||
) -> Tuple[torch.Tensor, float]:
|
||||
if NO_VIDEO_DEPTH_ANYTHING:
|
||||
raise ImportError(
|
||||
f"VideoDepthAnything library not found. Clone "
|
||||
f"https://github.com/DepthAnything/Video-Depth-Anything into {COMFYUI_ROOT!r} "
|
||||
f"(expected module path: {video_depth_path!r})."
|
||||
)
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
# if max input<1.5 normalize to 0-255
|
||||
if frames.max() < 1.5:
|
||||
frames = (frames * 255)
|
||||
cache_key = (model_checkpoint, str(device))
|
||||
model = _VIDEO_DEPTH_MODEL_CACHE.get(cache_key)
|
||||
if model is None:
|
||||
# Same checkpoint directory as computed in INPUT_TYPES
|
||||
model_dir = os.path.join(os.getcwd(), "models", "checkpoints")
|
||||
checkpoint_path = os.path.join(model_dir, model_checkpoint)
|
||||
if not os.path.isfile(checkpoint_path):
|
||||
raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}")
|
||||
model = VideoDepthAnything(**{"encoder": "vitl", "features": 256, "out_channels": [256,512,1024,1024]})
|
||||
state = torch.load(checkpoint_path, map_location='cpu')
|
||||
model.load_state_dict(state, strict=True)
|
||||
model = model.to(device).eval()
|
||||
_VIDEO_DEPTH_MODEL_CACHE[cache_key] = model
|
||||
np_frames = frames.cpu().numpy().astype(np.uint8)
|
||||
metric_depths, fps = model.infer_video_depth(np_frames, max_fps, input_size=input_size, device=device.type, fp32=False)
|
||||
return (torch.from_numpy(metric_depths), float(fps))
|
||||
|
||||
# Register nodes
|
||||
if NO_VIDEO_DEPTH_ANYTHING:
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
else:
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VideoCameraMotionSequence": VideoCameraMotionSequence,
|
||||
"VideoMetricDepthEstimate": VideoMetricDepthEstimate,
|
||||
"DepthFramesToVideo": DepthFramesToVideo,
|
||||
}
|
||||
@@ -1,289 +1 @@
|
||||
{
|
||||
"id": "dd56c0bf-7405-406e-924f-42b2feacb73f",
|
||||
"revision": 0,
|
||||
"last_node_id": 8,
|
||||
"last_link_id": 9,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 3,
|
||||
"type": "SaveWEBM",
|
||||
"pos": [
|
||||
606.5880737304688,
|
||||
1515.5966796875
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
437
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 5
|
||||
}
|
||||
],
|
||||
"outputs": [],
|
||||
"properties": {},
|
||||
"widgets_values": [
|
||||
"ComfyUI",
|
||||
"vp9",
|
||||
10.000000000000002,
|
||||
32
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 7,
|
||||
"type": "CameraMotionNode",
|
||||
"pos": [
|
||||
176.07933044433594,
|
||||
1504.22705078125
|
||||
],
|
||||
"size": [
|
||||
278.75,
|
||||
270
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "pointcloud",
|
||||
"type": "TENSOR",
|
||||
"link": 6
|
||||
},
|
||||
{
|
||||
"name": "trajectory",
|
||||
"type": "TENSOR",
|
||||
"link": 9
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "motion_frames",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
5
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "mask_frames",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "CameraMotionNode"
|
||||
},
|
||||
"widgets_values": [
|
||||
10,
|
||||
"PINHOLE",
|
||||
90,
|
||||
512,
|
||||
512,
|
||||
1,
|
||||
0,
|
||||
false,
|
||||
false
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "LoadPointCloud",
|
||||
"pos": [
|
||||
-357.24761962890625,
|
||||
1294.427734375
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
58
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "loaded pointcloud",
|
||||
"type": "TENSOR",
|
||||
"links": [
|
||||
6
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadPointCloud"
|
||||
},
|
||||
"widgets_values": [
|
||||
"ComfyUIPointCloud_00001.ply"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "TransformToMatrix",
|
||||
"pos": [
|
||||
-537.2319946289062,
|
||||
1491.40380859375
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
154
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "transformation matrix",
|
||||
"type": "MAT_4X4",
|
||||
"links": [
|
||||
7
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TransformToMatrix"
|
||||
},
|
||||
"widgets_values": [
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "TransformToMatrix",
|
||||
"pos": [
|
||||
-531.216796875,
|
||||
1701.1632080078125
|
||||
],
|
||||
"size": [
|
||||
315,
|
||||
154
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "transformation matrix",
|
||||
"type": "MAT_4X4",
|
||||
"links": [
|
||||
8
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TransformToMatrix"
|
||||
},
|
||||
"widgets_values": [
|
||||
0.10000000000000002,
|
||||
0,
|
||||
0,
|
||||
0,
|
||||
0
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "CameraInterpolationNode",
|
||||
"pos": [
|
||||
-121.91971588134766,
|
||||
1597.3514404296875
|
||||
],
|
||||
"size": [
|
||||
200.21640014648438,
|
||||
46
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "initial_matrix",
|
||||
"type": "MAT_4X4",
|
||||
"link": 7
|
||||
},
|
||||
{
|
||||
"name": "final_matrix",
|
||||
"type": "MAT_4X4",
|
||||
"link": 8
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "trajectory",
|
||||
"type": "TENSOR",
|
||||
"links": [
|
||||
9
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "CameraInterpolationNode"
|
||||
},
|
||||
"widgets_values": []
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
5,
|
||||
7,
|
||||
0,
|
||||
3,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
6,
|
||||
6,
|
||||
0,
|
||||
7,
|
||||
0,
|
||||
"TENSOR"
|
||||
],
|
||||
[
|
||||
7,
|
||||
1,
|
||||
0,
|
||||
8,
|
||||
0,
|
||||
"MAT_4X4"
|
||||
],
|
||||
[
|
||||
8,
|
||||
5,
|
||||
0,
|
||||
8,
|
||||
1,
|
||||
"MAT_4X4"
|
||||
],
|
||||
[
|
||||
9,
|
||||
8,
|
||||
0,
|
||||
7,
|
||||
1,
|
||||
"TENSOR"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 1.015255979947716,
|
||||
"offset": [
|
||||
636.7531305750655,
|
||||
-1186.3099424359816
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.21.7"
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
{"id":"dd56c0bf-7405-406e-924f-42b2feacb73f","revision":0,"last_node_id":6,"last_link_id":4,"nodes":[{"id":1,"type":"TransformToMatrix","pos":[-337.9580993652344,1500.5211181640625],"size":[315,154],"flags":{},"order":0,"mode":0,"inputs":[],"outputs":[{"localized_name":"MAT_4X4","name":"MAT_4X4","type":"MAT_4X4","links":[1]}],"properties":{"Node name for S&R":"TransformToMatrix"},"widgets_values":[0,0,0,0,0]},{"id":2,"type":"CameraMotion","pos":[130.0218505859375,1516.7567138671875],"size":[367.79998779296875,218],"flags":{},"order":3,"mode":0,"inputs":[{"localized_name":"pointcloud","name":"pointcloud","type":"TENSOR","link":4},{"localized_name":"initial_matrix","name":"initial_matrix","type":"MAT_4X4","link":1},{"localized_name":"final_matrix","name":"final_matrix","type":"MAT_4X4","link":2}],"outputs":[{"localized_name":"IMAGE","name":"IMAGE","type":"IMAGE","links":[3]}],"properties":{"Node name for S&R":"CameraMotion"},"widgets_values":[24,"PINHOLE",90,1024,1024,2]},{"id":3,"type":"SaveWEBM","pos":[606.5880737304688,1515.5966796875],"size":[315,437],"flags":{},"order":4,"mode":0,"inputs":[{"localized_name":"images","name":"images","type":"IMAGE","link":3}],"outputs":[],"properties":{},"widgets_values":["ComfyUI","vp9",10.000000000000002,32]},{"id":5,"type":"TransformToMatrix","pos":[-262.9134521484375,1740.8876953125],"size":[315,154],"flags":{},"order":1,"mode":0,"inputs":[],"outputs":[{"localized_name":"MAT_4X4","name":"MAT_4X4","type":"MAT_4X4","links":[2]}],"properties":{"Node name for S&R":"TransformToMatrix"},"widgets_values":[0.10000000000000002,0,0,0,0]},{"id":6,"type":"LoadPointCloud","pos":[-357.24761962890625,1294.427734375],"size":[315,58],"flags":{},"order":2,"mode":0,"inputs":[],"outputs":[{"localized_name":"TENSOR","name":"TENSOR","type":"TENSOR","links":[4]}],"properties":{"Node name for S&R":"LoadPointCloud"},"widgets_values":["ComfyUIPointCloud_00001.ply"]}],"links":[[1,1,0,2,1,"MAT_4X4"],[2,5,0,2,2,"MAT_4X4"],[3,2,0,3,0,"IMAGE"],[4,6,0,2,0,"TENSOR"]],"groups":[],"config":{},"extra":{"ds":{"scale":1.351305709310409,"offset":[-187.6257577580669,-1567.2143321744395]}},"version":0.4}
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
-632
@@ -1,632 +0,0 @@
|
||||
"""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