From ca7e351a70e506f69fd075ac8a08e8554c835dac Mon Sep 17 00:00:00 2001 From: IDGallagher Date: Thu, 8 May 2025 23:35:48 +0100 Subject: [PATCH] Stiching actually working --- __init__.py | 9 +- nodes/auto_stitch_rgb_tiles.py | 135 --------------- nodes/auto_stitch_rgb_tiles_any.py | 260 ----------------------------- nodes/stitcher_cv2.py | 77 +++++++++ requirements.txt | 3 +- 5 files changed, 82 insertions(+), 402 deletions(-) delete mode 100644 nodes/auto_stitch_rgb_tiles.py delete mode 100644 nodes/auto_stitch_rgb_tiles_any.py create mode 100644 nodes/stitcher_cv2.py diff --git a/__init__.py b/__init__.py index f95c52e..3ac9771 100644 --- a/__init__.py +++ b/__init__.py @@ -17,8 +17,7 @@ from .nodes.stitch_depth import * from .nodes.pointcloud_from_depth import * from .nodes.ply_export import * from .nodes.pointcloud_cylindrical import * -from .nodes.auto_stitch_rgb_tiles import * -from .nodes.auto_stitch_rgb_tiles_any import * +from .nodes.stitcher_cv2 import * NODE_CLASS_MAPPINGS = { "IG Multiply": IG_MultiplyNode, @@ -44,8 +43,7 @@ NODE_CLASS_MAPPINGS = { "IG PointCloud From Depth": IG_PointCloudFromDepth, "IG Save PLY PointCloud": IG_SavePLYPointCloud, "IG PointCloud From Cylindrical": IG_PointCloudCylindricalFromDepth, - "IG Auto Stitch RGB Tiles": IG_AutoStitchRGBTiles, - "IG Auto Stitch RGB Tiles Any": IG_AutoStitchRGBTilesAny, + "IG Stitch Images CV2": IG_StitchImagesCV2, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -72,6 +70,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "IG PointCloud From Depth": "🌐 IG PointCloud From Depth", "IG Save PLY PointCloud": "πŸ’Ύ IG Save PLY PointCloud", "IG PointCloud From Cylindrical": "🌐 IG PointCloud From Cylindrical", - "IG Auto Stitch RGB Tiles": "🧩 IG Auto Stitch RGB Tiles", - "IG Auto Stitch RGB Tiles Any": "🧩 IG Auto Stitch RGB Tiles (Any)", + "IG Stitch Images CV2": "🧩 IG Stitch Images (CV2)", } \ No newline at end of file diff --git a/nodes/auto_stitch_rgb_tiles.py b/nodes/auto_stitch_rgb_tiles.py deleted file mode 100644 index f2e7a6a..0000000 --- a/nodes/auto_stitch_rgb_tiles.py +++ /dev/null @@ -1,135 +0,0 @@ -import torch -import torch.nn.functional as F -from typing import List -from ..common.tree import * # gives TREE_IO - -class IG_AutoStitchRGBTiles: - """ - Given a *batch* of overlapping RGB tiles that together form a single - (wide) image, automatically discover the true horizontal overlap between - neighbouring tiles and stitch them into one seamless picture. - - Assumptions - ----------- - β€’ Tiles are already ordered **left β†’ right** in the batch. - β€’ All tiles share the same height H and width W. - β€’ Mis-alignments are horizontal only (no vertical shift, no rotation). - - Method - ------ - For each consecutive pair of tiles (Tα΅’, Tα΅’β‚Šβ‚) we search horizontal - offsets `s ∈ [min_overlap, max_search]` and choose the *s* that minimises - Mean-Squared-Error between - - Tα΅’[..., Wβˆ’s:W, :] and Tα΅’β‚Šβ‚[..., 0:s, :] - - We then blend with linear ramps so seams vanish. Each pair may discover - a **different** overlap width, so ramps are generated per-pair. - - Inputs - ------ - tiles : IMAGE – tensor [N, H, W, 3] (float 0-1) - min_overlap : INT – smallest overlap to test (default 64) - max_search : INT – largest overlap to test. ≀0 β†’ W-min_overlap - stride : INT – >1 evaluates every *stride*-th pixel when - computing MSE (for speed). - - Outputs - ------- - stitched : IMAGE – [1, H, W_out, 3] (float) - overlaps : INT – list[N-1] of discovered overlap px - """ - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "tiles": ("IMAGE",), - "min_overlap": ("INT", {"default": 64, "min": 1, "step": 1}), - "max_search": ("INT", {"default": 0, "step": 1}), - "stride": ("INT", {"default": 1, "min": 1, "step": 1}), - }, - } - - RETURN_TYPES = ("IMAGE", "INT") - RETURN_NAMES = ("stitched", "overlaps") - FUNCTION = "main" - CATEGORY = TREE_IO # πŸ“ IG Nodes / IO - - # --------------------------------------------------------------- # - # Main # - # --------------------------------------------------------------- # - def main(self, - tiles: torch.Tensor, - min_overlap: int, - max_search: int, - stride: int): - - # ---------- normalise input dims --------------------------- # - if tiles.dim() == 3: - tiles = tiles.unsqueeze(0) # [1,H,W,3] - if tiles.dim() != 4 or tiles.shape[-1] != 3: - raise ValueError("tiles must be IMAGE tensor [N,H,W,3].") - - N, H, W, C = tiles.shape - if N < 2: - return (tiles, []) # nothing to stitch - - if max_search <= 0 or max_search > W - min_overlap: - max_search = W - min_overlap - - # ---------- utility: quickly compute MSE ------------------- # - def mse(a, b): - return torch.mean((a - b) ** 2) - - # ---------- discover overlap for each pair ----------------- # - overlaps: List[int] = [] - for i in range(N - 1): - left = tiles[i, ::stride, ::stride, :] - right = tiles[i+1,::stride, ::stride, :] - - best_s = min_overlap - best_err = float("inf") - - for s in range(min_overlap, max_search + 1): - err = mse(left[:, -s:, :], right[:, :s, :]) - if err < best_err: - best_err, best_s = err.item(), s - - overlaps.append(best_s) - - # ---------- compute x-offsets ------------------------------- # - offsets = [0] - for s in overlaps: - offsets.append(offsets[-1] + (W - s)) - W_out = offsets[-1] + W - - # ---------- allocate output & blend ramps ------------------ # - device, dtype = tiles.device, tiles.dtype - stitched = torch.zeros((1, H, W_out, C), dtype=dtype, device=device) - weight = torch.zeros_like(stitched) - - for idx in range(N): - start = offsets[idx] - end = start + W - mask_1d = torch.ones(W, dtype=dtype, device=device) - - # ramp on left edge (overlap with previous tile) - if idx > 0: - s_left = overlaps[idx-1] - ramp = torch.linspace(0.0, 1.0, s_left, device=device, dtype=dtype) - mask_1d[:s_left] = ramp - - # ramp on right edge (overlap with next tile) - if idx < N - 1: - s_right = overlaps[idx] - ramp = torch.linspace(1.0, 0.0, s_right, device=device, dtype=dtype) - curr = mask_1d[-s_right:] - mask_1d[-s_right:] = torch.minimum(curr, ramp) - - mask = mask_1d.view(1, 1, W, 1) # broadcast over H - stitched[..., start:end, :] += tiles[idx] * mask - weight [..., start:end, :] += mask - - stitched /= torch.clamp_min(weight, 1e-8) - return (stitched, overlaps) \ No newline at end of file diff --git a/nodes/auto_stitch_rgb_tiles_any.py b/nodes/auto_stitch_rgb_tiles_any.py deleted file mode 100644 index 1839643..0000000 --- a/nodes/auto_stitch_rgb_tiles_any.py +++ /dev/null @@ -1,260 +0,0 @@ -import cv2 -import numpy as np -import torch -import torch.nn.functional as F -from typing import List, Tuple - -from ..common.tree import * # πŸ“ TREE_IO constant - - -class IG_AutoStitchRGBTilesAny: - """ - Automatically stitches an unordered batch of horizontally-overlapping RGB - tiles into a single image. The node iteratively grows a 'combined' canvas: - - 1. Pick one tile as the initial canvas. - 2. For every remaining tile, uses the ENTIRE candidate tile as - cv2.matchTemplate template, sliding it over left & right bands - (max_search px) of the current canvas. - 3. Select the tile/side with the *lowest* error below `max_err`. - Blend it in with linear ramps, extend the canvas, repeat (2). - 4. Stop when no tile matches within `max_err`. - - Assumptions - ----------- - β€’ Tiles share the same height 𝐇 (widths may differ). - β€’ Overlaps are purely horizontal (no vertical shift / rotation). - - Inputs - ------ - tiles : IMAGE – [N,H,W,3] or [H,W,3] float 0-1 - min_overlap : INT – smallest overlap to test (β‰₯ 1) - max_search : INT – largest overlap to test (0β‡’auto) (≀W) - stride : INT – subsamples every `stride`-th pixel (β‰₯ 1) - max_err : FLOAT – maximum TM_SQDIFF_NORMED accepted (default 0.1) - debug : BOOL – print detailed progress messages - - Outputs - ------- - stitched : IMAGE – [1,H,W_out,3] - order : INT – list of tile indices in the stitched order - overlaps : INT – list of per-junction overlap widths - """ - - # ----------------------------------------------------------- # - # Tiny helper so we can sprinkle prints without clutter # - # ----------------------------------------------------------- # - def _log(self, dbg: bool, msg: str): - if dbg: - print("[IG-Stitch-DBG]", msg) - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "tiles": ("IMAGE",), - "min_overlap": ("INT", {"default": 64, "min": 1, "max": 99999999, "step": 1}), - "max_search": ("INT", {"default": 0, "min": 0, "max": 99999999, "step": 1}), - "stride": ("INT", {"default": 1, "min": 1, "max": 99999999, "step": 1}), - "max_err": ("FLOAT", {"default": 0.1, "step": 0.01}), - "debug": ("BOOLEAN", {"default": False}), - }, - } - - RETURN_TYPES = ("IMAGE", "INT", "INT") - RETURN_NAMES = ("stitched", "order", "overlaps") - FUNCTION = "main" - CATEGORY = TREE_IO # appears under πŸ“ IG Nodes / IO - - # ----------------------------------------------------------- # - # Helpers # - # ----------------------------------------------------------- # - def _to_gray_u8(self, img: torch.Tensor) -> np.ndarray: - """RGB float tensor [H,W,3] β†’ uint8 grayscale on CPU.""" - img_u8 = (img.cpu().numpy() * 255).astype(np.uint8) - return cv2.cvtColor(img_u8, cv2.COLOR_RGB2GRAY) - - def _slide_template( - self, - template: np.ndarray, - search_band: np.ndarray, - ) -> tuple[float, int]: - """ - Run cv2.matchTemplate (TM_SQDIFF_NORMED) and return - (best_score, best_shift). - best_shift is the horizontal offset of the top-left corner of the - template inside the search band. - """ - res = cv2.matchTemplate(search_band, template, cv2.TM_SQDIFF_NORMED) - min_val, _, min_loc, _ = cv2.minMaxLoc(res) - return float(min_val), int(min_loc[0]) - - # ----------------------------------------------------------- # - - def main( - self, - tiles: torch.Tensor, - min_overlap: int, - max_search: int, - stride: int, - max_err: float, - debug: bool = False, - ): - self._log(debug, f"Starting stitch: {tiles.shape[0]} tiles, min_overlap={min_overlap}, max_search={max_search}") - - # -------- normalise input -------------------------------- # - if tiles.dim() == 3: - tiles = tiles.unsqueeze(0) # [1,H,W,3] - N, H, W0, C = tiles.shape - if N == 1: - return (tiles, [0], []) - - if max_search <= 0: - max_search = max(t.shape[1] for t in tiles) - min_overlap - - # -------- init canvas with first tile -------------------- # - order : List[int] = [0] - overlaps : List[int] = [] - placed : List[bool] = [False] * N - placed[0] = True - - device, dtype = tiles.device, tiles.dtype - stitched = tiles[0:1].clone() # [1,H,W,3] - weight = torch.ones_like(stitched) - - # pre-compute grayscale images for speed - gray_imgs = [self._to_gray_u8(tiles[i, ...]) for i in range(N)] - gray_canvas = gray_imgs[0] - - # -------- iterative placement ---------------------------- # - while not all(placed): - self._log(debug, f"--- iteration {len(order)} | canvas_w={stitched.shape[2]} ---") - - best_idx = -1 - best_side = None # 'left' or 'right' - best_err = max_err - best_shift = 0 - best_band_w = 0 - - # =========================================================== # - # Search all remaining tiles # - # =========================================================== # - canvas_left_band = gray_canvas[:, :max_search] # H Γ— W_band - canvas_right_band = gray_canvas[:, -max_search:] # H Γ— W_band - - for idx in range(N): - if placed[idx]: - continue - tpl = gray_imgs[idx] # whole tile as template - h, w_tpl = tpl.shape - - self._log(debug, f" tile {idx:02d} | w={w_tpl} px") - - # Skip if template wider than band - if w_tpl < min_overlap: - continue - band_w = min(max_search, w_tpl) - band_w = min(max_search, w_tpl, canvas_right_band.shape[1]) - - # ---------- try placing tile to the RIGHT of canvas ----- - score, shift = self._slide_template( - template=tpl, - search_band=canvas_right_band[:, -band_w:], - ) - self._log(debug, f" RIGHT score={score:.4f} shift={shift} band_w={band_w}") - if score < best_err: - best_idx, best_side = idx, "right" - best_err, best_shift = score, shift - best_band_w = band_w - - # ---------- try placing tile to the LEFT of canvas ------ - score, shift = self._slide_template( - template=tpl, - search_band=canvas_left_band[:, :band_w], - ) - self._log(debug, f" LEFT score={score:.4f} shift={shift} band_w={band_w}") - if score < best_err: - best_idx, best_side = idx, "left" - best_err, best_shift = score, shift - best_band_w = band_w - - # -------------------------------------------------------------- # - # Determine overlap & new-columns for the chosen tile # - # -------------------------------------------------------------- # - if best_idx == -1 or best_err > max_err: - self._log(debug, "No candidate beat max_err β€” stopping.") - break # no suitable tile - - tile_w = tiles[best_idx].shape[2] - overlap_px = best_band_w - best_shift - new_cols = tile_w - overlap_px - if new_cols < min_overlap: - self._log(debug, f"Best candidate would add only {new_cols} px β€” stopping.") - break # adds nothing useful - - self._log(debug, f"Chosen tile {best_idx} on {best_side} | err={best_err:.4f} | overlap={overlap_px} | new_cols={new_cols}") - - # -------- blend the chosen tile into canvas ---------- # - tile = tiles[best_idx : best_idx + 1] # keep batch dim - - if best_side == "right": - # ------------------------------------------------------- # - # append tile on the RIGHT side of the canvas # - # ------------------------------------------------------- # - offset = stitched.shape[2] - overlap_px # where tile starts - extra = new_cols # new columns needed - - if extra > 0: # grow canvas - zeros_img = torch.zeros( - (1, H, extra, C), dtype=dtype, device=device - ) - stitched = torch.cat([stitched, zeros_img.clone()], dim=2) - weight = torch.cat([weight, zeros_img.clone()], dim=2) - - # blend with linear ramp - ramp = torch.linspace(1.0, 0.0, overlap_px, device=device, dtype=dtype) - mask_1d = torch.cat( - [torch.ones(tile_w - overlap_px, device=device, dtype=dtype), ramp] - ) - mask = mask_1d.view(1, 1, tile_w, 1) # broadcast β†’ NHWC - - stitched[..., offset : offset + tile_w, :] += tile * mask - weight [..., offset : offset + tile_w, :] += mask - - order.append(best_idx) - overlaps.append(overlap_px) - - else: - # ------------------------------------------------------- # - # prepend tile on the LEFT side of the canvas # - # ------------------------------------------------------- # - extra = new_cols # new columns on left - offset = extra # where old canvas shifts - - if extra > 0: # grow canvas - zeros_img = torch.zeros( - (1, H, extra, C), dtype=dtype, device=device - ) - stitched = torch.cat([zeros_img.clone(), stitched], dim=2) - weight = torch.cat([zeros_img.clone(), weight ], dim=2) - - # blend with linear ramp - ramp = torch.linspace(0.0, 1.0, overlap_px, device=device, dtype=dtype) - mask_1d = torch.cat( - [ramp, torch.ones(tile_w - overlap_px, device=device, dtype=dtype)] - ) - mask = mask_1d.view(1, 1, tile_w, 1) - - stitched[..., :tile_w, :] += tile * mask - weight [..., :tile_w, :] += mask - - order.insert(0, best_idx) - overlaps.insert(0, overlap_px) - - placed[best_idx] = True - # keep a greyscale version of the stitched canvas for next iteration - gray_canvas = self._to_gray_u8(stitched[0, ...]) - - stitched /= torch.clamp_min(weight, 1e-8) - self._log(debug, f"Finished. final_w={stitched.shape[2]} order={order} overlaps={overlaps}") - return (stitched, order, overlaps) \ No newline at end of file diff --git a/nodes/stitcher_cv2.py b/nodes/stitcher_cv2.py new file mode 100644 index 0000000..4d5aea8 --- /dev/null +++ b/nodes/stitcher_cv2.py @@ -0,0 +1,77 @@ +import cv2 +import torch +import numpy as np + +from ..common.tree import * # πŸ“ TREE_IO + + +class IG_StitchImagesCV2: + """ + Wraps OpenCV's high-level `Stitcher` class so you can feed a batch of + partially-overlapping RGB tiles in *any* order and receive a stitched + panorama. + + Inputs + ------ + images : IMAGE – Tensor [N,H,W,3] or [H,W,3] (0-1 float) + mode : ["PANORAMA", "SCANS"] + PANORAMA (default) = camera rotates in place + SCANS = camera translates (e.g. flatbed scan) + + Outputs + ------- + stitched : IMAGE – [1,H_out,W_out,3] (0-1 float) + + Notes + ----- + * Relies on OpenCV contrib (`opencv-contrib-python>=4.5`) because the + Stitcher class lives in the contrib module. + * If stitching fails (status != OK) the node raises a RuntimeError. + """ + + MODES = { + "PANORAMA": cv2.Stitcher_PANORAMA, + "SCANS": cv2.Stitcher_SCANS, + } + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "images": ("IMAGE",), + "mode": (list(cls.MODES.keys()),), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("stitched",) + FUNCTION = "main" + CATEGORY = TREE_IO # shows up under πŸ“ IG Nodes / IO + + # --------------------------------------------------------- # + def _tensor_to_bgr_u8(self, img: torch.Tensor) -> np.ndarray: + """[H,W,3] 0-1 tensor ➜ uint8 BGR ndarray (CPU).""" + img = (img.cpu().numpy() * 255).clip(0, 255).astype(np.uint8) + return cv2.cvtColor(img, cv2.COLOR_RGB2BGR) + + def _bgr_to_tensor(self, img_bgr: np.ndarray) -> torch.Tensor: + """uint8 BGR ➜ float RGB tensor [1,H,W,3].""" + rgb = cv2.cvtColor(img_bgr, cv2.COLOR_BGR2RGB).astype(np.float32) / 255.0 + return torch.from_numpy(rgb).unsqueeze(0) # add batch dim + + # --------------------------------------------------------- # + def main(self, images: torch.Tensor, mode: str): + # normalise dims to [N,H,W,3] + if images.dim() == 3: + images = images.unsqueeze(0) + + imgs_bgr = [self._tensor_to_bgr_u8(im) for im in images] + + stitcher = cv2.Stitcher_create(self.MODES[mode]) + status, pano = stitcher.stitch(imgs_bgr) + + if status != cv2.Stitcher_OK: + raise RuntimeError(f"OpenCV Stitcher failed with status {status}") + + stitched_tensor = self._bgr_to_tensor(pano) + return (stitched_tensor,) \ No newline at end of file diff --git a/requirements.txt b/requirements.txt index 0ab4c53..bef65c4 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1 +1,2 @@ -diffusers>=0.30.0 \ No newline at end of file +diffusers>=0.30.0 +opencv-contrib-python>=4.5 \ No newline at end of file