From 7fb6cec7eae662e2ecbeb8e61c95b0f3c6e5cb26 Mon Sep 17 00:00:00 2001
From: Maxed-Out-99
Date: Sun, 22 Feb 2026 08:22:25 -0800
Subject: [PATCH] Add video support and WAN22 video prep + SVI node
Add video handling and prep utilities across nodes and UI.
- mediacomparers.py: add MxdVideoComparer (saves previews, shows A/B videos) and import folder_paths/os/time; register the new node and display name.
- wan22nodes.py: refactor WAN 2.2 scaling with _wan22_pick_bucket and _wan22_scale_image_core, add _select_frames_start_end, expose faster passthrough for scale(), add video prep nodes (WAN22_I2V_Video_Prep_MXD and advanced variant) that scale videos and return start/end/selected frames, improve CombineVideos_MXD audio handling to support multiple audio payload shapes and batching, and mark PreviewVideoMXD as an output node; register new nodes and display names.
- Add new node file wan_svi_first_last_mxd.py: WAN 2.2 First/Last I2V SVI node that prepares SVI conditioning (anchors, padding, masks) and returns latent samples.
- __init__.py: include wan_svi_first_last_mxd in dynamic imports so new node is discovered.
- web/js/image_comparer.js: extend UI to support video/media (playback, sizing, URL builder, mixed image/video lists) and update widget rendering logic for media types.
- pyproject.toml: bump package version to 1.8.1.
These changes add end-to-end video support (preview, compare, prep for WAN 2.2 workflows), improve robustness of audio concatenation, and introduce an SVI-first/last conditioning node for WAN pipelines.
---
__init__.py | 2 +-
mediacomparers.py | 65 ++++++++
pyproject.toml | 2 +-
wan22nodes.py | 330 +++++++++++++++++++++++++++++++++++++-
wan_svi_first_last_mxd.py | 132 +++++++++++++++
web/js/image_comparer.js | 122 ++++++++++----
6 files changed, 618 insertions(+), 35 deletions(-)
create mode 100644 wan_svi_first_last_mxd.py
diff --git a/__init__.py b/__init__.py
index 8a96d49..faf1ba9 100644
--- a/__init__.py
+++ b/__init__.py
@@ -19,7 +19,7 @@ def _get_mappings(mod):
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
-for _name in ("maxedoutnodes", "mediacomparers", "wan22nodes", "loraloader_mxd"):
+for _name in ("maxedoutnodes", "mediacomparers", "wan22nodes", "loraloader_mxd", "wan_svi_first_last_mxd"):
_mod = _safe_import(_name)
_class_map, _display_map = _get_mappings(_mod)
NODE_CLASS_MAPPINGS.update(_class_map)
diff --git a/mediacomparers.py b/mediacomparers.py
index bd767d4..2010021 100644
--- a/mediacomparers.py
+++ b/mediacomparers.py
@@ -1,6 +1,10 @@
"""
Standalone Image Comparer Node extracted from the original project.
"""
+import os
+import time
+
+import folder_paths
from nodes import PreviewImage, SaveImage
# --- Constants ---
@@ -118,13 +122,74 @@ class MxdImageComparerSave(PreviewImage):
return {"ui": result_ui}
return {}
+
+class MxdVideoComparer:
+ """Compare two videos in-node and pass through the new video."""
+
+ NAME = "Video Comparer MXD"
+ CATEGORY = "MXD/video"
+ FUNCTION = "compare_videos"
+ OUTPUT_NODE = True
+ DESCRIPTION = "Compares original/new videos in-node; new video displays first and is passed through."
+ RETURN_TYPES = ("VIDEO",)
+ RETURN_NAMES = ("new_video",)
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ return {
+ "required": {
+ "new_video": ("VIDEO", {"tooltip": "New video to compare. This output is passed through."}),
+ },
+ "optional": {
+ "original_video": ("VIDEO", {"tooltip": "Original video for comparison."}),
+ },
+ }
+
+ def _save_preview_video(self, video_obj, filename):
+ if video_obj is None:
+ return None
+
+ previews_dir = os.path.join(folder_paths.get_output_directory(), "previews")
+ os.makedirs(previews_dir, exist_ok=True)
+
+ preview_path = os.path.join(previews_dir, filename)
+ video_obj.save_to(preview_path, format="mp4", codec="h264")
+ return {
+ "filename": filename,
+ "subfolder": "previews",
+ "type": "output",
+ }
+
+ def compare_videos(self, new_video, original_video=None):
+ stamp = int(time.time() * 1000)
+
+ new_preview = self._save_preview_video(new_video, f"mxd_video_compare_new_{stamp}.mp4")
+ original_preview = self._save_preview_video(
+ original_video,
+ f"mxd_video_compare_original_{stamp}.mp4",
+ )
+
+ result_ui = {
+ "a_videos": [new_preview] if new_preview else [],
+ "b_videos": [original_preview] if original_preview else [],
+ "videos": [],
+ }
+ result_ui["videos"] = result_ui["a_videos"] + result_ui["b_videos"]
+
+ return {
+ "ui": result_ui,
+ "result": (new_video,),
+ }
+
# --- Registration ---
NODE_CLASS_MAPPINGS = {
MxdImageComparerSave.NAME: MxdImageComparerSave,
+ MxdVideoComparer.NAME: MxdVideoComparer,
}
NODE_DISPLAY_NAME_MAPPINGS = {
MxdImageComparerSave.NAME: "Image Comparer + Save MXD",
+ MxdVideoComparer.NAME: "Video Comparer MXD",
}
WEB_DIRECTORY = "."
diff --git a/pyproject.toml b/pyproject.toml
index 00cab17..e037c58 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,7 +1,7 @@
[project]
name = "ComfyUI-MaxedOut"
description = "Custom ComfyUI nodes used in Maxed Out workflows (SDXL, Flux, Wan 2.2, etc.)"
-version = "1.8.0"
+version = "1.8.1"
license = {file = "LICENSE"}
# classifiers = [
# # For OS-independent nodes (works on all operating systems)
diff --git a/wan22nodes.py b/wan22nodes.py
index 8d3d75e..b3c318b 100644
--- a/wan22nodes.py
+++ b/wan22nodes.py
@@ -1266,6 +1266,107 @@ def _wan22_is_valid_dim(w, h):
return (w, h) in _WAN22_VALID_RES
+def _wan22_pick_bucket(iw, ih, tier, crop_to_fit):
+ is_squareish = _is_squareish(iw, ih)
+ is_landscape = iw >= ih
+
+ # --- Square handling ---
+ if is_squareish:
+ if tier == "720p":
+ return (720, 720)
+ return (624, 624)
+
+ # --- Explicit tiers ---
+ if tier == "480p":
+ return _closest_bucket(iw, ih, [(832, 480)] if is_landscape else [(480, 832)], cover=crop_to_fit)
+ if tier == "720p":
+ return _closest_bucket(iw, ih, [(1280, 720)] if is_landscape else [(720, 1280)], cover=crop_to_fit)
+
+ # --- Auto tier logic ---
+ buckets_480 = [(832, 480)] if is_landscape else [(480, 832)]
+ buckets_720 = [(1280, 720)] if is_landscape else [(720, 1280)]
+ iw_ih = iw * ih
+ area_480, area_720 = 832 * 480, 1280 * 720
+ scale_to_480 = abs(iw_ih - area_480) / area_480
+ scale_to_720 = abs(iw_ih - area_720) / area_720
+
+ # prefer minimal scaling
+ if iw <= 832 and ih <= 480:
+ return _closest_bucket(iw, ih, buckets_480, cover=crop_to_fit)
+ return _closest_bucket(iw, ih, buckets_480 if scale_to_480 <= scale_to_720 else buckets_720, cover=crop_to_fit)
+
+
+def _wan22_scale_image_core(image, tier="Auto", crop_to_fit=False):
+ """
+ Shared WAN 2.2 scaler core.
+ Returns (scaled_image, out_w, out_h, did_passthrough).
+ """
+ _, ih, iw, _ = image.shape
+
+ # --- Safe Auto logic ---
+ if tier == "Safe Auto":
+ # passthrough if already WAN-safe
+ if _wan22_is_valid_dim(iw, ih):
+ return image, iw, ih, True
+
+ area = iw * ih
+ area_480, area_720 = 832 * 480, 1280 * 720
+ min_area, max_area = int(area_480 * 0.5), int(area_720 * 1.8)
+
+ if area < min_area or area > max_area:
+ size_label = "small" if area < min_area else "large"
+ raise ValueError(
+ f"[WAN22_I2V_Image_Scaler_MXD] Input resolution {iw}x{ih} is too {size_label} for WAN 2.2 video buckets.\n"
+ "WAN 2.2 works best around:\n"
+ " - 480p tier ~= 832x480 (or 480x832)\n"
+ " - 720p tier ~= 1280x720 (or 720x1280)\n"
+ " - Squares: 624x624 or 720x720\n\n"
+ "Please use a source closer to 480p/720p, or first process it "
+ "through your WAN 2.2 workflow. This ensures extend runs without mismatch."
+ )
+ # fallback to Auto scaling
+ tier = "Auto"
+
+ # --- Normal path (Auto / 480p / 720p) ---
+ bw, bh = _wan22_pick_bucket(iw, ih, tier, crop_to_fit)
+ is_squareish = _is_squareish(iw, ih)
+
+ if is_squareish:
+ crop_to_fit = False
+
+ if crop_to_fit:
+ bw, bh = _safe_hw(_ceil16(bw), _ceil16(bh))
+ out = _resize_then_center_crop(image, bw, bh)
+ else:
+ bw, bh = _safe_hw(_floor16(bw), _floor16(bh))
+ out, _, _ = _resize_fit_inside(image, bw, bh)
+
+ return out, int(out.shape[2]), int(out.shape[1]), False
+
+
+def _select_frames_start_end(frames, count=1, offset=1, mode="end"):
+ total = int(frames.shape[0])
+ if total <= 0:
+ raise ValueError("No frames available for selection.")
+
+ # Clamp offset and count
+ offset = max(1, min(offset, total))
+ count = max(1, min(count, total - offset + 1))
+
+ if mode == "start":
+ start_idx = offset - 1
+ end_idx = start_idx + count
+ selected = frames[start_idx:end_idx].clone()
+ elif mode == "end":
+ start_idx = max(0, total - offset - count + 1)
+ end_idx = start_idx + count
+ selected = frames[start_idx:end_idx].clone()
+ else:
+ raise ValueError(f"Invalid mode '{mode}'. Expected 'start' or 'end'.")
+
+ return selected
+
+
class WAN22_I2V_Image_Scaler_MXD:
"""
MXD Image Scaler for WAN 2.2 (NO PADDING)
@@ -1338,6 +1439,9 @@ class WAN22_I2V_Image_Scaler_MXD:
# Main function
# -----------------------------
def scale(self, image, tier="Auto", crop_to_fit=False):
+ out, _, _, _ = _wan22_scale_image_core(image, tier=tier, crop_to_fit=crop_to_fit)
+ return (out,)
+
_, ih, iw, _ = image.shape
# --- Safe Auto logic ---
@@ -1415,6 +1519,9 @@ class Frames_Select_StartEnd_MXD:
CATEGORY = "MXD/images"
def main(self, frames=None, count=1, offset=1, mode="end"):
+ selected = _select_frames_start_end(frames, count=count, offset=offset, mode=mode)
+ return (selected,)
+
total = frames.shape[0]
# Clamp offset and count
@@ -1493,14 +1600,95 @@ if HAVE_COMFY_API:
# ✅ Correct way: concatenate frame tensors along batch/time dimension (dim=0)
frames_a = torch.stack(comp_a.images) if isinstance(comp_a.images, list) else comp_a.images
frames_b = torch.stack(comp_b.images) if isinstance(comp_b.images, list) else comp_b.images
+ if frames_a.shape[1] != frames_b.shape[1] or frames_a.shape[2] != frames_b.shape[2]:
+ raise ValueError(
+ "Resolution mismatch in CombineVideos_MXD: "
+ f"front_video={frames_a.shape[2]}x{frames_a.shape[1]}, "
+ f"back_video={frames_b.shape[2]}x{frames_b.shape[1]}. "
+ "Use 'WAN 2.2 Video Prep I2V MXD' before WAN generation so scaled base video and generated clip match."
+ )
combined_images = torch.cat([frames_a, frames_b], dim=0)
# ✅ Combine audio sequentially
combined_audio = None
if comp_a.audio is not None or comp_b.audio is not None:
- audio_a = comp_a.audio if comp_a.audio is not None else torch.zeros((1, 0))
- audio_b = comp_b.audio if comp_b.audio is not None else torch.zeros((1, 0))
- combined_audio = torch.cat([audio_a, audio_b], dim=1)
+ def _extract_audio(audio_obj):
+ if audio_obj is None:
+ return None, None, None, None
+ if torch.is_tensor(audio_obj):
+ return audio_obj, None, "tensor", None
+ if isinstance(audio_obj, dict):
+ wave_key = "waveform" if "waveform" in audio_obj else ("samples" if "samples" in audio_obj else None)
+ if wave_key is None or not torch.is_tensor(audio_obj.get(wave_key)):
+ raise TypeError(f"Unsupported audio dict format. Keys: {list(audio_obj.keys())}")
+ return audio_obj[wave_key], audio_obj.get("sample_rate"), "dict", wave_key
+ waveform = getattr(audio_obj, "waveform", None)
+ sample_rate = getattr(audio_obj, "sample_rate", None)
+ if torch.is_tensor(waveform):
+ return waveform, sample_rate, "object", None
+ raise TypeError(f"Unsupported audio payload type: {type(audio_obj).__name__}")
+
+ wave_a, sr_a, kind_a, wave_key_a = _extract_audio(comp_a.audio)
+ wave_b, sr_b, kind_b, wave_key_b = _extract_audio(comp_b.audio)
+ rank_a = wave_a.ndim if wave_a is not None else None
+ rank_b = wave_b.ndim if wave_b is not None else None
+
+ def _to_bct(w):
+ if w is None:
+ return None
+ if w.ndim == 1:
+ return w.unsqueeze(0).unsqueeze(0) # [1,1,T]
+ if w.ndim == 2:
+ return w.unsqueeze(0) # [1,C,T]
+ if w.ndim == 3:
+ return w # [B,C,T]
+ raise ValueError(f"Unsupported audio tensor rank: {w.ndim}")
+
+ wave_a = _to_bct(wave_a)
+ wave_b = _to_bct(wave_b)
+
+ if wave_a is None and wave_b is not None:
+ wave_a = torch.zeros((wave_b.shape[0], wave_b.shape[1], 0), dtype=wave_b.dtype, device=wave_b.device)
+ if wave_b is None and wave_a is not None:
+ wave_b = torch.zeros((wave_a.shape[0], wave_a.shape[1], 0), dtype=wave_a.dtype, device=wave_a.device)
+
+ if wave_a is not None and wave_b is not None:
+ if wave_a.shape[0] != wave_b.shape[0]:
+ if wave_a.shape[0] == 1:
+ wave_a = wave_a.expand(wave_b.shape[0], -1, -1)
+ elif wave_b.shape[0] == 1:
+ wave_b = wave_b.expand(wave_a.shape[0], -1, -1)
+ else:
+ raise ValueError(f"Audio batch mismatch: {wave_a.shape[0]} vs {wave_b.shape[0]}")
+
+ if wave_a.shape[1] != wave_b.shape[1]:
+ if wave_a.shape[1] == 1:
+ wave_a = wave_a.expand(-1, wave_b.shape[1], -1)
+ elif wave_b.shape[1] == 1:
+ wave_b = wave_b.expand(-1, wave_a.shape[1], -1)
+ else:
+ raise ValueError(f"Audio channel mismatch: {wave_a.shape[1]} vs {wave_b.shape[1]}")
+
+ if sr_a is not None and sr_b is not None and sr_a != sr_b:
+ raise ValueError(f"Audio sample-rate mismatch: {sr_a} vs {sr_b}")
+
+ combined_wave = torch.cat([wave_a, wave_b], dim=2)
+ out_sr = sr_a if sr_a is not None else sr_b
+
+ target_rank = rank_a if rank_a is not None else rank_b
+ if target_rank == 1 and combined_wave.shape[0] == 1 and combined_wave.shape[1] == 1:
+ combined_wave = combined_wave.squeeze(0).squeeze(0)
+ elif target_rank == 2 and combined_wave.shape[0] == 1:
+ combined_wave = combined_wave.squeeze(0)
+
+ out_kind = kind_a if kind_a is not None else kind_b
+ if out_kind == "dict":
+ out_key = wave_key_a if kind_a == "dict" else wave_key_b
+ combined_audio = {out_key or "waveform": combined_wave}
+ if out_sr is not None:
+ combined_audio["sample_rate"] = out_sr
+ else:
+ combined_audio = combined_wave
@@ -1513,6 +1701,133 @@ if HAVE_COMFY_API:
)
return (combined_video,)
+
+ class WAN22_I2V_Video_Prep_MXD:
+ """
+ Prepare a source video for iterative WAN 2.2 extension:
+ - scale entire video using WAN bucket logic
+ - output start/end frames from the full scaled video
+ - keep default workflow simple for common use
+ """
+ CATEGORY = "MXD/video"
+ FUNCTION = "prepare"
+ RETURN_TYPES = ("VIDEO", "IMAGE", "IMAGE", "INT", "INT", "FLOAT")
+ RETURN_NAMES = ("scaled_video", "start_image", "end_image", "width", "height", "fps")
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ return {
+ "required": {
+ "video": ("VIDEO",),
+ "tier": (["Auto", "480p", "720p", "Safe Auto"], {"default": "Auto"}),
+ "crop_to_fit": ("BOOLEAN", {
+ "default": True,
+ "label_on": "Perfect Fit (Crops Edges)",
+ "label_off": "Closest Fit (No Crop)"
+ }),
+ },
+ }
+
+ def prepare(self, video, tier="Auto", crop_to_fit=True):
+ comp = video.get_components()
+ if isinstance(comp.images, list):
+ if len(comp.images) == 0:
+ raise ValueError("[WAN22_I2V_Video_Prep_MXD] Input video has zero frames.")
+ frames = torch.stack(comp.images)
+ else:
+ frames = comp.images
+
+ if frames is None:
+ raise ValueError("[WAN22_I2V_Video_Prep_MXD] Input video has no frames.")
+ if frames.ndim == 3:
+ frames = frames.unsqueeze(0)
+ if frames.ndim != 4:
+ raise ValueError(f"[WAN22_I2V_Video_Prep_MXD] Unexpected frame tensor shape: {tuple(frames.shape)}")
+ if frames.shape[0] <= 0:
+ raise ValueError("[WAN22_I2V_Video_Prep_MXD] Input video has zero frames.")
+
+ scaled_frames, out_w, out_h, _ = _wan22_scale_image_core(
+ frames, tier=tier, crop_to_fit=crop_to_fit
+ )
+
+ start_image = scaled_frames[0:1].clone()
+ end_image = scaled_frames[-1:].clone()
+
+ scaled_video = VideoFromComponents(
+ VideoComponents(
+ images=scaled_frames,
+ audio=comp.audio,
+ frame_rate=comp.frame_rate,
+ )
+ )
+
+ fps = float(comp.frame_rate) if comp.frame_rate is not None else 0.0
+ return (scaled_video, start_image, end_image, out_w, out_h, fps)
+
+ class WAN22_I2V_Video_Prep_Advanced_MXD:
+ """
+ Advanced variant of WAN22_I2V_Video_Prep_MXD with frame-selection controls.
+ """
+ CATEGORY = "MXD/video"
+ FUNCTION = "prepare"
+ RETURN_TYPES = ("VIDEO", "IMAGE", "IMAGE", "IMAGE", "INT", "INT", "FLOAT")
+ RETURN_NAMES = ("scaled_video", "selected_frames", "start_image", "end_image", "width", "height", "fps")
+
+ @classmethod
+ def INPUT_TYPES(cls):
+ return {
+ "required": {
+ "video": ("VIDEO",),
+ "tier": (["Auto", "480p", "720p", "Safe Auto"], {"default": "Auto"}),
+ "crop_to_fit": ("BOOLEAN", {
+ "default": True,
+ "label_on": "Perfect Fit (Crops Edges)",
+ "label_off": "Closest Fit (No Crop)"
+ }),
+ "mode": (["start", "end"], {"default": "end"}),
+ "count": ("INT", {"default": 1, "min": 1, "max": 10000}),
+ "offset": ("INT", {"default": 1, "min": 1, "max": 10000}),
+ },
+ }
+
+ def prepare(self, video, tier="Auto", crop_to_fit=True, mode="end", count=1, offset=1):
+ comp = video.get_components()
+ if isinstance(comp.images, list):
+ if len(comp.images) == 0:
+ raise ValueError("[WAN22_I2V_Video_Prep_Advanced_MXD] Input video has zero frames.")
+ frames = torch.stack(comp.images)
+ else:
+ frames = comp.images
+
+ if frames is None:
+ raise ValueError("[WAN22_I2V_Video_Prep_Advanced_MXD] Input video has no frames.")
+ if frames.ndim == 3:
+ frames = frames.unsqueeze(0)
+ if frames.ndim != 4:
+ raise ValueError(f"[WAN22_I2V_Video_Prep_Advanced_MXD] Unexpected frame tensor shape: {tuple(frames.shape)}")
+ if frames.shape[0] <= 0:
+ raise ValueError("[WAN22_I2V_Video_Prep_Advanced_MXD] Input video has zero frames.")
+
+ scaled_frames, out_w, out_h, _ = _wan22_scale_image_core(
+ frames, tier=tier, crop_to_fit=crop_to_fit
+ )
+
+ selected_frames = _select_frames_start_end(
+ scaled_frames, count=count, offset=offset, mode=mode
+ )
+ start_image = selected_frames[0:1].clone()
+ end_image = selected_frames[-1:].clone()
+
+ scaled_video = VideoFromComponents(
+ VideoComponents(
+ images=scaled_frames,
+ audio=comp.audio,
+ frame_rate=comp.frame_rate,
+ )
+ )
+
+ fps = float(comp.frame_rate) if comp.frame_rate is not None else 0.0
+ return (scaled_video, selected_frames, start_image, end_image, out_w, out_h, fps)
# ---------- Load Video MXD (video-only picker with refresh) ----------
class LoadVideoMXD:
@@ -1535,7 +1850,6 @@ if HAVE_COMFY_API:
"remote": {
"route": "/mxd/videos/input",
"refresh_button": True,
- "control_after_refresh": "first",
},
}),
}
@@ -1692,13 +2006,15 @@ if HAVE_COMFY_API:
node_id="PreviewVideoMXD",
display_name="Preview Video MXD",
category="image/video",
- description="Preview a video without saving output.",
+ description="Preview a video without saving output (optional pass-through).",
inputs=[
io.Video.Input("input_video", tooltip="Video to preview."),
],
outputs=[
io.Video.Output("output_video", tooltip="Passes the same video forward."),
],
+ # Allow this node to run even when output_video is not connected.
+ is_output_node=True,
)
@classmethod
@@ -1845,6 +2161,8 @@ NODE_CLASS_MAPPINGS = {
if HAVE_COMFY_API:
NODE_CLASS_MAPPINGS.update({
"Wan22ImageToVideoMXD": Wan22ImageToVideoMXD,
+ "WAN22_I2V_Video_Prep_MXD": WAN22_I2V_Video_Prep_MXD,
+ "WAN22_I2V_Video_Prep_Advanced_MXD": WAN22_I2V_Video_Prep_Advanced_MXD,
"CombineVideos_MXD": CombineVideos_MXD,
"LoadVideoMXD": LoadVideoMXD,
"SaveVideoMXD": SaveVideoMXD,
@@ -1870,6 +2188,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
if HAVE_COMFY_API:
NODE_DISPLAY_NAME_MAPPINGS.update({
"Wan22ImageToVideoMXD": "Wan 2.2 Image to Video MXD",
+ "WAN22_I2V_Video_Prep_MXD": "WAN 2.2 Video Prep I2V MXD",
+ "WAN22_I2V_Video_Prep_Advanced_MXD": "WAN 2.2 Video Prep I2V MXD Advanced",
"CombineVideos_MXD": "Combine Videos MXD",
"LoadVideoMXD": "Load Video MXD",
"SaveVideoMXD": "Save Video MXD",
diff --git a/wan_svi_first_last_mxd.py b/wan_svi_first_last_mxd.py
new file mode 100644
index 0000000..57c07a2
--- /dev/null
+++ b/wan_svi_first_last_mxd.py
@@ -0,0 +1,132 @@
+from comfy_api.latest import io
+import torch
+import node_helpers
+import comfy
+import comfy.latent_formats
+
+
+class Wan22FirstLastImageToVideoSVIMXD(io.ComfyNode):
+ @classmethod
+ def define_schema(cls) -> io.Schema:
+ return io.Schema(
+ node_id="Wan22FirstLastImageToVideoSVIMXD",
+ display_name="WAN 2.2 First/Last I2V SVI MXD",
+ category="ComfyUI-Wan22FMLF",
+ inputs=[
+ io.Conditioning.Input("positive"),
+ io.Conditioning.Input("negative"),
+ io.Vae.Input("vae"),
+ io.Int.Input("length", default=81, min=1, max=8192, step=4),
+ io.Int.Input("batch_size", default=1, min=1, max=4096),
+ io.Image.Input("start_image", optional=True),
+ io.Image.Input("end_image", optional=True),
+ io.Latent.Input("prev_latent", optional=True),
+ io.Int.Input("continue_frames_count", default=5, min=0, max=20, step=1, optional=True),
+ ],
+ outputs=[
+ io.Conditioning.Output(display_name="positive"),
+ io.Conditioning.Output(display_name="negative"),
+ io.Latent.Output(display_name="latent"),
+ ],
+ )
+
+ @classmethod
+ def execute(
+ cls,
+ positive,
+ negative,
+ vae,
+ length,
+ batch_size,
+ start_image=None,
+ end_image=None,
+ prev_latent=None,
+ continue_frames_count=5,
+ ) -> io.NodeOutput:
+ _ = end_image # SVI-only node keeps this input for wiring compatibility and intentionally ignores it.
+
+ spacial_scale = vae.spacial_compression_encode()
+ latent_channels = vae.latent_channels
+ total_latents = ((length - 1) // 4) + 1
+ device = comfy.model_management.intermediate_device()
+
+ prev_samples = None
+ if prev_latent is not None:
+ if not isinstance(prev_latent, dict) or prev_latent.get("samples") is None:
+ raise ValueError("prev_latent was provided but does not contain a valid 'samples' tensor.")
+ prev_samples = prev_latent["samples"]
+
+ if start_image is None and prev_samples is None:
+ raise ValueError("SVI node requires either start_image or prev_latent.")
+
+ if start_image is not None:
+ anchor_latent = vae.encode(start_image[:1, :, :, :3])
+ h = anchor_latent.shape[-2]
+ w = anchor_latent.shape[-1]
+ else:
+ h = prev_samples.shape[-2]
+ w = prev_samples.shape[-1]
+ anchor_latent = torch.zeros(
+ [1, latent_channels, 1, h, w],
+ device=device,
+ )
+
+ latent = torch.zeros(
+ [batch_size, latent_channels, total_latents, h, w],
+ device=device,
+ )
+
+ cond_parts = [anchor_latent]
+ anchor_t = anchor_latent.shape[2]
+ motion_t_limit = max(0, total_latents - anchor_t)
+
+ if prev_samples is not None and continue_frames_count > 0 and motion_t_limit > 0:
+ motion_t = min(continue_frames_count, prev_samples.shape[2])
+ motion_latent = prev_samples[:, :, -motion_t:].clone()
+
+ if motion_latent.shape[-2] != h or motion_latent.shape[-1] != w:
+ raise ValueError("prev_latent spatial size does not match current SVI anchor latent size.")
+
+ if motion_latent.shape[2] > motion_t_limit:
+ motion_latent = motion_latent[:, :, -motion_t_limit:]
+
+ cond_parts.append(motion_latent)
+
+ image_cond_latent = torch.cat(cond_parts, dim=2)
+ padding_size = total_latents - image_cond_latent.shape[2]
+
+ if padding_size > 0:
+ padding = torch.zeros(
+ [1, latent_channels, padding_size, h, w],
+ dtype=image_cond_latent.dtype,
+ device=image_cond_latent.device,
+ )
+ padding = comfy.latent_formats.Wan21().process_out(padding)
+ image_cond_latent = torch.cat([image_cond_latent, padding], dim=2)
+
+ # SVI uses latent-time masks directly (1,1,T,H,W), unlike frame-domain 4-phase mask reshaping used by non-SVI nodes.
+ mask_svi = torch.ones(
+ [1, 1, total_latents, h, w],
+ device=image_cond_latent.device,
+ dtype=image_cond_latent.dtype,
+ )
+ mask_svi[:, :, :1] = 0.0
+
+ conditioning_values = {
+ "concat_latent_image": image_cond_latent,
+ "concat_mask": mask_svi,
+ }
+ positive_out = node_helpers.conditioning_set_values(positive, conditioning_values)
+ negative_out = node_helpers.conditioning_set_values(negative, conditioning_values)
+
+ out_latent = {"samples": latent}
+ return io.NodeOutput(positive_out, negative_out, out_latent)
+
+
+NODE_CLASS_MAPPINGS = {
+ "Wan22FirstLastImageToVideoSVIMXD": Wan22FirstLastImageToVideoSVIMXD,
+}
+
+NODE_DISPLAY_NAME_MAPPINGS = {
+ "Wan22FirstLastImageToVideoSVIMXD": "WAN 2.2 First/Last I2V SVI MXD",
+}
diff --git a/web/js/image_comparer.js b/web/js/image_comparer.js
index c370251..8e551fa 100644
--- a/web/js/image_comparer.js
+++ b/web/js/image_comparer.js
@@ -4,7 +4,11 @@ import { ComfyWidgets } from "../../../scripts/widgets.js";
// --- Constants ---
const NODE_TYPE_STRING = "Image Comparer + Save MXD";
-const NODE_TYPE_STRINGS = new Set([NODE_TYPE_STRING]);
+const NODE_TYPE_STRINGS = new Set([
+ NODE_TYPE_STRING,
+ "Video Comparer MXD",
+ "VideoComparerMXD",
+]);
// --- Canvas Utilities (Inline) ---
function measureText(ctx, str) {
@@ -122,8 +126,12 @@ class MxdBaseWidget {
}
// --- Helper Functions ---
-function imageDataToUrl(data) {
- return api.apiURL(`/view?filename=${encodeURIComponent(data.filename)}&type=${data.type}&subfolder=${data.subfolder}${app.getPreviewFormatParam()}${app.getRandParam()}`);
+function mediaDataToUrl(data, mediaType = "image") {
+ const previewParam = mediaType === "image" ? app.getPreviewFormatParam() : "";
+ const subfolder = data.subfolder || "";
+ return api.apiURL(
+ `/view?filename=${encodeURIComponent(data.filename)}&type=${data.type}&subfolder=${encodeURIComponent(subfolder)}${previewParam}${app.getRandParam()}`
+ );
}
@@ -222,13 +230,14 @@ class MxdImageComparerWidget extends MxdBaseWidget {
if (Array.isArray(v)) {
cleanedVal = v.map((d, i) => {
if (!d || typeof d === "string") {
- d = { url: d, name: i == 0 ? "A" : "B", selected: true };
+ d = { url: d, name: i == 0 ? "A" : "B", selected: true, mediaType: "image" };
}
+ d.mediaType = d.mediaType || "image";
return d;
});
}
else {
- cleanedVal = v.images || [];
+ cleanedVal = (v.images || []).map((d) => ({ mediaType: "image", ...d }));
}
if (cleanedVal.length > 2) {
const hasAAndB = cleanedVal.some((i) => i.name.startsWith("A")) &&
@@ -252,15 +261,49 @@ class MxdImageComparerWidget extends MxdBaseWidget {
get value() {
return this._value;
}
+ createMediaElement(sel) {
+ if (sel.mediaType === "video") {
+ const video = document.createElement("video");
+ video.src = sel.url;
+ video.muted = true;
+ video.loop = true;
+ video.autoplay = true;
+ video.playsInline = true;
+ video.preload = "auto";
+ video.onloadeddata = () => this.node.setDirtyCanvas(true, true);
+ video.play().catch(() => { });
+ return video;
+ }
+ const img = new Image();
+ img.src = sel.url;
+ return img;
+ }
+ getMediaSize(media) {
+ if (!media) {
+ return null;
+ }
+ if (media instanceof HTMLVideoElement) {
+ if (!media.videoWidth || !media.videoHeight) {
+ return null;
+ }
+ return { width: media.videoWidth, height: media.videoHeight };
+ }
+ if (!media.naturalWidth || !media.naturalHeight) {
+ return null;
+ }
+ return { width: media.naturalWidth, height: media.naturalHeight };
+ }
setSelected(selected) {
this._value.images.forEach((d) => (d.selected = false));
this.node.imgs.length = 0;
for (const sel of selected) {
if (!sel.img) {
- sel.img = new Image();
- sel.img.src = sel.url;
+ sel.img = this.createMediaElement(sel);
this.node.imgs.push(sel.img);
}
+ if (sel.mediaType === "video" && sel.img instanceof HTMLVideoElement) {
+ sel.img.play().catch(() => { });
+ }
sel.selected = true;
}
this.selected = selected;
@@ -302,28 +345,31 @@ class MxdImageComparerWidget extends MxdBaseWidget {
if (isClickMode) {
const image = this.selected[this.node.isPointerDown ? 1 : 0];
this.updateAutoSize(image, y);
- this.drawImage(ctx, image, y);
+ this.drawMedia(ctx, image, y);
}
else {
const image = this.selected[0];
this.updateAutoSize(image, y);
- this.drawImage(ctx, image, y);
+ this.drawMedia(ctx, image, y);
if (node.isPointerOver) {
- this.drawImage(ctx, this.selected[1], y, this.node.pointerOverPos[0]);
+ this.drawMedia(ctx, this.selected[1], y, this.node.pointerOverPos[0]);
}
}
+ if (this.selected.some((item) => item.mediaType === "video")) {
+ this.node.setDirtyCanvas(true, false);
+ }
}
updateAutoSize(image, y) {
- var _a, _b;
const initial = this.node._mxdInitialSize;
if (initial && (this.node.size[0] !== initial[0] || this.node.size[1] !== initial[1])) {
return;
}
- if (!((_a = image === null || image === void 0 ? void 0 : image.img) === null || _a === void 0 ? void 0 : _a.naturalWidth) || !((_b = image === null || image === void 0 ? void 0 : image.img) === null || _b === void 0 ? void 0 : _b.naturalHeight)) {
+ const mediaSize = this.getMediaSize(image === null || image === void 0 ? void 0 : image.img);
+ if (!mediaSize) {
return;
}
const nodeWidth = this.node.size[0];
- const imageAspect = image.img.naturalWidth / image.img.naturalHeight;
+ const imageAspect = mediaSize.width / mediaSize.height;
const desiredImageHeight = Math.round(nodeWidth / imageAspect);
const desiredHeight = y + desiredImageHeight;
if (desiredHeight > this.node.size[1]) {
@@ -343,13 +389,13 @@ class MxdImageComparerWidget extends MxdBaseWidget {
}
this.setSelected(selected);
}
- drawImage(ctx, image, y, cropX) {
- var _a, _b;
- if (!((_a = image === null || image === void 0 ? void 0 : image.img) === null || _a === void 0 ? void 0 : _a.naturalWidth) || !((_b = image === null || image === void 0 ? void 0 : image.img) === null || _b === void 0 ? void 0 : _b.naturalHeight)) {
+ drawMedia(ctx, image, y, cropX) {
+ const mediaSize = this.getMediaSize(image === null || image === void 0 ? void 0 : image.img);
+ if (!mediaSize) {
return;
}
let [nodeWidth, nodeHeight] = this.node.size;
- const imageAspect = (image === null || image === void 0 ? void 0 : image.img.naturalWidth) / (image === null || image === void 0 ? void 0 : image.img.naturalHeight);
+ const imageAspect = mediaSize.width / mediaSize.height;
let height = nodeHeight - y;
const widgetAspect = nodeWidth / height;
let targetWidth, targetHeight;
@@ -363,11 +409,11 @@ class MxdImageComparerWidget extends MxdBaseWidget {
targetWidth = height * imageAspect;
offsetX = (nodeWidth - targetWidth) / 2;
}
- const widthMultiplier = (image === null || image === void 0 ? void 0 : image.img.naturalWidth) / targetWidth;
+ const widthMultiplier = mediaSize.width / targetWidth;
const sourceX = 0;
const sourceY = 0;
- const sourceWidth = cropX != null ? (cropX - offsetX) * widthMultiplier : image === null || image === void 0 ? void 0 : image.img.naturalWidth;
- const sourceHeight = image === null || image === void 0 ? void 0 : image.img.naturalHeight;
+ const sourceWidth = cropX != null ? (cropX - offsetX) * widthMultiplier : mediaSize.width;
+ const sourceHeight = mediaSize.height;
const destX = (nodeWidth - targetWidth) / 2;
const destY = y + (height - targetHeight) / 2;
const destWidth = cropX != null ? cropX - offsetX : targetWidth;
@@ -427,7 +473,8 @@ class MxdImageComparer extends MxdBaseServerNode {
return {
name: i === 0 ? "A" : "B",
selected: true,
- url: imageDataToUrl(d),
+ mediaType: "image",
+ url: mediaDataToUrl(d, "image"),
};
}),
};
@@ -435,20 +482,41 @@ class MxdImageComparer extends MxdBaseServerNode {
else {
output.a_images = output.a_images || [];
output.b_images = output.b_images || [];
+ output.a_videos = output.a_videos || [];
+ output.b_videos = output.b_videos || [];
const imagesToChoose = [];
- const multiple = output.a_images.length + output.b_images.length > 2;
+ const total = output.a_images.length + output.b_images.length + output.a_videos.length + output.b_videos.length;
+ const multiple = total > 2;
for (const [i, d] of output.a_images.entries()) {
imagesToChoose.push({
name: output.a_images.length > 1 || multiple ? `A${i + 1}` : "A",
selected: i === 0,
- url: imageDataToUrl(d),
+ mediaType: "image",
+ url: mediaDataToUrl(d, "image"),
});
}
for (const [i, d] of output.b_images.entries()) {
imagesToChoose.push({
name: output.b_images.length > 1 || multiple ? `B${i + 1}` : "B",
selected: i === 0,
- url: imageDataToUrl(d),
+ mediaType: "image",
+ url: mediaDataToUrl(d, "image"),
+ });
+ }
+ for (const [i, d] of output.a_videos.entries()) {
+ imagesToChoose.push({
+ name: output.a_videos.length > 1 || multiple ? `A${i + 1}` : "A",
+ selected: i === 0 && output.a_images.length === 0,
+ mediaType: "video",
+ url: mediaDataToUrl(d, "video"),
+ });
+ }
+ for (const [i, d] of output.b_videos.entries()) {
+ imagesToChoose.push({
+ name: output.b_videos.length > 1 || multiple ? `B${i + 1}` : "B",
+ selected: i === 0 && output.b_images.length === 0,
+ mediaType: "video",
+ url: mediaDataToUrl(d, "video"),
});
}
this.canvasWidget.value = { images: imagesToChoose };
@@ -505,7 +573,7 @@ class MxdImageComparer extends MxdBaseServerNode {
getHelp() {
return `
- The Image Comparer node compares two images on top of each other.
+ The MXD comparer node overlays two media inputs (image or video) for quick A/B checks.
-
@@ -514,9 +582,7 @@ class MxdImageComparer extends MxdBaseServerNode {
- The right-click menu may show image options (Open Image, Save Image, etc.) which will
- correspond to the first image (image_a) if clicked on the left-half of the node, or
- the second image if on the right half of the node.
+ In Slide mode, hover the node to reveal media B. In Click mode, press and hold to toggle A/B.