diff --git a/nodes/WASHybridLatentUpscale.py b/nodes/WASHybridLatentUpscale.py index a54762b..70699b4 100644 --- a/nodes/WASHybridLatentUpscale.py +++ b/nodes/WASHybridLatentUpscale.py @@ -1,11 +1,10 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Optional, Tuple +from typing import Optional, Tuple, Literal, Dict, Any import torch import torch.nn.functional as F - import numpy as np try: @@ -17,50 +16,139 @@ except Exception as e: @dataclass class EdgeBlendConfig: - # upscale factor scale: float = 2.0 - # edge detection pre_blur_sigma_px: float = 1.0 - canny_threshold1: int = 100 - canny_threshold2: int = 200 + canny_threshold1: int = 25 + canny_threshold2: int = 155 canny_l2gradient: bool = True - # mask shaping - dilate_radius_px: int = 2 - feather_sigma_px: float = 1.0 + dilate_radius_px: int = 8 + feather_sigma_px: float = 6.0 - # mask clamp mask_min: float = 0.0 mask_max: float = 1.0 - # upscale nearest_exact: bool = True align_corners: Optional[bool] = False - # mask size output_mask_resolution: str = "image" + video_decode_horizontal_tiles: int = 2 + video_decode_vertical_tiles: int = 2 + video_decode_overlap_latent: int = 4 + video_decode_last_frame_fix: bool = False + video_decode_enable_cudnn: bool = True -def upscaleLatentNearestExact(x: torch.Tensor, size: Tuple[int, int]) -> torch.Tensor: + +LatentLayout = Literal["4d_bchw", "5d_bcthw", "5d_btchw"] + + +def is_probable_latent_channels(v: int) -> bool: + return int(v) in (4, 8, 16) + + +def normalize_latent_to_bchw(x: torch.Tensor) -> Tuple[torch.Tensor, Dict[str, Any]]: + """ + Normalize latents to 4D BCHW. + + Accepts: + 4D: [B,C,H,W] + 5D: [B,C,T,H,W] + 5D: [B,T,C,H,W] + + Returns: + x4: [B',C,H,W] where B' = B (image) or B*T (video) + meta: dict used to restore original layout + """ + if not isinstance(x, torch.Tensor): + raise TypeError("latent_samples must be a torch.Tensor") + + if x.dim() == 4: + b, c, h, w = x.shape + return x, {"layout": "4d_bchw", "B": int(b), "C": int(c), "H": int(h), "W": int(w)} + + if x.dim() != 5: + raise ValueError(f"Expected latent 4D or 5D, got {tuple(x.shape)}") + + b = int(x.shape[0]) + + if is_probable_latent_channels(int(x.shape[1])): + c = int(x.shape[1]) + t = int(x.shape[2]) + h = int(x.shape[3]) + w = int(x.shape[4]) + x4 = x.permute(0, 2, 1, 3, 4).contiguous().reshape(b * t, c, h, w) + return x4, {"layout": "5d_bcthw", "B": b, "C": c, "T": t, "H": h, "W": w} + + if is_probable_latent_channels(int(x.shape[2])): + t = int(x.shape[1]) + c = int(x.shape[2]) + h = int(x.shape[3]) + w = int(x.shape[4]) + x4 = x.contiguous().reshape(b * t, c, h, w) + return x4, {"layout": "5d_btchw", "B": b, "C": c, "T": t, "H": h, "W": w} + + c = int(x.shape[1]) + t = int(x.shape[2]) + h = int(x.shape[3]) + w = int(x.shape[4]) + x4 = x.permute(0, 2, 1, 3, 4).contiguous().reshape(b * t, c, h, w) + return x4, {"layout": "5d_bcthw", "B": b, "C": c, "T": t, "H": h, "W": w} + + +def restore_latent_from_bchw(x4: torch.Tensor, meta: Dict[str, Any]) -> torch.Tensor: + """ + Restore latents back to original 4D/5D layout described by meta. + """ + layout = meta["layout"] + + if layout == "4d_bchw": + return x4 + + if x4.dim() != 4: + raise ValueError(f"Expected 4D [B',C,H,W], got {tuple(x4.shape)}") + + b = int(meta["B"]) + c = int(meta["C"]) + t = int(meta["T"]) + h = int(x4.shape[-2]) + w = int(x4.shape[-1]) + + if int(x4.shape[0]) != b * t: + raise ValueError(f"Expected batch {b*t}, got {int(x4.shape[0])}") + if int(x4.shape[1]) != c: + raise ValueError(f"Expected channels {c}, got {int(x4.shape[1])}") + + if layout == "5d_bcthw": + return x4.reshape(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous() + + if layout == "5d_btchw": + return x4.reshape(b, t, c, h, w).contiguous() + + raise ValueError(f"Unknown latent layout: {layout}") + + +def upscale_latent_nearest_exact(x: torch.Tensor, size: Tuple[int, int]) -> torch.Tensor: + """ + Upscale latents with nearest-exact when available. + """ try: return F.interpolate(x, size=size, mode="nearest-exact") except Exception: return F.interpolate(x, size=size, mode="nearest") -def upscaleLatentBilinear(x: torch.Tensor, size: Tuple[int, int], align_corners: Optional[bool]) -> torch.Tensor: +def upscale_latent_bilinear(x: torch.Tensor, size: Tuple[int, int], align_corners: Optional[bool]) -> torch.Tensor: + """ + Upscale latents with bilinear interpolation. + """ return F.interpolate(x, size=size, mode="bilinear", align_corners=align_corners) -def decodeLatentToImageBHWCViaVAE(vae, latent_samples: torch.Tensor) -> torch.Tensor: - images = vae.decode(latent_samples) - if images.dim() == 5: - images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) - return images - - -def _makeEllipticalKernel(radius_px: int): +def make_elliptical_kernel(radius_px: int): + if cv2 is None: + return None r = int(radius_px) if r <= 0: return None @@ -68,66 +156,332 @@ def _makeEllipticalKernel(radius_px: int): return cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (k, k)) -def _clamp01_u8_to_f32(mask_u8: np.ndarray) -> np.ndarray: +def clamp01_u8_to_f32(mask_u8: np.ndarray) -> np.ndarray: return (mask_u8.astype(np.float32) / 255.0).clip(0.0, 1.0) -def _decodeAndComputeTargetPixelSize( +def latent_mask_to_comfy_mask(mask_b1hw: torch.Tensor) -> torch.Tensor: + """ + Convert [B,1,H,W] to Comfy MASK [B,H,W]. + """ + if mask_b1hw.dim() != 4 or int(mask_b1hw.shape[1]) != 1: + raise ValueError(f"Expected [B,1,H,W], got {tuple(mask_b1hw.shape)}") + return torch.clamp(mask_b1hw[:, 0, :, :], 0.0, 1.0) + + +def get_vae_scale_factors(vae) -> Tuple[int, int, int]: + """ + Return (time_scale, width_scale, height_scale) for decode. + """ + df = getattr(vae, "downscale_index_formula", None) + if df: + try: + t, w, h = df + return max(1, int(t)), max(1, int(w)), max(1, int(h)) + except Exception: + pass + + spatial = 1 + temporal = 1 + + scd = getattr(vae, "spacial_compression_decode", None) + if callable(scd): + try: + v = scd() + spatial = 1 if v is None else int(v) + except Exception: + spatial = 1 + + tcd = getattr(vae, "temporal_compression_decode", None) + if callable(tcd): + try: + v = tcd() + temporal = 1 if v is None else int(v) + except Exception: + temporal = 1 + + return max(1, int(temporal)), max(1, int(spatial)), max(1, int(spatial)) + + +def decode_image_latent_regular(vae, latent_bchw: torch.Tensor) -> torch.Tensor: + """ + Decode 4D image latents. + + Args: + vae: ComfyUI VAE object + latent_bchw: [B,C,H,W] + + Returns: + images_bhwc: [B,H,W,C] + """ + if latent_bchw.dim() != 4: + raise ValueError(f"Expected 4D [B,C,H,W], got {tuple(latent_bchw.shape)}") + images = vae.decode(latent_bchw) + if not isinstance(images, torch.Tensor): + raise ValueError("vae.decode did not return a torch.Tensor") + if images.dim() == 5 and int(images.shape[1]) == 1: + images = images[:, 0, :, :, :] + if images.dim() != 4: + raise ValueError(f"Expected decoded [B,H,W,C], got {tuple(images.shape)}") + if int(images.shape[-1]) < 3: + raise ValueError(f"Decoded channels must be >=3, got {int(images.shape[-1])}") + return images + + +def decode_video_latent_lazy_tiled( vae, - latent_samples: torch.Tensor, + latent_bcthw: torch.Tensor, + horizontal_tiles: int, + vertical_tiles: int, + overlap_latent: int, + last_frame_fix: bool, + enable_cudnn: bool, +) -> torch.Tensor: + """ + Decode 5D video latents with spatial tiling. + + Args: + latent_bcthw: [B,C,T,H,W] + + Returns: + images_btHWC: [B,T_out,H_px,W_px,C] + """ + if latent_bcthw.dim() != 5: + raise ValueError(f"Expected 5D [B,C,T,H,W], got {tuple(latent_bcthw.shape)}") + + with torch.backends.cudnn.flags(enabled=bool(enable_cudnn)): + samples = latent_bcthw + b, _, t, h, w = samples.shape + + time_sf, w_sf, h_sf = get_vae_scale_factors(vae) + + if last_frame_fix and t > 0: + last_frame = samples[:, :, -1:, :, :] + samples = torch.cat([samples, last_frame], dim=2) + t = int(samples.shape[2]) + + t_out = 1 + (t - 1) * int(time_sf) + out_h = int(h) * int(h_sf) + out_w = int(w) * int(w_sf) + + horizontal_tiles = max(1, int(horizontal_tiles)) + vertical_tiles = max(1, int(vertical_tiles)) + overlap_latent = max(0, int(overlap_latent)) + + base_tile_h = (int(h) + (vertical_tiles - 1) * overlap_latent) // vertical_tiles + base_tile_w = (int(w) + (horizontal_tiles - 1) * overlap_latent) // horizontal_tiles + + output = None + weights = None + + for vv in range(vertical_tiles): + for hh in range(horizontal_tiles): + w_start = hh * (base_tile_w - overlap_latent) + h_start = vv * (base_tile_h - overlap_latent) + w_end = min(w_start + base_tile_w, int(w)) if hh < horizontal_tiles - 1 else int(w) + h_end = min(h_start + base_tile_h, int(h)) if vv < vertical_tiles - 1 else int(h) + + tile = samples[:, :, :, h_start:h_end, w_start:w_end] + decoded_tile = vae.decode(tile) + if not isinstance(decoded_tile, torch.Tensor): + raise ValueError("vae.decode did not return a torch.Tensor for video tile") + + if decoded_tile.dim() == 4: + decoded_tile = decoded_tile.unsqueeze(1) + elif decoded_tile.dim() != 5: + raise RuntimeError(f"Unexpected decoded tile shape: {tuple(decoded_tile.shape)}") + + if int(decoded_tile.shape[0]) != int(b): + raise RuntimeError("Decoded tile batch mismatch") + + c_out = int(decoded_tile.shape[-1]) + if c_out < 3: + raise RuntimeError("Decoded tile channels must be >=3") + + if output is None: + output = torch.zeros( + (b, t_out, out_h, out_w, c_out), + device=decoded_tile.device, + dtype=decoded_tile.dtype, + ) + weights = torch.zeros( + (b, t_out, out_h, out_w, 1), + device=decoded_tile.device, + dtype=decoded_tile.dtype, + ) + + out_h_start = int(h_start) * int(h_sf) + out_h_end = int(h_end) * int(h_sf) + out_w_start = int(w_start) * int(w_sf) + out_w_end = int(w_end) * int(w_sf) + + expected_h = out_h_end - out_h_start + expected_w = out_w_end - out_w_start + + dec_h = int(decoded_tile.shape[2]) + dec_w = int(decoded_tile.shape[3]) + + if dec_h != expected_h or dec_w != expected_w: + mh = min(dec_h, expected_h) + mw = min(dec_w, expected_w) + decoded_tile = decoded_tile[:, :, :mh, :mw, :] + expected_h = mh + expected_w = mw + out_h_end = out_h_start + mh + out_w_end = out_w_start + mw + + tile_weights = torch.ones( + (b, t_out, expected_h, expected_w, 1), + device=decoded_tile.device, + dtype=decoded_tile.dtype, + ) + + overlap_out_h = min(int(overlap_latent) * int(h_sf), expected_h) + overlap_out_w = min(int(overlap_latent) * int(w_sf), expected_w) + + if hh > 0 and overlap_out_w > 0: + hb = torch.linspace(0, 1, overlap_out_w, device=decoded_tile.device, dtype=decoded_tile.dtype) + tile_weights[:, :, :, :overlap_out_w, :] *= hb.view(1, 1, 1, -1, 1) + if hh < horizontal_tiles - 1 and overlap_out_w > 0: + hb = torch.linspace(1, 0, overlap_out_w, device=decoded_tile.device, dtype=decoded_tile.dtype) + tile_weights[:, :, :, -overlap_out_w:, :] *= hb.view(1, 1, 1, -1, 1) + + if vv > 0 and overlap_out_h > 0: + vb = torch.linspace(0, 1, overlap_out_h, device=decoded_tile.device, dtype=decoded_tile.dtype) + tile_weights[:, :, :overlap_out_h, :, :] *= vb.view(1, 1, -1, 1, 1) + if vv < vertical_tiles - 1 and overlap_out_h > 0: + vb = torch.linspace(1, 0, overlap_out_h, device=decoded_tile.device, dtype=decoded_tile.dtype) + tile_weights[:, :, -overlap_out_h:, :, :] *= vb.view(1, 1, -1, 1, 1) + + t_dec = int(decoded_tile.shape[1]) + if t_dec == t_out: + decoded_for_add = decoded_tile + elif t_dec == 1: + decoded_for_add = decoded_tile.repeat(1, t_out, 1, 1, 1) + else: + if t_out % t_dec == 0: + factor = t_out // t_dec + decoded_for_add = decoded_tile.repeat(1, factor, 1, 1, 1) + else: + if t_dec > t_out: + decoded_for_add = decoded_tile[:, :t_out, :, :, :] + else: + reps = (t_out + t_dec - 1) // t_dec + decoded_for_add = decoded_tile.repeat(1, reps, 1, 1, 1)[:, :t_out, :, :, :] + + output[:, :, out_h_start:out_h_end, out_w_start:out_w_end, :] += decoded_for_add * tile_weights + weights[:, :, out_h_start:out_h_end, out_w_start:out_w_end, :] += tile_weights + + output = output / (weights + 1e-8) + + if bool(last_frame_fix) and int(time_sf) > 0: + output = output[:, :-int(time_sf), :, :, :] + + return output.contiguous() + + +def decode_for_edge_detection( + vae, + latent_bchw_or_flat: torch.Tensor, + meta: Dict[str, Any], + cfg: EdgeBlendConfig, +) -> Tuple[torch.Tensor, Tuple[int, int], int]: + """ + Decode latents for edge detection. + + Returns: + images_bhwc: [B',H,W,C] + (himg, wimg): per-frame decoded size + frames_out: decoded frames per batch (1 for images) + """ + layout = meta.get("layout", "4d_bchw") + + if layout == "4d_bchw": + images = decode_image_latent_regular(vae, latent_bchw_or_flat) + _, himg, wimg, _ = images.shape + return images, (int(himg), int(wimg)), 1 + + b = int(meta["B"]) + c = int(meta["C"]) + t = int(meta["T"]) + h = int(meta["H"]) + w = int(meta["W"]) + + if latent_bchw_or_flat.dim() != 4: + raise ValueError(f"Expected flattened [B*T,C,H,W], got {tuple(latent_bchw_or_flat.shape)}") + if int(latent_bchw_or_flat.shape[0]) != b * t or int(latent_bchw_or_flat.shape[1]) != c: + raise ValueError(f"Video flatten mismatch: expected {(b*t, c, h, w)}, got {tuple(latent_bchw_or_flat.shape)}") + + latent_bcthw = latent_bchw_or_flat.reshape(b, t, c, h, w).permute(0, 2, 1, 3, 4).contiguous() + + images_bt = decode_video_latent_lazy_tiled( + vae=vae, + latent_bcthw=latent_bcthw, + horizontal_tiles=cfg.video_decode_horizontal_tiles, + vertical_tiles=cfg.video_decode_vertical_tiles, + overlap_latent=cfg.video_decode_overlap_latent, + last_frame_fix=cfg.video_decode_last_frame_fix, + enable_cudnn=cfg.video_decode_enable_cudnn, + ) + + b2, t_out, himg, wimg, ch = images_bt.shape + if int(b2) != b: + raise ValueError(f"Decoded batch mismatch: got {int(b2)} expected {b}") + if int(ch) < 3: + raise ValueError("Decoded channels must be >=3") + + images_bhwc = images_bt.reshape(b * t_out, int(himg), int(wimg), int(ch)).contiguous() + return images_bhwc, (int(himg), int(wimg)), int(t_out) + + +def build_edge_masks_opencv( + vae, + latent_samples_bchw_or_flat: torch.Tensor, + meta: Dict[str, Any], target_latent_size: Tuple[int, int], -): + cfg: EdgeBlendConfig, +) -> Tuple[torch.Tensor, torch.Tensor]: + """ + Build edge masks from decoded pixels. - images = decodeLatentToImageBHWCViaVAE(vae, latent_samples) - if images.dim() != 4 or images.shape[-1] < 3: - raise ValueError(f"VAE decode must return [B,H,W,C>=3], got {tuple(images.shape)}") + Returns: + mask_img: [B',1,target_himg,target_wimg] + mask_lat: [B',1,Ht,Wt] + """ + if cv2 is None: + raise RuntimeError(f"OpenCV (cv2) is required but could not be imported: {_CV2_IMPORT_ERROR}") - b, himg, wimg, c = images.shape + images, (himg, wimg), _ = decode_for_edge_detection( + vae=vae, + latent_bchw_or_flat=latent_samples_bchw_or_flat, + meta=meta, + cfg=cfg, + ) - _, _, h_lat, w_lat = latent_samples.shape - if h_lat <= 0 or w_lat <= 0: - raise ValueError("Invalid latent shape.") + if images.dim() != 4 or int(images.shape[-1]) < 3: + raise ValueError(f"Decoded images must be [B',H,W,C>=3], got {tuple(images.shape)}") + + b_prime = int(images.shape[0]) + + if meta["layout"] == "4d_bchw": + _, _, h_lat, w_lat = latent_samples_bchw_or_flat.shape + else: + h_lat = int(meta["H"]) + w_lat = int(meta["W"]) - # pixel-per-latent scaling scale_y = float(himg) / float(h_lat) scale_x = float(wimg) / float(w_lat) ht, wt = target_latent_size - target_himg = int(round(ht * scale_y)) - target_wimg = int(round(wt * scale_x)) + target_himg = max(1, int(round(ht * scale_y))) + target_wimg = max(1, int(round(wt * scale_x))) - target_himg = max(1, target_himg) - target_wimg = max(1, target_wimg) - - return images, (himg, wimg), (target_himg, target_wimg) - - -def buildEdgeMasksOpenCV( - vae, - latent_samples: torch.Tensor, - target_latent_size: Tuple[int, int], - cfg: EdgeBlendConfig, -) -> Tuple[torch.Tensor, torch.Tensor]: - - if cv2 is None: - raise RuntimeError(f"OpenCV (cv2) is required but could not be imported: {_CV2_IMPORT_ERROR}") - - images, (himg, wimg), (target_himg, target_wimg) = _decodeAndComputeTargetPixelSize( - vae=vae, - latent_samples=latent_samples, - target_latent_size=target_latent_size, - ) - - b, _, _, c = images.shape - - dilate_kernel = _makeEllipticalKernel(cfg.dilate_radius_px) + dilate_kernel = make_elliptical_kernel(cfg.dilate_radius_px) masks_img = [] masks_lat = [] - ht, wt = target_latent_size - - for i in range(b): + for i in range(b_prime): img = images[i].detach().cpu().numpy() img_u8 = (np.clip(img, 0.0, 1.0) * 255.0).astype(np.uint8) @@ -138,44 +492,29 @@ def buildEdgeMasksOpenCV( pre_sigma = float(cfg.pre_blur_sigma_px) if pre_sigma > 0.0: - gray = cv2.GaussianBlur( - gray, - (0, 0), - sigmaX=pre_sigma, - sigmaY=pre_sigma, - borderType=cv2.BORDER_REPLICATE, - ) + gray = cv2.GaussianBlur(gray, (0, 0), sigmaX=pre_sigma, sigmaY=pre_sigma, borderType=cv2.BORDER_REPLICATE) edges = cv2.Canny( gray, int(cfg.canny_threshold1), int(cfg.canny_threshold2), L2gradient=bool(cfg.canny_l2gradient), - ) # uint8 0/255 + ) if dilate_kernel is not None and int(cfg.dilate_radius_px) > 0: edges = cv2.dilate(edges, dilate_kernel, iterations=1) feather_sigma = float(cfg.feather_sigma_px) if feather_sigma > 0.0: - edges = cv2.GaussianBlur( - edges, - (0, 0), - sigmaX=feather_sigma, - sigmaY=feather_sigma, - borderType=cv2.BORDER_REPLICATE, - ) + edges = cv2.GaussianBlur(edges, (0, 0), sigmaX=feather_sigma, sigmaY=feather_sigma, borderType=cv2.BORDER_REPLICATE) - # float32 0..1 - mask_f = _clamp01_u8_to_f32(edges) + mask_f = clamp01_u8_to_f32(edges) - # Resize to target decoded pixel size if (target_himg != himg) or (target_wimg != wimg): mask_img_f = cv2.resize(mask_f, (target_wimg, target_himg), interpolation=cv2.INTER_LINEAR) else: mask_img_f = mask_f - # Downsample to target latent size for blending if (target_himg != ht) or (target_wimg != wt): mask_lat_f = cv2.resize(mask_img_f, (wt, ht), interpolation=cv2.INTER_AREA) else: @@ -187,64 +526,89 @@ def buildEdgeMasksOpenCV( masks_img.append(mask_img_f) masks_lat.append(mask_lat_f) - mask_img_np = np.stack(masks_img, axis=0).astype(np.float32) # [B,target_himg,target_wimg] - mask_lat_np = np.stack(masks_lat, axis=0).astype(np.float32) # [B,Ht,Wt] + mask_img_np = np.stack(masks_img, axis=0).astype(np.float32) + mask_lat_np = np.stack(masks_lat, axis=0).astype(np.float32) - mask_img = torch.from_numpy(mask_img_np).unsqueeze(1).to(device=latent_samples.device, dtype=torch.float32) - mask_lat = torch.from_numpy(mask_lat_np).unsqueeze(1).to(device=latent_samples.device, dtype=torch.float32) + mask_img = torch.from_numpy(mask_img_np).unsqueeze(1).to(device=latent_samples_bchw_or_flat.device, dtype=torch.float32) + mask_lat = torch.from_numpy(mask_lat_np).unsqueeze(1).to(device=latent_samples_bchw_or_flat.device, dtype=torch.float32) return mask_img, mask_lat -def latentMaskToComfyMask(mask_b1hw: torch.Tensor) -> torch.Tensor: - if mask_b1hw.dim() != 4 or mask_b1hw.shape[1] != 1: - raise ValueError(f"Expected [B,1,H,W], got {tuple(mask_b1hw.shape)}") - return torch.clamp(mask_b1hw[:, 0, :, :], 0.0, 1.0) - - -def runHybridUpscaleWithOpenCVMask( +def run_hybrid_upscale_with_opencv_mask( latent_samples: torch.Tensor, cfg: EdgeBlendConfig, vae, donor_latent: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """ + Hybrid upscale with edge-based mask blending. - if latent_samples.dim() != 4: - raise ValueError(f"Expected latent [B,C,H,W], got {tuple(latent_samples.shape)}") + Returns: + up_latent: upscaled latent in original layout + mask_img: [B',1,Himg,Wimg] + mask_lat: [B',1,Ht,Wt] + """ + latent_bchw, meta = normalize_latent_to_bchw(latent_samples) - _, _, h, w = latent_samples.shape + donor_bchw = None + if donor_latent is not None: + donor_bchw, donor_meta = normalize_latent_to_bchw(donor_latent) + for k in ("layout", "B", "C"): + if donor_meta.get(k) != meta.get(k): + raise ValueError(f"donor_latent mismatch on {k}: latent={meta.get(k)} donor={donor_meta.get(k)}") + if meta["layout"] != "4d_bchw" and int(donor_meta.get("T", -1)) != int(meta.get("T", -1)): + raise ValueError(f"donor_latent mismatch on T: latent={meta.get('T')} donor={donor_meta.get('T')}") + + _, _, h, w = latent_bchw.shape ht = int(round(h * float(cfg.scale))) wt = int(round(w * float(cfg.scale))) if ht <= 0 or wt <= 0: raise ValueError("Invalid target size computed from scale.") target_latent_size = (ht, wt) - # Base if cfg.nearest_exact: - base = upscaleLatentNearestExact(latent_samples, target_latent_size) + base = upscale_latent_nearest_exact(latent_bchw, target_latent_size) else: - base = F.interpolate(latent_samples, size=target_latent_size, mode="nearest") - - # Optional donor - donor = None - if donor_latent is not None: - if donor_latent.dim() != 4: - raise ValueError("donor_latent must be [B,C,H,W]") - donor = F.interpolate(donor_latent, size=target_latent_size, mode="bilinear", align_corners=cfg.align_corners) - else: - donor = upscaleLatentBilinear(latent_samples, target_latent_size, cfg.align_corners) + base = F.interpolate(latent_bchw, size=target_latent_size, mode="nearest") - # Masks - mask_img, mask_lat = buildEdgeMasksOpenCV( + if donor_bchw is not None: + donor = F.interpolate(donor_bchw, size=target_latent_size, mode="bilinear", align_corners=cfg.align_corners) + else: + donor = upscale_latent_bilinear(latent_bchw, target_latent_size, cfg.align_corners) + + mask_img, mask_lat = build_edge_masks_opencv( vae=vae, - latent_samples=latent_samples, + latent_samples_bchw_or_flat=latent_bchw, + meta=meta, target_latent_size=target_latent_size, cfg=cfg, ) + if meta["layout"] != "4d_bchw": + b = int(meta["B"]) + t = int(meta["T"]) + b_flat = int(latent_bchw.shape[0]) + b_prime = int(mask_lat.shape[0]) + + if b_prime != b_flat: + time_sf, _, _ = get_vae_scale_factors(vae) + time_sf = max(1, int(time_sf)) + + t_out = b_prime // b + m = mask_lat.reshape(b, t_out, 1, ht, wt) + + idx = torch.arange(0, t * time_sf, step=time_sf, device=m.device) + idx = torch.clamp(idx, 0, t_out - 1) + m_sel = torch.index_select(m, dim=1, index=idx) + + mask_lat = m_sel.reshape(b_flat, 1, ht, wt) + mask_lat = mask_lat.to(dtype=base.dtype) - out = base * (1.0 - mask_lat) + donor * mask_lat - return out, mask_img, mask_lat + out_bchw = base * (1.0 - mask_lat) + donor * mask_lat + + out_latent = restore_latent_from_bchw(out_bchw, meta) + return out_latent, mask_img, mask_lat class WASLatentUpscaleHybrid: @@ -270,6 +634,12 @@ class WASLatentUpscaleHybrid: "use_nearest_exact": ("BOOLEAN", {"default": True}), "output_mask_resolution": (["image", "latent"], {"default": "image"}), + + "video_decode_horizontal_tiles": ("INT", {"default": 2, "min": 1, "max": 8}), + "video_decode_vertical_tiles": ("INT", {"default": 2, "min": 1, "max": 8}), + "video_decode_overlap_latent": ("INT", {"default": 4, "min": 0, "max": 32}), + "video_decode_last_frame_fix": ("BOOLEAN", {"default": False}), + "video_decode_enable_cudnn": ("BOOLEAN", {"default": True}), }, "optional": { "donor_latent": ("LATENT",), @@ -296,6 +666,11 @@ class WASLatentUpscaleHybrid: mask_max: float, use_nearest_exact: bool, output_mask_resolution: str, + video_decode_horizontal_tiles: int, + video_decode_vertical_tiles: int, + video_decode_overlap_latent: int, + video_decode_last_frame_fix: bool, + video_decode_enable_cudnn: bool, donor_latent=None, ): latent_samples = latent["samples"] @@ -313,9 +688,14 @@ class WASLatentUpscaleHybrid: mask_max=float(mask_max), nearest_exact=bool(use_nearest_exact), output_mask_resolution=str(output_mask_resolution).strip().lower(), + video_decode_horizontal_tiles=int(video_decode_horizontal_tiles), + video_decode_vertical_tiles=int(video_decode_vertical_tiles), + video_decode_overlap_latent=int(video_decode_overlap_latent), + video_decode_last_frame_fix=bool(video_decode_last_frame_fix), + video_decode_enable_cudnn=bool(video_decode_enable_cudnn), ) - up_latent, mask_img, mask_lat = runHybridUpscaleWithOpenCVMask( + up_latent, mask_img, mask_lat = run_hybrid_upscale_with_opencv_mask( latent_samples=latent_samples, cfg=cfg, vae=vae, @@ -326,9 +706,9 @@ class WASLatentUpscaleHybrid: out["samples"] = up_latent if cfg.output_mask_resolution == "latent": - edge_mask = latentMaskToComfyMask(mask_lat) + edge_mask = latent_mask_to_comfy_mask(mask_lat) else: - edge_mask = latentMaskToComfyMask(mask_img) + edge_mask = latent_mask_to_comfy_mask(mask_img) return (out, edge_mask) diff --git a/nodes/WASWanExposureStabilizer.py b/nodes/WASWanExposureStabilizer.py new file mode 100644 index 0000000..2bc5232 --- /dev/null +++ b/nodes/WASWanExposureStabilizer.py @@ -0,0 +1,472 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Dict, Any, Tuple + +import torch +import torch.nn.functional as F + + +@dataclass +class ExposureStats: + log_mean: torch.Tensor + log_std: torch.Tensor + + +def compute_luma(rgb: torch.Tensor) -> torch.Tensor: + r = rgb[..., 0] + g = rgb[..., 1] + b = rgb[..., 2] + return (0.2126 * r) + (0.7152 * g) + (0.0722 * b) + + +def downscale_for_stats(images_bhwc: torch.Tensor, proxy_size: int) -> torch.Tensor: + b, h, w, c = images_bhwc.shape + if proxy_size <= 0: + return images_bhwc + if h == proxy_size and w == proxy_size: + return images_bhwc + x = images_bhwc.permute(0, 3, 1, 2) + x = F.interpolate(x, size=(proxy_size, proxy_size), mode="area") + return x.permute(0, 2, 3, 1) + + +def compute_exposure_stats( + images_bhwc: torch.Tensor, + eps: float, + clip_low: float, + clip_high: float, +) -> ExposureStats: + rgb = images_bhwc[..., :3].clamp(0.0, 1.0) + luma = compute_luma(rgb).clamp(0.0, 1.0) + + if clip_low > 0.0 or clip_high < 1.0: + luma = luma.clamp(clip_low, clip_high) + + log_luma = torch.log(luma + eps) + log_mean = log_luma.mean(dim=(1, 2)) + log_std = log_luma.std(dim=(1, 2), unbiased=False) + return ExposureStats(log_mean=log_mean, log_std=log_std) + + +def smooth_1d(x: torch.Tensor, window: int) -> torch.Tensor: + window = int(window) + if window <= 1: + return x + if window % 2 == 0: + window += 1 + pad = window // 2 + v = x.view(1, 1, -1) + v = F.pad(v, (pad, pad), mode="replicate") + kernel = torch.ones((1, 1, window), device=x.device, dtype=x.dtype) / float(window) + y = F.conv1d(v, kernel) + return y.view(-1) + + +def find_settle_index( + log_mean: torch.Tensor, + ref_log_mean: float, + tolerance_log: float, + stable_count: int, +) -> int: + b = int(log_mean.numel()) + if b == 0: + return 0 + stable_count = max(int(stable_count), 1) + + diff = (log_mean - float(ref_log_mean)).abs() + within = diff <= float(tolerance_log) + + run = 0 + for i in range(b): + if bool(within[i].item()): + run += 1 + if run >= stable_count: + return i - stable_count + 1 + else: + run = 0 + return b + + +def apply_exposure_correction(images_bhwc: torch.Tensor, gains: torch.Tensor) -> torch.Tensor: + b, h, w, c = images_bhwc.shape + g = gains.view(b, 1, 1, 1).to(dtype=images_bhwc.dtype, device=images_bhwc.device) + rgb = (images_bhwc[..., :3] * g).clamp(0.0, 1.0) + if c > 3: + rest = images_bhwc[..., 3:] + return torch.cat([rgb, rest], dim=-1) + return rgb + + +def ev_to_log(ev: float) -> float: + return float(ev) * float(torch.log(torch.tensor(2.0)).item()) + + +def log_to_ev(logv: torch.Tensor) -> torch.Tensor: + return logv / float(torch.log(torch.tensor(2.0)).item()) + + +def build_anchor_range( + b: int, + anchor_mode: str, + ref_tail_frames: int, + anchor_center: float, + anchor_window: int, +) -> Tuple[int, int]: + b = int(b) + anchor_mode = str(anchor_mode).strip().lower() + + if anchor_mode == "tail": + n = max(int(ref_tail_frames), 1) + n = min(n, b) + return (b - n, b) + + w = max(int(anchor_window), 1) + w = min(w, b) + center = float(anchor_center) + if center < 0.0: + center = 0.0 + if center > 1.0: + center = 1.0 + + cidx = int(round(center * (b - 1))) + start = cidx - (w // 2) + end = start + w + if start < 0: + start = 0 + end = w + if end > b: + end = b + start = b - w + if start < 0: + start = 0 + return (start, end) + + +class WASWanExposureStabilizer: + @classmethod + def INPUT_TYPES(cls) -> Dict[str, Any]: + return { + "required": { + "images": ("IMAGE",), + + "anchor_mode": ( + ["middle", "tail"], + { + "default": "middle", + "tooltip": ( + "How the exposure reference is chosen.\n" + "middle: use a window around anchor_center as the reference (recommended for WAN drift at BOTH start/end).\n" + "tail: use the last ref_tail_frames as the reference (useful when the end is known-stable)." + ), + }, + ), + "ref_tail_frames": ( + "INT", + { + "default": 12, + "min": 1, + "max": 256, + "step": 1, + "tooltip": ( + "Only used when anchor_mode=tail.\n" + "Number of final frames sampled to compute the reference exposure (median log-luma).\n" + "Increase if the tail is stable but noisy; decrease if the tail contains fades/changes." + ), + }, + ), + "anchor_center": ( + "FLOAT", + { + "default": 0.55, + "min": 0.00, + "max": 1.00, + "step": 0.01, + "tooltip": ( + "Only used when anchor_mode=middle.\n" + "Normalized position (0..1) for the center of the anchor window.\n" + "0.50 anchors the middle; 0.55 biases slightly later if the beginning is more unstable." + ), + }, + ), + "anchor_window": ( + "INT", + { + "default": 16, + "min": 1, + "max": 256, + "step": 1, + "tooltip": ( + "Only used when anchor_mode=middle.\n" + "Number of frames in the anchor window used to compute the reference exposure.\n" + "Larger is more robust, but avoid spanning major scene changes." + ), + }, + ), + + "correct_ends": ( + ["start_and_end", "start_only"], + { + "default": "start_and_end", + "tooltip": ( + "Which regions to correct.\n" + "start_only: correct only the initial transient until it settles.\n" + "start_and_end: also correct tail drift if the last stable_count frames are outside tolerance_ev." + ), + }, + ), + "tolerance_ev": ( + "FLOAT", + { + "default": 0.10, + "min": 0.00, + "max": 2.00, + "step": 0.01, + "tooltip": ( + "Stability tolerance in exposure stops (EV).\n" + "Lower = stricter (detects drift longer, may correct more frames).\n" + "Higher = more forgiving (corrects fewer frames, less risk of reacting to content changes)." + ), + }, + ), + "stable_count": ( + "INT", + { + "default": 4, + "min": 1, + "max": 64, + "step": 1, + "tooltip": ( + "How many consecutive frames must be within tolerance_ev to be considered stable.\n" + "Higher values reduce false-stability on noisy sequences but may delay settle detection." + ), + }, + ), + "max_correct_frames": ( + "INT", + { + "default": 20, + "min": 0, + "max": 512, + "step": 1, + "tooltip": ( + "Hard cap on how many frames can be corrected at the start and (if enabled) at the end.\n" + "0 disables correction entirely (stats/report only)." + ), + }, + ), + + "proxy_size": ( + "INT", + { + "default": 96, + "min": 0, + "max": 512, + "step": 1, + "tooltip": ( + "Downscale size used to compute luminance statistics.\n" + "0 uses full resolution (slower). 64–128 is usually sufficient and much faster." + ), + }, + ), + "gain_min": ( + "FLOAT", + { + "default": 0.70, + "min": 0.05, + "max": 2.00, + "step": 0.01, + "tooltip": ( + "Minimum allowed exposure gain applied to any frame.\n" + "Lower values allow stronger darkening correction but can crush highlights if too low." + ), + }, + ), + "gain_max": ( + "FLOAT", + { + "default": 1.30, + "min": 0.05, + "max": 4.00, + "step": 0.01, + "tooltip": ( + "Maximum allowed exposure gain applied to any frame.\n" + "Higher values allow stronger brightening correction but can clip highlights if too high." + ), + }, + ), + "gain_smooth_window": ( + "INT", + { + "default": 5, + "min": 1, + "max": 51, + "step": 2, + "tooltip": ( + "Temporal smoothing window for the per-frame gain curve.\n" + "Use odd values. Larger values reduce pumping but can lag real drift.\n" + "The anchor window is forced to gain=1.0 after smoothing." + ), + }, + ), + + "clip_low": ( + "FLOAT", + { + "default": 0.00, + "min": 0.00, + "max": 0.50, + "step": 0.01, + "tooltip": ( + "Luminance clamp floor used ONLY for computing stats (not applied to output pixels).\n" + "Raise slightly (e.g., 0.02–0.05) to reduce influence of deep blacks/noise on exposure estimation." + ), + }, + ), + "clip_high": ( + "FLOAT", + { + "default": 1.00, + "min": 0.50, + "max": 1.00, + "step": 0.01, + "tooltip": ( + "Luminance clamp ceiling used ONLY for computing stats (not applied to output pixels).\n" + "Lower slightly (e.g., 0.98–0.995) to reduce influence of specular peaks on exposure estimation." + ), + }, + ), + } + } + + RETURN_TYPES = ("IMAGE", "STRING") + RETURN_NAMES = ("images", "report") + FUNCTION = "stabilize" + CATEGORY = "WAS/Video" + + def stabilize( + self, + images: torch.Tensor, + anchor_mode: str = "middle", + ref_tail_frames: int = 12, + anchor_center: float = 0.55, + anchor_window: int = 16, + correct_ends: str = "start_and_end", + tolerance_ev: float = 0.10, + stable_count: int = 4, + max_correct_frames: int = 96, + proxy_size: int = 96, + gain_min: float = 0.70, + gain_max: float = 1.30, + gain_smooth_window: int = 5, + clip_low: float = 0.00, + clip_high: float = 1.00, + ) -> Tuple[torch.Tensor, str]: + if images is None or images.ndim != 4 or images.shape[-1] < 3: + return images, "WASWanExposureStabilizer: invalid IMAGE input" + + b = int(images.shape[0]) + if b <= 1: + return images, "WASWanExposureStabilizer: batch too small (no temporal stabilization needed)" + + device = images.device + images_f32 = images.to(torch.float32) + + proxy = downscale_for_stats(images_f32, int(proxy_size)) + stats = compute_exposure_stats(proxy, eps=1e-6, clip_low=float(clip_low), clip_high=float(clip_high)) + + a0, a1 = build_anchor_range( + b=b, + anchor_mode=str(anchor_mode), + ref_tail_frames=int(ref_tail_frames), + anchor_center=float(anchor_center), + anchor_window=int(anchor_window), + ) + anchor_slice = stats.log_mean[a0:a1] + ref_log_mean = float(anchor_slice.median().item()) + + tolerance_log = ev_to_log(float(tolerance_ev)) + + settle_index = find_settle_index( + log_mean=stats.log_mean, + ref_log_mean=ref_log_mean, + tolerance_log=tolerance_log, + stable_count=int(stable_count), + ) + + tail_diff = (stats.log_mean[-int(stable_count):] - ref_log_mean).abs() + tail_within = bool((tail_diff <= tolerance_log).all().item()) + + gains = torch.ones((b,), device=device, dtype=torch.float32) + needed = torch.exp(torch.tensor(ref_log_mean, device=device) - stats.log_mean.to(device)) + needed = needed.clamp(float(gain_min), float(gain_max)) + + if int(max_correct_frames) <= 0: + max_correct_frames = 0 + + correct_start_upto = min(settle_index, int(max_correct_frames)) if int(max_correct_frames) > 0 else settle_index + do_end = (str(correct_ends).strip().lower() == "start_and_end") + + if correct_start_upto > 0: + gains[:correct_start_upto] = needed[:correct_start_upto] + + end_start = b + if do_end and not tail_within and b > int(stable_count): + within = ((stats.log_mean - ref_log_mean).abs() <= tolerance_log).detach().cpu() + run = 0 + for i in range(b - 1, -1, -1): + if bool(within[i].item()): + run += 1 + if run >= int(stable_count): + end_start = i + int(stable_count) + break + else: + run = 0 + + if end_start >= b: + end_start = b - 1 + + end_len = b - end_start + if int(max_correct_frames) > 0: + end_len = min(end_len, int(max_correct_frames)) + end_start = b - end_len + + if end_len > 0 and end_start < b: + gains[end_start:] = needed[end_start:] + + if int(gain_smooth_window) > 1: + if bool((gains != 1.0).any().item()): + gains = smooth_1d(gains, int(gain_smooth_window)).clamp(float(gain_min), float(gain_max)) + gains[a0:a1] = 1.0 + + corrected = apply_exposure_correction(images, gains) + + ev_delta = log_to_ev(stats.log_mean - ref_log_mean) + ev_delta_cpu = ev_delta.detach().cpu().tolist() + gains_cpu = gains.detach().cpu().tolist() + + tail_last = min(12, b) + tail_ev = ev_delta[-tail_last:].detach().cpu() + tail_min = float(tail_ev.min().item()) + tail_max = float(tail_ev.max().item()) + + report = ( + f"WASWanExposureStabilizer: b={b}, anchor_mode={anchor_mode}, anchor=[{a0},{a1}), " + f"ref_log_mean={ref_log_mean:.6f}, tolerance_ev={tolerance_ev:.3f}, stable_count={stable_count}, " + f"settle_index={settle_index}, corrected_start={correct_start_upto}, " + f"tail_within={tail_within}, correct_ends={correct_ends}, " + f"gain_range=[{min(gains_cpu):.4f},{max(gains_cpu):.4f}], tail_ev_range_last{tail_last}=[{tail_min:+.3f},{tail_max:+.3f}]\n" + f"per_frame_ev_delta_vs_ref={['{:+.3f}'.format(x) for x in ev_delta_cpu]}\n" + f"per_frame_gain={['{:.4f}'.format(x) for x in gains_cpu]}" + ) + + return corrected, report + + +NODE_CLASS_MAPPINGS = { + "WASWanExposureStabilizer": WASWanExposureStabilizer, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "WASWanExposureStabilizer": "WAN 2.2 Exposure Stabilizer", +} diff --git a/pyproject.toml b/pyproject.toml index 90bbb59..f43b3f0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "was-extras" -version = "1.2.4" +version = "1.2.5" description = "A collection of experimental WAS nodes and utilities for ComfyUI." readme = "README.md" requires-python = ">=3.10"