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.
- 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.