Add WASWanExposureStabilizer - Stabilizing exposure gain/loss at beginning/end of frame batches.
This commit is contained in:
+497
-117
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
}
|
||||
+1
-1
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user