diff --git a/__init__.py b/__init__.py index 847afd9..75a0bea 100644 --- a/__init__.py +++ b/__init__.py @@ -20,18 +20,12 @@ NODE_CLASS_MAPPINGS = {} NODE_DISPLAY_NAME_MAPPINGS = {} for _name in ( - "maxedoutnodes", - "mediacomparers", - "wan22nodes", + "nodes", "loraloader_mxd", "CharacterPrompts", - "ltxnodes", - "video_preview_mxd", - "model_paths_autoregister_mxd", - "combine_materials_ffgo_mxd", - "save_checkpoint_mxd", - "checkpoint_loader_mxd", "smart_loaders_mxd", + "system.live_preview", + "system.model_paths", ): _mod = _safe_import(_name) _class_map, _display_map = _get_mappings(_mod) diff --git a/checkpoint_loader_mxd.py b/checkpoint_loader_mxd.py deleted file mode 100644 index 53dabdc..0000000 --- a/checkpoint_loader_mxd.py +++ /dev/null @@ -1,41 +0,0 @@ -import folder_paths -import comfy.sd - - -class LoadCheckpointMXD: - DESCRIPTION = ( - "Loads a diffusion model checkpoint, same as the core Load Checkpoint node, " - "with the MXD info-icon UI (CivitAI lookup, cached metadata, local notes)." - ) - TITLE = "Load Checkpoint MXD" - CATEGORY = "MXD/Loaders" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), - } - } - - RETURN_TYPES = ("MODEL", "CLIP", "VAE") - FUNCTION = "load_checkpoint" - - def load_checkpoint(self, ckpt_name): - ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) - out = comfy.sd.load_checkpoint_guess_config( - ckpt_path, - output_vae=True, - output_clip=True, - embedding_directory=folder_paths.get_folder_paths("embeddings"), - ) - return out[:3] - - -NODE_CLASS_MAPPINGS = { - "LoadCheckpointMXD": LoadCheckpointMXD, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "LoadCheckpointMXD": "Load Checkpoint MXD", -} diff --git a/ltxnodes.py b/ltxnodes.py deleted file mode 100644 index 25b8a9b..0000000 --- a/ltxnodes.py +++ /dev/null @@ -1,684 +0,0 @@ -from __future__ import annotations -import os -import re -import base64 -import time -import urllib.error -import urllib.request -from fractions import Fraction -from io import BytesIO -from PIL import Image -from threading import Lock, Thread - -import torch -import torch.nn.functional as F - -import comfy -import comfy.model_management -import comfy.patcher_extension -import comfy.samplers -import comfy.sample -import comfy.utils -import latent_preview -import server -import folder_paths -from comfy_api.latest import VideoFromComponents, VideoComponents - -_serv = server.PromptServer.instance - - -######################################################################################################################## -# LTX Video Empty Latent Image -class LTXVideoEmptyLatentMXD: - DESCRIPTION = "Create an LTX Video empty latent batch from connected width/height and frame count." - TITLE = "LTX Empty Latent Video MXD" - CATEGORY = "MXD/Latent" - - # All dimensions must be multiples of 32 (LTX 32× spatial compression). - # Lengths must be 8n+1 for LTX's 8× temporal compression. - RESOLUTIONS = { - "16:9 Landscape": None, - "16:9 512×288": (512, 288), - "16:9 768×448": (768, 448), - "16:9 832×480": (832, 480), - "16:9 1024×576": (1024, 576), - "16:9 1280×736": (1280, 736), - - "9:16 Portrait": None, - "9:16 288×512": (288, 512), - "9:16 448×768": (448, 768), - "9:16 480×832": (480, 832), - "9:16 576×1024": (576, 1024), - - "4:3 Standard": None, - "4:3 512×384": (512, 384), - "4:3 768×576": (768, 576), - - "1:1 Square": None, - "1:1 512×512": (512, 512), - "1:1 768×768": (768, 768), - } - - def __init__(self): - self.device = comfy.model_management.intermediate_device() - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "length": ( - "INT", - { - "default": 97, - "min": 9, - "max": 1025, - "step": 8, - "tooltip": "Number of frames. Must be 8n+1 (e.g. 25, 49, 73, 97, 121, 201).", - }, - ), - "batch_size": ( - "INT", - { - "default": 1, - "min": 1, - "max": 4096, - "tooltip": "Number of latent videos in the batch.", - }, - ), - }, - "optional": { - "width": ("INT", { - "default": 640, "min": 32, "max": 8192, "step": 32, - "tooltip": "Stage 1 width. Connect the LTX Image Scaler stage1_width output for I2V workflows.", - }), - "height": ("INT", { - "default": 384, "min": 32, "max": 8192, "step": 32, - "tooltip": "Stage 1 height. Connect the LTX Image Scaler stage1_height output for I2V workflows.", - }), - }, - } - - RETURN_TYPES = ("LATENT", "INT") - RETURN_NAMES = ("latent", "length") - FUNCTION = "generate" - - def generate(self, length, batch_size=1, width=640, height=384): - # LTX latent: 128 channels, 32× spatial compression, 8× temporal compression - width = max(32, int(width) // 32 * 32) - height = max(32, int(height) // 32 * 32) - length = max(9, 1 + 8 * round((int(length) - 1) / 8)) - - t = ((length - 1) // 8) + 1 - h = height // 32 - w = width // 32 - latent = torch.zeros([batch_size, 128, t, h, w], device=self.device) - return ({"samples": latent}, length) - - -######################################################################################################################## -# Shared noise helper — equivalent to ComfyUI RandomNoise -class _LTXNoise: - def __init__(self, seed: int): - self.seed = seed - - def generate_noise(self, latent: dict) -> torch.Tensor: - samples = latent["samples"] - batch_inds = latent.get("batch_index", None) - return comfy.sample.prepare_noise(samples, self.seed, batch_inds) - - -######################################################################################################################## -# LTX video preview (taeltx TAE decode) -# -# Core ComfyUI has no preview for the LTXAV format used by LTX 2.3, and the -# latent2rgb approximation looks awful for video. This installs a previewer that -# decodes latent frames with the tiny "taeltx" autoencoder for accurate previews. -# The taeltx model is auto-discovered in the vae / vae_approx model folders. If -# it isn't found, it is downloaded to the configured vae model folder. -# -# TAE decode path borrowed from kjnodes / VideoHelperSuite. - -_TAELTX_FILENAME = "taeltx2_3.safetensors" -_TAELTX_URL = "https://huggingface.co/Kijai/LTX2.3_comfy/resolve/main/vae/taeltx2_3.safetensors?download=true" -_TAELTX_DOWNLOAD_LOCK = Lock() - - -def _find_taeltx_path(folder_paths): - for folder in ("vae", "vae_approx"): - try: - names = folder_paths.get_filename_list(folder) - except Exception: - continue - name = next((fn for fn in names if "taeltx" in fn.lower()), None) - if name is not None: - path = folder_paths.get_full_path(folder, name) - if path: - return path - return None - - -def _download_taeltx(folder_paths): - try: - vae_dirs = folder_paths.get_folder_paths("vae") - except Exception as exc: - print(f"[MXD LTX preview] cannot find ComfyUI vae model folder: {exc}") - return None - - if not vae_dirs: - print("[MXD LTX preview] cannot find ComfyUI vae model folder.") - return None - - target_dir = vae_dirs[0] - target_path = os.path.join(target_dir, _TAELTX_FILENAME) - partial_path = f"{target_path}.part" - - with _TAELTX_DOWNLOAD_LOCK: - if os.path.isfile(target_path): - return target_path - - try: - os.makedirs(target_dir, exist_ok=True) - print(f"[MXD LTX preview] downloading {_TAELTX_FILENAME} to {target_path}") - request = urllib.request.Request(_TAELTX_URL, headers={"User-Agent": "ComfyUI-MaxedOut"}) - with urllib.request.urlopen(request, timeout=120) as response, open(partial_path, "wb") as out: - while True: - chunk = response.read(1024 * 1024) - if not chunk: - break - out.write(chunk) - if not os.path.isfile(partial_path) or os.path.getsize(partial_path) == 0: - raise RuntimeError("downloaded file is empty") - os.replace(partial_path, target_path) - try: - folder_paths.get_filename_list("vae") - except Exception: - pass - print(f"[MXD LTX preview] downloaded {_TAELTX_FILENAME}") - return target_path - except (OSError, RuntimeError, urllib.error.URLError) as exc: - try: - if os.path.exists(partial_path): - os.remove(partial_path) - except OSError: - pass - print(f"[MXD LTX preview] failed to download {_TAELTX_FILENAME}: {exc}") - return None - - -def _load_taeltx(): - """Load the taeltx TAE from the vae / vae_approx model folders. Returns a VAE or None.""" - try: - import folder_paths - from comfy.sd import VAE - except Exception: - return None - - path = _find_taeltx_path(folder_paths) - if not path: - path = _download_taeltx(folder_paths) - if not path: - return None - - try: - taeltx = VAE(comfy.utils.load_torch_file(path)) - taeltx.first_stage_model.show_progress_bar = False - except Exception as exc: - print(f"[MXD LTX preview] failed to load taeltx ({path}): {exc}") - return None - return taeltx - - -class _LTXTAEPreviewer: - """Cycles through LTX video latent frames during sampling, decoding with taeltx.""" - - def __init__(self, taeltx, rate=8): - self.first_preview = True - self.last_time = 0.0 - self.c_index = 0 - self.rate = rate - self.taeltx = taeltx - - def decode_latent_to_preview_image(self, preview_format, x0): - if x0.ndim == 5: - x0 = x0.movedim(2, 1) - x0 = x0.reshape((-1,) + x0.shape[-3:]) - num_images = x0.size(0) - new_time = time.time() - num_previews = int((new_time - self.last_time) * self.rate) - self.last_time += num_previews / self.rate - if num_previews > num_images: - num_previews = num_images - elif num_previews <= 0: - return None - if self.first_preview: - self.first_preview = False - _serv.send_sync( - 'MXD_live_preview_start', - {'length': num_images, 'rate': self.rate, 'id': _serv.last_node_id}, - ) - self.last_time = new_time + 1.0 / self.rate - if self.c_index + num_previews > num_images: - frames = x0.roll(-self.c_index, 0)[:num_previews] - else: - frames = x0[self.c_index:self.c_index + num_previews] - Thread(target=self._send_frames, args=(frames, self.c_index, num_images)).run() - self.c_index = (self.c_index + num_previews) % num_images - return None - - def _send_frames(self, image_tensor, ind, leng): - max_size, min_size = 512, 256 - image_tensor = self._decode(image_tensor) - if image_tensor.size(1) < min_size or image_tensor.size(2) < min_size: - image_tensor = F.interpolate( - image_tensor.movedim(-1, 0), scale_factor=4, mode='nearest' - ).movedim(0, -1) - if image_tensor.size(1) > max_size or image_tensor.size(2) > max_size: - t = image_tensor.movedim(-1, 0) - if t.size(2) < t.size(3): - h = (max_size * t.size(2)) // t.size(3) - t = F.interpolate(t, (h, max_size), mode='nearest') - else: - w = (max_size * t.size(3)) // t.size(2) - t = F.interpolate(t, (max_size, w), mode='nearest') - image_tensor = t.movedim(0, -1) - previews = image_tensor.clamp(0, 1).mul(0xFF).to(device="cpu", dtype=torch.uint8) - node_id = _serv.last_node_id - for preview in previews: - img = Image.fromarray(preview.numpy()) - buf = BytesIO() - img.save(buf, format="JPEG", quality=90) - data_url = "data:image/jpeg;base64," + base64.b64encode(buf.getvalue()).decode("ascii") - _serv.send_sync('MXD_live_preview_frame', {'id': node_id, 'index': ind, 'length': leng, 'data': data_url}) - # taeltx expands the 8× temporal compression on decode - ind = (ind + 1) % ((leng - 1) * 8 + 1) - - def _decode(self, x0): - dev = comfy.model_management.get_torch_device() - dtype = self.taeltx.first_stage_model.decoder[1].weight.dtype - x0 = x0.unsqueeze(0).to(dtype=dtype, device=dev) - return self.taeltx.first_stage_model.decode(x0)[0].permute(1, 2, 3, 0) - - -def _save_final_ltx_preview(node_id, previewer, x0_v, rate): - """Decode the full final clip with taeltx and save it as an mp4 to output/live_previews.""" - try: - frames = x0_v.movedim(2, 1) - frames = frames.reshape((-1,) + frames.shape[-3:]) - frames = previewer._decode(frames).clamp(0, 1).to(device="cpu", dtype=torch.float32) - if frames.ndim != 4 or frames.size(0) == 0: - return - out_dir = os.path.join(folder_paths.get_output_directory(), "live_previews") - os.makedirs(out_dir, exist_ok=True) - safe_id = str(node_id).replace(":", "_").replace("/", "_") - filename = f"{safe_id}_{int(time.time())}.mp4" - path = os.path.join(out_dir, filename) - video = VideoFromComponents(VideoComponents(images=frames, frame_rate=Fraction(max(1, round(rate))))) - video.save_to(path) - print(f"[MXD LTX preview] Saved live preview to {path}") - _serv.send_sync("MXD_live_preview_saved", { - "node_id": node_id, "filename": filename, "subfolder": "live_previews", "type": "output", - }) - except Exception as e: - print(f"[MXD LTX preview] Failed to save live preview: {e}") - - -class _LTXPreviewWrapper: - """OUTER_SAMPLE wrapper that installs the taeltx video previewer during sampling.""" - - def __init__(self, taeltx): - self.taeltx = taeltx - - def __call__(self, executor, noise, latent_image, sampler, sigmas, - denoise_mask, callback, disable_pbar, seed, latent_shapes): - guider = executor.class_obj - device = comfy.model_management.get_torch_device() - self.taeltx.first_stage_model.to(device) - - previewer = _LTXTAEPreviewer(self.taeltx, rate=8) - pbar = comfy.utils.ProgressBar(len(sigmas) - 1) - node_id = _serv.last_node_id - - # Strip I2V guide frames appended at the end of the latent before previewing. - num_keyframes = 0 - if 'positive' in guider.conds and guider.conds['positive']: - kf = guider.conds['positive'][0].get('keyframe_idxs') - if kf is not None: - num_keyframes = len(torch.unique(kf[0, 0, :, 0])) - - def ltx_callback(step, x0, x, total_steps): - x0_v = x0 - if x0_v is not None and len(latent_shapes) > 1: - # Audio+video latents are packed into [B, 1, total]; unpack and - # take the video tensor (the 5D one). Audio is a lower-rank entry. - x0_v = next( - (p for p in comfy.utils.unpack_latents(x0, latent_shapes) if p.ndim == 5), - None, - ) - if x0_v is not None and x0_v.ndim == 5 and num_keyframes > 0: - x0_v = x0_v[:, :, :-num_keyframes] - preview = ( - previewer.decode_latent_to_preview_image("JPEG", x0_v) - if x0_v is not None and x0_v.ndim == 5 else None - ) - pbar.update_absolute(step + 1, total_steps, preview) - if step + 1 >= total_steps and x0_v is not None and x0_v.ndim == 5: - _save_final_ltx_preview(node_id, previewer, x0_v, previewer.rate) - if callback is not None: - callback(step, x0, x, total_steps) - - try: - return executor( - noise, latent_image, sampler, sigmas, denoise_mask, - ltx_callback, disable_pbar, seed, latent_shapes=latent_shapes, - ) - finally: - self.taeltx.first_stage_model.to(comfy.model_management.unet_offload_device()) - - -######################################################################################################################## -# LTX KSampler — Stage 1 (T2V / I2V generation at base resolution) -class LTXKSamplerMXD: - DESCRIPTION = ( - "LTX-Video Stage 1 sampler for the distilled workflow. Use Distilled 8 Step " - "for the trained schedule, or Custom Sigmas when intentionally testing a " - "manual schedule." - ) - TITLE = "LTX Stage 1 Sampler MXD" - CATEGORY = "MXD/Sampling" - - MODES = ["Distilled 8 Step", "Custom Sigmas"] - _DISTILLED_SIGMAS = [1.0, 0.99375, 0.9875, 0.98125, 0.975, - 0.909375, 0.725, 0.421875, 0.0] - _CUSTOM_SIGMAS_DEFAULT = "1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": ("MODEL",), - "positive": ("CONDITIONING",), - "negative": ("CONDITIONING",), - "latent_image": ("LATENT",), - "mode": (cls.MODES, {"default": "Distilled 8 Step"}), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "control_after_generate": True}), - "cfg": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.1}), - "sampler_name": ( - ["euler_ancestral_cfg_pp", "euler_cfg_pp", "euler"], - {"default": "euler_ancestral_cfg_pp"}, - ), - "custom_sigmas": ( - "STRING", - { - "default": cls._CUSTOM_SIGMAS_DEFAULT, - "multiline": True, - "tooltip": "Only used when mode is Custom Sigmas. Enter comma, space, or newline separated sigma values.", - }, - ), - "ltx_preview": ("BOOLEAN", {"default": True, "tooltip": "Show LTX video previews during sampling. Downloads the taeltx VAE to your vae model folder if it is missing."}), - }, - } - - RETURN_TYPES = ("LATENT",) - RETURN_NAMES = ("latent",) - FUNCTION = "sample" - OUTPUT_NODE = False - - def sample( - self, - model, - positive, - negative, - latent_image, - mode="Distilled 8 Step", - seed=0, - cfg=2.0, - sampler_name="euler_ancestral_cfg_pp", - custom_sigmas=_CUSTOM_SIGMAS_DEFAULT, - ltx_preview=True, - ): - sigmas = _select_sigmas( - mode, - { - "Distilled 8 Step": self._DISTILLED_SIGMAS, - }, - custom_sigmas, - "LTX Stage 1 Sampler MXD", - ) - return _run_sampling(model, positive, negative, latent_image, seed, cfg, sampler_name, sigmas, ltx_preview) - - -######################################################################################################################## -# LTX KSampler 2 — Stage 2 (refinement at 2× resolution with distilled LoRA) -class LTXKSampler2MXD: - DESCRIPTION = ( - "LTX-Video Stage 2 refiner for the distilled workflow. Official Refine " - "matches the Lightricks 2.3 two-stage example (start sigma 0.85). " - "Custom Sigmas is for manual testing." - ) - TITLE = "LTX Stage 2 Refiner MXD" - CATEGORY = "MXD/Sampling" - - # Exact stage-2 refine schedule from the official Lightricks 2.3 two-stage - # workflow (LTX-2.3_T2V_I2V_Two_Stage_Distilled.json, euler_cfg_pp, cfg 1). - # Only the starting sigma (denoise strength) is meant to vary; use Custom - # Sigmas for that. - MODES = ["Official Refine", "Custom Sigmas"] - _OFFICIAL_REFINE_SIGMAS = [0.85, 0.725, 0.4219, 0.0] - _CUSTOM_SIGMAS_DEFAULT = "0.85, 0.725, 0.4219, 0.0" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": ("MODEL",), - "positive": ("CONDITIONING",), - "negative": ("CONDITIONING",), - "latent_image": ("LATENT",), - "mode": (cls.MODES, {"default": "Official Refine"}), - "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "control_after_generate": True}), - "cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1}), - "sampler_name": ( - ["euler_cfg_pp", "euler_ancestral_cfg_pp", "euler"], - {"default": "euler_cfg_pp"}, - ), - "custom_sigmas": ( - "STRING", - { - "default": cls._CUSTOM_SIGMAS_DEFAULT, - "multiline": True, - "tooltip": "Only used when mode is Custom Sigmas. Enter comma, space, or newline separated sigma values.", - }, - ), - "ltx_preview": ("BOOLEAN", {"default": True, "tooltip": "Show LTX video previews during sampling. Downloads the taeltx VAE to your vae model folder if it is missing."}), - }, - } - - RETURN_TYPES = ("LATENT",) - RETURN_NAMES = ("latent",) - FUNCTION = "sample" - OUTPUT_NODE = False - - def sample( - self, - model, - positive, - negative, - latent_image, - mode="Official Refine", - seed=0, - cfg=1.0, - sampler_name="euler_cfg_pp", - custom_sigmas=_CUSTOM_SIGMAS_DEFAULT, - ltx_preview=True, - ): - sigmas = _select_sigmas( - mode, - { - "Official Refine": self._OFFICIAL_REFINE_SIGMAS, - }, - custom_sigmas, - "LTX Stage 2 Refiner MXD", - ) - return _run_sampling(model, positive, negative, latent_image, seed, cfg, sampler_name, sigmas, ltx_preview) - - -######################################################################################################################## -# Sigma schedule helpers -_SIGMA_RE = re.compile(r"[-+]?(?:\d*\.\d+|\d+\.?)(?:[eE][-+]?\d+)?") - - -def _select_sigmas(mode, presets, custom_sigmas, node_name): - if mode == "Custom Sigmas": - values = _parse_custom_sigmas(custom_sigmas, node_name) - else: - try: - values = presets[mode] - except KeyError as exc: - allowed = ", ".join([*presets.keys(), "Custom Sigmas"]) - raise ValueError(f"{node_name}: unknown mode '{mode}'. Expected one of: {allowed}.") from exc - - return torch.tensor(values, dtype=torch.float32) - - -def _parse_custom_sigmas(custom_sigmas, node_name): - text = str(custom_sigmas or "") - values = [float(match.group(0)) for match in _SIGMA_RE.finditer(text)] - - if len(values) < 2: - raise ValueError(f"{node_name}: Custom Sigmas needs at least two sigma values, ending with 0.0.") - - for index, (left, right) in enumerate(zip(values, values[1:]), start=1): - if right > left: - raise ValueError( - f"{node_name}: Custom Sigmas must be in descending order. " - f"Value {index + 1} ({right}) is greater than value {index} ({left})." - ) - - if abs(values[-1]) > 1e-8: - raise ValueError(f"{node_name}: Custom Sigmas must end with 0.0.") - - return values - - -######################################################################################################################## -# Shared sampling logic -def _run_sampling(model, positive, negative, latent_image, seed, cfg, sampler_name, sigmas, ltx_preview=False): - taeltx = _load_taeltx() if ltx_preview else None - if ltx_preview and taeltx is None: - print("[MXD LTX preview] taeltx model not found in vae / vae_approx — skipping preview.") - - if taeltx is not None: - model = model.clone() - model.add_wrapper_with_key( - comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, - "ltx_mxd_preview", - _LTXPreviewWrapper(taeltx), - ) - - guider = comfy.samplers.CFGGuider(model) - guider.set_conds(positive, negative) - guider.set_cfg(cfg) - - sampler = comfy.samplers.sampler_object(sampler_name) - - latent = latent_image.copy() - latent_samples = latent["samples"] - - try: - latent_samples = comfy.sample.fix_empty_latent_channels( - guider.model_patcher, latent_samples, - latent.get("downscale_ratio_spacial", None), - ) - except AttributeError: - pass - - latent["samples"] = latent_samples - noise_mask = latent.get("noise_mask", None) - - noise = _LTXNoise(seed) - - if taeltx is not None: - # The preview wrapper owns the progress bar / callback. - callback = None - else: - x0_output = {} - callback = latent_preview.prepare_callback(guider.model_patcher, sigmas.shape[-1] - 1, x0_output) - - disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED - - samples = guider.sample( - noise.generate_noise(latent), - latent_samples, - sampler, - sigmas, - denoise_mask=noise_mask, - callback=callback, - disable_pbar=disable_pbar, - seed=seed, - ) - samples = samples.to(comfy.model_management.intermediate_device()) - - out = latent.copy() - out.pop("downscale_ratio_spacial", None) - out["samples"] = samples - return (out,) - - -######################################################################################################################## -# LTX Preview — attach the taeltx previewer to any model -class LTXPreviewMXD: - DESCRIPTION = ( - "Enables taeltx video previews during sampling for ANY sampler node " - "(SamplerCustomAdvanced, KSampler, etc.), not just the MXD LTX samplers. " - "LTX 2.3 (LTXAV) ships no built-in preview decoder, so core ComfyUI shows " - "nothing; this attaches a wrapper to the model that decodes latent frames " - "with the tiny taeltx autoencoder. Wire it between your model loader and " - "the sampler's model input. Downloads taeltx to your vae folder if missing." - ) - TITLE = "LTX Preview MXD" - CATEGORY = "MXD/Sampling" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "model": ("MODEL",), - "enabled": ("BOOLEAN", {"default": True, "tooltip": "Turn taeltx previews on/off without unwiring the node."}), - }, - } - - RETURN_TYPES = ("MODEL",) - RETURN_NAMES = ("model",) - FUNCTION = "apply" - OUTPUT_NODE = False - - def apply(self, model, enabled=True): - if not enabled: - return (model,) - taeltx = _load_taeltx() - if taeltx is None: - print("[MXD LTX preview] taeltx model not found in vae / vae_approx — skipping preview.") - return (model,) - model = model.clone() - model.add_wrapper_with_key( - comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, - "ltx_mxd_preview", - _LTXPreviewWrapper(taeltx), - ) - return (model,) - - -######################################################################################################################## -NODE_CLASS_MAPPINGS = { - "LTXVideoEmptyLatent_MXD": LTXVideoEmptyLatentMXD, - "LTXKSampler_MXD": LTXKSamplerMXD, - "LTXKSampler2_MXD": LTXKSampler2MXD, - "LTXPreview_MXD": LTXPreviewMXD, -} - -NODE_DISPLAY_NAME_MAPPINGS = { - "LTXVideoEmptyLatent_MXD": "LTX Empty Latent Video MXD", - "LTXKSampler_MXD": "LTX Stage 1 Sampler MXD", - "LTXKSampler2_MXD": "LTX Stage 2 Refiner MXD", - "LTXPreview_MXD": "LTX Preview MXD", -} diff --git a/maxedoutnodes.py b/maxedoutnodes.py deleted file mode 100644 index aaf6856..0000000 --- a/maxedoutnodes.py +++ /dev/null @@ -1,2196 +0,0 @@ -from __future__ import annotations -import torch, math, comfy, os, folder_paths, node_helpers, comfy.model_management, comfy.utils, json, hashlib, re, random -import torch.nn.functional as F -from comfy.comfy_types import IO, ComfyNodeABC, InputTypeDict -import numpy as np -from PIL import Image, ImageOps, ImageSequence, ImageFilter, ImageColor -from nodes import SaveImage -try: - from comfy_api.latest import io - HAVE_COMFY_API = True -except Exception as _e: - io = None - HAVE_COMFY_API = False - print(f"[ComfyUI-MaxedOut] comfy_api not available in maxedoutnodes: {_e}") - -try: - from comfy_api.input_impl import VideoFromFile - HAVE_COMFY_API_VIDEO = True -except Exception as _e: - VideoFromFile = None - HAVE_COMFY_API_VIDEO = False - print(f"[ComfyUI-MaxedOut] comfy_api video I/O not available in maxedoutnodes: {_e}") - -######################################################################################################################## -# Flux Empty Latent Image (SD3-compatible) -class FluxEmptyLatentImage: - DESCRIPTION = """Select a Flux resolution and create an empty latent batch.""" - TITLE = "Flux Empty Latent Image" - CATEGORY = "MXD/Latent" - - RESOLUTIONS = { - "— High Resolutions —": None, - "Square (1:1) 1408x1408": (1408, 1408), - "Standard (4:3) 1664x1216": (1664, 1216), - "Landscape (3:2) 1728x1152": (1728, 1152), - "Widescreen (16:9) 1920x1088": (1920, 1088), - "Ultrawide (21:9) 2176x960": (2176, 960), - - "— Standard Resolutions —": None, - "Square (1:1) 1024x1024": (1024, 1024), - "Standard (4:3) 1152x896": (1152, 896), - "Landscape (3:2) 1216x832": (1216, 832), - "Widescreen (16:9) 1344x768": (1344, 768), - "Ultrawide (21:9) 1536x640": (1536, 640), - - "— Low Resolutions —": None, - "Square (1:1) 320x320": (320, 320), - "Standard (4:3) 448x320": (448, 320), - "Landscape (3:2) 384x256": (384, 256), - "Widescreen (16:9) 448x256": (448, 256), - "Ultrawide (21:9) 576x256": (576, 256), - } - - def __init__(self): - self.device = comfy.model_management.intermediate_device() - - @classmethod - def INPUT_TYPES(cls) -> dict: - return { - "required": { - "resolution": ( - list(cls.RESOLUTIONS.keys()), - {"default": "Square (1:1) 1024x1024"} - ), - "vertical": ("BOOLEAN", {"default": False}), - "batch_size": ( - "INT", - { - "default": 1, - "min": 1, - "max": 4096, - "tooltip": "The number of latent images in the batch." - } - ) - } - } - - RETURN_TYPES = ("LATENT",) - OUTPUT_TOOLTIPS = ("The empty latent image batch.",) - FUNCTION = "generate" - - def generate(self, resolution, vertical, batch_size=1) -> tuple: - size = self.RESOLUTIONS.get(resolution) - if size is None: - raise ValueError(f"'{resolution}' is a header or invalid option.") - - width, height = size - if vertical: - width, height = height, width - - latent = torch.zeros([batch_size, 16, height // 8, width // 8], device=self.device) - return ({"samples": latent},) - -######################################################################################################################## -# Flux 2 Empty Latent Image (Flux2-compatible) -class Flux2EmptyLatentImage: - DESCRIPTION = """Select a Flux resolution and create an empty Flux 2 latent batch.""" - TITLE = "Flux 2 Empty Latent Image" - CATEGORY = "MXD/Latent" - - RESOLUTIONS = FluxEmptyLatentImage.RESOLUTIONS - - def __init__(self): - self.device = comfy.model_management.intermediate_device() - - @classmethod - def INPUT_TYPES(cls) -> dict: - return { - "required": { - "resolution": ( - list(cls.RESOLUTIONS.keys()), - {"default": "Square (1:1) 1024x1024"} - ), - "vertical": ("BOOLEAN", {"default": False}), - "batch_size": ( - "INT", - { - "default": 1, - "min": 1, - "max": 4096, - "tooltip": "The number of latent images in the batch." - } - ) - } - } - - RETURN_TYPES = ("LATENT",) - OUTPUT_TOOLTIPS = ("The empty Flux 2 latent image batch.",) - FUNCTION = "generate" - - def generate(self, resolution, vertical, batch_size=1) -> tuple: - size = self.RESOLUTIONS.get(resolution) - if size is None: - raise ValueError(f"'{resolution}' is a header or invalid option.") - - width, height = size - if vertical: - width, height = height, width - - latent = torch.zeros([batch_size, 128, height // 16, width // 16], device=self.device) - return ({"samples": latent},) - -######################################################################################################################## -# Flux Resolution Selector (for feeding into FluxEmptyLatentImage) -class FluxResolutionSelector: - DESCRIPTION = """Pick a Flux resolution string for Flux Empty Latent Image.""" - TITLE = "Flux Resolution Selector" - CATEGORY = "MXD/Latent" - - @classmethod - def INPUT_TYPES(cls) -> dict: - return { - "required": { - "resolution": ( - list(FluxEmptyLatentImage.RESOLUTIONS.keys()), # Include ALL keys including headers - {"default": "Square (1:1) 1024x1024"} - ), - } - } - - RETURN_TYPES = (list(FluxEmptyLatentImage.RESOLUTIONS.keys()),) - RETURN_NAMES = ("resolution",) - OUTPUT_TOOLTIPS = ("The selected resolution string for FluxEmptyLatentImage.",) - FUNCTION = "select_resolution" - - def select_resolution(self, resolution) -> tuple: - return (resolution,) - -######################################################################################################################## -# Sdxl Empty Latent Image -class SdxlEmptyLatentImage: - DESCRIPTION = """Select an SDXL resolution and create an empty latent batch.""" - TITLE = "Sdxl Empty Latent Image (With Resolutions)" - CATEGORY = "MXD/Latent" - - # SDXL predefined resolutions (width, height) - RESOLUTIONS = { - "Square (1:1) 1024x1024": (1024, 1024), - "Standard (4:3) 1152x896": (1152, 896), - "Landscape (3:2) 1216x832": (1216, 832), - "Widescreen (16:9) 1344x768": (1344, 768), - "Ultra-Wide (21:9) 1536x640": (1536, 640), - } - - def __init__(self): - # Retrieve the intermediate device (usually the GPU) from ComfyUI's model management. - self.device = comfy.model_management.intermediate_device() - - @classmethod - def INPUT_TYPES(cls) -> dict: - return { - "required": { - # Dropdown selection for one of the predefined SDXL resolutions. - "resolution": (list(cls.RESOLUTIONS.keys()),), - # Toggle for vertical mode (swaps width and height). - "vertical": ("BOOLEAN", {"default": False}), - # Number of latent images to create in the batch. - "batch_size": ( - "INT", - { - "default": 1, - "min": 1, - "max": 4096, - "tooltip": "The number of latent images in the batch." - } - ) - } - } - - RETURN_TYPES = ("LATENT",) - OUTPUT_TOOLTIPS = ("The empty latent image batch.",) - FUNCTION = "generate" - - def generate(self, resolution, vertical, batch_size=1) -> tuple: - # Get the selected resolution tuple (width, height) - width, height = self.RESOLUTIONS[resolution] - # If vertical mode is enabled, swap width and height. - if vertical: - width, height = height, width - - # Create an empty latent tensor. - # Typically, the latent space has 4 channels and each spatial dimension is 1/8th of the image. - latent = torch.zeros([batch_size, 4, height // 8, width // 8], device=self.device) - return ({"samples": latent},) - -######################################################################################################################## -# Z-Image Turbo Empty Latent Image (SD3-compatible) — Flux-style grouping -class ZImageTurboEmptyLatentImage: - DESCRIPTION = """Select a Z-Image Turbo resolution and create an empty latent batch.""" - TITLE = "Z-Image Turbo Empty Latent Image" - CATEGORY = "MXD/Latent" - - # Tuned for Z-Image Turbo: - # - Rule of 64: every dimension is a multiple of 64 - # - 1MP baseline: 1024x1024 in the standard tier - # - Ceiling: keep presets below 6.5MP - MAX_TOTAL_PIXELS = 6_500_000 - MIN_BLOCK = 64 - RESOLUTIONS = { - "— High Resolutions —": None, - "Square (1:1) 1536x1536": (1536, 1536), - "Photo (4:3) 1792x1344": (1792, 1344), - "Landscape (3:2) 1920x1280": (1920, 1280), - "Widescreen (16:9) 2048x1152": (2048, 1152), - "Ultrawide (21:9) 2304x1024": (2304, 1024), - - "— Standard Resolutions —": None, - "Square (1:1) 1024x1024": (1024, 1024), - "Photo (4:3) 1152x896": (1152, 896), - "Landscape (3:2) 1280x832": (1280, 832), - "Widescreen (16:9) 1344x768": (1344, 768), - "Ultrawide (21:9) 1536x640": (1536, 640), - - "— Low Resolutions —": None, - "Square (1:1) 512x512": (512, 512), - "Photo (4:3) 576x448": (576, 448), - "Landscape (3:2) 640x448": (640, 448), - "Widescreen (16:9) 704x384": (704, 384), - "Ultrawide (21:9) 768x320": (768, 320), - } - - def __init__(self): - self.device = comfy.model_management.intermediate_device() - - @classmethod - def INPUT_TYPES(cls) -> dict: - return { - "required": { - "resolution": ( - list(cls.RESOLUTIONS.keys()), - {"default": "Square (1:1) 1024x1024"} - ), - "vertical": ( - "BOOLEAN", - {"default": False, "tooltip": "Swap width and height."} - ), - "batch_size": ( - "INT", - {"default": 1, "min": 1, "max": 4096, "tooltip": "Number of latent images in the batch."} - ) - } - } - - RETURN_TYPES = ("LATENT",) - OUTPUT_TOOLTIPS = ("The empty Z-Image Turbo latent batch.",) - FUNCTION = "generate" - - def generate(self, resolution, vertical, batch_size=1) -> tuple: - size = self.RESOLUTIONS.get(resolution) - if size is None: - raise ValueError(f"'{resolution}' is a header or invalid option.") - - width, height = size - if vertical: - width, height = height, width - - if (width % self.MIN_BLOCK) != 0 or (height % self.MIN_BLOCK) != 0: - raise ValueError( - f"Invalid preset {width}x{height}. Z-Image Turbo requires multiples of {self.MIN_BLOCK}." - ) - if (width * height) > self.MAX_TOTAL_PIXELS: - raise ValueError( - f"Invalid preset {width}x{height}. Z-Image Turbo presets must stay at or below {self.MAX_TOTAL_PIXELS:,} pixels." - ) - - latent = torch.zeros([batch_size, 16, height // 8, width // 8], device=self.device) - return ({"samples": latent},) - -######################################################################################################################## -# Image Scale To Total Pixels (SDXL Safe) -class SDXLImageScaleToTotalPixelsSafe: - DESCRIPTION = """Scale to a target megapixel count and keep aspect ratio. Skips SDXL-safe sizes.""" - upscale_methods = ["bilinear", "bicubic", "lanczos", "nearest-exact", "area"] - - # SDXL-safe resolutions (width, height) – store one orientation only, - # the code will check both (w, h) and (h, w) - SDXL_SAFE_RESOLUTIONS = [ - (1024, 1024), - (1152, 896), - (1216, 832), - (1344, 768), - (1536, 640), - ] - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "upscale_method": (cls.upscale_methods, {"default": "bilinear"}), - "total_megapixels": ( - "FLOAT", - { - "default": 1.0, - "min": 0.01, - "max": 128.0, - "step": 0.01, - "tooltip": "Set the total megapixels (e.g., 1.0 = 1 MP)", - }, - ), - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "upscale" - CATEGORY = "MXD/Upscaling" - - def upscale(self, image, upscale_method, total_megapixels): - if upscale_method in ["nearest-exact", "area"]: - raise Exception( - f"❌ '{upscale_method}' gives poor results.\n\n" - f"👉 Go to the Scale SDXL Image MXD node and switch to another like 'lanczos'.\n\n" - f"Node may be hidden behind KSampler." - ) - - b, h, w, c = image.shape - - # Skip scaling if the image already matches an SDXL-safe resolution - if (w, h) in self.SDXL_SAFE_RESOLUTIONS or (h, w) in self.SDXL_SAFE_RESOLUTIONS: - return (image,) - - # ComfyUI-native megapixel math - samples = image.movedim(-1, 1) - orig_h, orig_w = samples.shape[2], samples.shape[3] - - target_pixels = int(round(total_megapixels * 1024 * 1024)) - scale_by = math.sqrt(target_pixels / (orig_w * orig_h)) - - new_w = max(1, round(orig_w * scale_by)) - new_h = max(1, round(orig_h * scale_by)) - - scaled = comfy.utils.common_upscale(samples, new_w, new_h, upscale_method, "disabled") - scaled = scaled.movedim(1, -1) - return (scaled,) - -######################################################################################################################## -# Flux Image Scale To Total Pixels (Flux Safe) -class FluxImageScaleToTotalPixelsSafe: - DESCRIPTION = """Scale to a target megapixel count and keep aspect ratio. Skips Flux-safe sizes.""" - upscale_methods = ["bilinear", "bicubic", "lanczos", "nearest-exact", "area"] - - # Flux-safe resolutions (width, height) – stored in one orientation only - FLUX_SAFE_RESOLUTIONS = [ - (1408, 1408), - (1728, 1152), - (1664, 1216), - (1920, 1088), - (2176, 960), - (1024, 1024), - (1216, 832), - (1152, 896), - (1344, 768), - (1536, 640), - (320, 320), - (384, 256), - (448, 320), - (448, 256), - (576, 256), - ] - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "upscale_method": (cls.upscale_methods, {"default": "bilinear"}), - "total_megapixels": ( - "FLOAT", - { - "default": 1.0, - "min": 0.01, - "max": 128.0, - "step": 0.01, - "tooltip": "Set the total megapixels (e.g., 1.0 = 1 MP)", - }, - ), - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "upscale" - CATEGORY = "MXD/Upscaling" - - def upscale(self, image, upscale_method, total_megapixels): - if upscale_method in ["nearest-exact", "area"]: - raise Exception( - f"❌ '{upscale_method}' gives poor results.\n\n" - f"👉 Go to the Scale Flux Image MXD node and switch to another like 'lanczos'.\n\n" - f"Node may be hidden behind KSampler." - ) - - b, h, w, c = image.shape - - # Skip scaling if image matches any Flux-safe resolution - if (w, h) in self.FLUX_SAFE_RESOLUTIONS or (h, w) in self.FLUX_SAFE_RESOLUTIONS: - return (image,) - - samples = image.movedim(-1, 1) - orig_h, orig_w = samples.shape[2], samples.shape[3] - - target_pixels = int(round(total_megapixels * 1024 * 1024)) - scale_by = math.sqrt(target_pixels / (orig_w * orig_h)) - - new_w = max(1, round(orig_w * scale_by)) - new_h = max(1, round(orig_h * scale_by)) - - scaled = comfy.utils.common_upscale(samples, new_w, new_h, upscale_method, "disabled") - scaled = scaled.movedim(1, -1) - return (scaled,) - -######################################################################################################################## -# Prompt with Guidance (Flux) -class PromptWithGuidance(ComfyNodeABC): - DESCRIPTION = """Encode text and apply Flux guidance in one node.""" - @classmethod - def INPUT_TYPES(cls) -> InputTypeDict: - return { - "required": { - "text": (IO.STRING, {"multiline": True, "dynamicPrompts": True}), - "clip": (IO.CLIP, {"tooltip": "The CLIP model used for encoding the text."}), - "guidance": ("FLOAT", {"default": 3.5, "min": 0.0, "max": 100.0, "step": 0.1}) - } - } - - RETURN_TYPES = (IO.CONDITIONING,) - FUNCTION = "encode_and_guide" - CATEGORY = "MXD/conditioning" - - def encode_and_guide(self, text, clip, guidance): - if clip is None: - raise RuntimeError("CLIP model is None. Your checkpoint may not contain a text encoder.") - - tokens = clip.tokenize(text) - conditioning = clip.encode_from_tokens_scheduled(tokens) - conditioning = node_helpers.conditioning_set_values(conditioning, {"guidance": guidance}) - return (conditioning,) - -######################################################################################################################## -if HAVE_COMFY_API: - class QwenImageEditSingleMXD(io.ComfyNode): - @classmethod - def define_schema(cls): - return io.Schema( - node_id="QwenImageEditSingleMXD", - display_name="Qwen Image Edit + Latent MXD", - category="MXD/conditioning", - description="Encode prompt/image and output a matching empty latent.", - inputs=[ - io.Clip.Input("clip"), - io.String.Input("prompt", multiline=True, dynamic_prompts=True), - io.Vae.Input("vae", optional=True), - io.Image.Input("image", optional=True), - io.Int.Input("batch_size", default=1, min=1, max=4096), - ], - outputs=[ - io.Conditioning.Output(), - io.Latent.Output(), # New Output - ], - ) - - @classmethod - def execute(cls, clip, prompt, vae=None, image=None, batch_size=1) -> io.NodeOutput: - ref_latents = [] - images_vl = [] - llama_template = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" - image_prompt = "" - - # Default fallback size if no image is provided (1024x1024) - final_width, final_height = 1024, 1024 - - if image is not None: - samples = image.movedim(-1, 1) - - # --- VISION SCALING (384px area) --- - total_vl = int(384 * 384) - scale_vl = math.sqrt(total_vl / (samples.shape[3] * samples.shape[2])) - width_vl = round(samples.shape[3] * scale_vl) - height_vl = round(samples.shape[2] * scale_vl) - - s_vl = comfy.utils.common_upscale(samples, width_vl, height_vl, "area", "disabled") - images_vl.append(s_vl.movedim(1, -1)) - - # --- LATENT/VAE SCALING (1024px area) --- - total_lat = int(1024 * 1024) - scale_lat = math.sqrt(total_lat / (samples.shape[3] * samples.shape[2])) - # Calculate final dimensions to be multiples of 8 - final_width = round(samples.shape[3] * scale_lat / 8.0) * 8 - final_height = round(samples.shape[2] * scale_lat / 8.0) * 8 - - if vae is not None: - s_lat = comfy.utils.common_upscale(samples, final_width, final_height, "area", "disabled") - ref_latents.append(vae.encode(s_lat.movedim(1, -1)[:, :, :, :3])) - - image_prompt += "Picture 1: <|vision_start|><|image_pad|><|vision_end|>" - - # 1. Generate the Empty Latent (SD3 Style: 16 channels, 1/8th resolution) - # This replaces the need for the separate EmptySD3LatentImage node - latent_tensor = torch.zeros( - [batch_size, 16, final_height // 8, final_width // 8], - device=comfy.model_management.intermediate_device() - ) - latent_output = {"samples": latent_tensor} - - # 2. Process Conditioning - tokens = clip.tokenize(image_prompt + prompt, images=images_vl, llama_template=llama_template) - conditioning = clip.encode_from_tokens_scheduled(tokens) - - if len(ref_latents) > 0: - conditioning = node_helpers.conditioning_set_values( - conditioning, - {"reference_latents": ref_latents}, - append=True, - ) - - return io.NodeOutput(conditioning, latent_output) - - ######################################################################################################################## - class QwenImageEditTripleMXD(io.ComfyNode): - @classmethod - def define_schema(cls): - return io.Schema( - node_id="QwenImageEditTripleMXD", - display_name="Qwen Image Edit Prompt MXD (Triple)", - category="advanced/conditioning", - inputs=[ - io.Clip.Input("clip"), - io.String.Input("prompt", multiline=True, dynamic_prompts=True), - io.Vae.Input("vae", optional=True), - io.Image.Input("image1", optional=True), - io.Image.Input("image2", optional=True), - io.Image.Input("image3", optional=True), - io.Int.Input("batch_size", default=1, min=1, max=4096), - ], - outputs=[ - io.Conditioning.Output(), - io.Latent.Output(), - ], - ) - - @classmethod - def execute(cls, clip, prompt, vae=None, image1=None, image2=None, image3=None, batch_size=1) -> io.NodeOutput: - ref_latents = [] - images = [image1, image2, image3] - images_vl = [] - llama_template = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" - image_prompt = "" - - # Default fallback - latent_width = 1024 - latent_height = 1024 - - for i, image in enumerate(images): - if image is not None: - samples = image.movedim(-1, 1) - - # 1. VL Model Scaling (LLM Vision) - total_vl = int(384 * 384) - scale_by_vl = math.sqrt(total_vl / (samples.shape[3] * samples.shape[2])) - width_vl = round(samples.shape[3] * scale_by_vl) - height_vl = round(samples.shape[2] * scale_by_vl) - s_vl = comfy.utils.common_upscale(samples, width_vl, height_vl, "area", "disabled") - images_vl.append(s_vl.movedim(1, -1)) - - # 2. VAE Scaling (Synchronized to 16-step for SD3 compatibility) - if vae is not None: - total_ref = int(1024 * 1024) - scale_by_ref = math.sqrt(total_ref / (samples.shape[3] * samples.shape[2])) - - # Pixels as multiple of 16 ensures Latent (Pixels/8) is always even - width_ref = round(samples.shape[3] * scale_by_ref / 16.0) * 16 - height_ref = round(samples.shape[2] * scale_by_ref / 16.0) * 16 - - if i == 0: - latent_width = width_ref - latent_height = height_ref - - s_ref = comfy.utils.common_upscale(samples, width_ref, height_ref, "area", "disabled") - ref_latents.append(vae.encode(s_ref.movedim(1, -1)[:, :, :, :3])) - - image_prompt += "Picture {}: <|vision_start|><|image_pad|><|vision_end|>".format(i + 1) - - # Process tokens and conditioning - tokens = clip.tokenize(image_prompt + prompt, images=images_vl, llama_template=llama_template) - conditioning = clip.encode_from_tokens_scheduled(tokens) - - if len(ref_latents) > 0: - conditioning = node_helpers.conditioning_set_values(conditioning, {"reference_latents": ref_latents}, append=True) - - # Create Output Latent - latent = torch.zeros([batch_size, 16, latent_height // 8, latent_width // 8], device=comfy.model_management.intermediate_device()) - - # FIXED: Return outputs positionally to match the schema defined above - # Output 1: Conditioning, Output 2: Latent Dictionary - return io.NodeOutput(conditioning, {"samples": latent}) - -######################################################################################################################## -class FluxResolutionMatcher: - DESCRIPTION = """Match the closest Flux resolution and orientation for the input image.""" - CATEGORY = "MXD/Latent" - FUNCTION = "match_resolution" - RETURN_NAMES = ("resolution", "vertical") - - # Full set kept for compatibility (enum list must match FluxEmptyLatentImage) - RESOLUTIONS = { - "— High Resolutions —": None, - "Square (1:1) 1408x1408": (1408, 1408), - "Standard (4:3) 1664x1216": (1664, 1216), - "Landscape (3:2) 1728x1152": (1728, 1152), - "Widescreen (16:9) 1920x1088": (1920, 1088), - "Ultrawide (21:9) 2176x960": (2176, 960), - - "— Standard Resolutions —": None, - "Square (1:1) 1024x1024": (1024, 1024), - "Standard (4:3) 1152x896": (1152, 896), - "Landscape (3:2) 1216x832": (1216, 832), - "Widescreen (16:9) 1344x768": (1344, 768), - "Ultrawide (21:9) 1536x640": (1536, 640), - - "— Low Resolutions —": None, - "Square (1:1) 320x320": (320, 320), - "Standard (4:3) 448x320": (448, 320), - "Landscape (3:2) 384x256": (384, 256), - "Widescreen (16:9) 448x256": (448, 256), - "Ultrawide (21:9) 576x256": (576, 256), - } - - # Keep same enum type so it connects to FluxEmptyLatentImage - RETURN_TYPES = (list(RESOLUTIONS.keys()), "BOOLEAN") - - # Precompute aspect ratio groups (only for standard resolutions) - ASPECT_RATIO_GROUPS = {} - for res_str, dims in RESOLUTIONS.items(): - if dims is None: - continue - # ✅ Skip high and low groups for logic - if "High" in res_str or "Low" in res_str: - continue - group_name = " ".join(res_str.split(' ')[:-1]) - if group_name not in ASPECT_RATIO_GROUPS: - w, h = dims - ratio = w / h - ASPECT_RATIO_GROUPS[group_name] = {'ratio': ratio, 'resolutions': []} - ASPECT_RATIO_GROUPS[group_name]['resolutions'].append(res_str) - - @classmethod - def INPUT_TYPES(cls): - return {"required": {"image": ("IMAGE",)}} - - def match_resolution(self, image: torch.Tensor): - if image.dim() < 4 or image.shape[1] < 1 or image.shape[2] < 1: - print("Warning: Invalid image tensor received. Falling back to default resolution.") - return ("Square (1:1) 1024x1024", False) - - _batch, height, width, _channels = image.shape - is_vertical = height > width - img_aspect_ratio = (height / width) if is_vertical else (width / height) - img_area = height * width - - best_ar_group_name = min( - self.ASPECT_RATIO_GROUPS.keys(), - key=lambda name: abs(img_aspect_ratio - self.ASPECT_RATIO_GROUPS[name]['ratio']) - ) - - candidate_res_strings = self.ASPECT_RATIO_GROUPS[best_ar_group_name]['resolutions'] - - best_res_string = min( - candidate_res_strings, - key=lambda res_str: abs(img_area - (self.RESOLUTIONS[res_str][0] * self.RESOLUTIONS[res_str][1])) - ) - - return (best_res_string, is_vertical) -######################################################################################################################## - -class SDXLResolutionMatcher: - DESCRIPTION = """Match the closest SDXL resolution and orientation for the input image.""" - CATEGORY = "MXD/Latent" - FUNCTION = "match_resolution" - RETURN_NAMES = ("resolution", "vertical") - - # Use the exact same enum list as SdxlEmptyLatentImage - RESOLUTIONS = SdxlEmptyLatentImage.RESOLUTIONS - - RETURN_TYPES = (list(RESOLUTIONS.keys()), "BOOLEAN") - - ASPECT_RATIO_GROUPS = {} - for res_str, dims in RESOLUTIONS.items(): - if dims is None: - continue - group_name = " ".join(res_str.split(" ")[:-1]) - if group_name not in ASPECT_RATIO_GROUPS: - w, h = dims - ratio = w / h - ASPECT_RATIO_GROUPS[group_name] = {"ratio": ratio, "resolutions": []} - ASPECT_RATIO_GROUPS[group_name]["resolutions"].append(res_str) - - @classmethod - def INPUT_TYPES(cls): - return {"required": {"image": ("IMAGE",)}} - - def match_resolution(self, image: torch.Tensor): - if image.dim() < 4 or image.shape[1] < 1 or image.shape[2] < 1: - print("Warning: Invalid image tensor received. Falling back to default resolution.") - return ("Square (1:1) 1024x1024", False) - - _batch, height, width, _channels = image.shape - is_vertical = height > width - img_aspect_ratio = (height / width) if is_vertical else (width / height) - img_area = height * width - - best_ar_group_name = min( - self.ASPECT_RATIO_GROUPS.keys(), - key=lambda name: abs(img_aspect_ratio - self.ASPECT_RATIO_GROUPS[name]["ratio"]) - ) - - candidate_res_strings = self.ASPECT_RATIO_GROUPS[best_ar_group_name]["resolutions"] - - best_res_string = min( - candidate_res_strings, - key=lambda res_str: abs(img_area - (self.RESOLUTIONS[res_str][0] * self.RESOLUTIONS[res_str][1])) - ) - - return (best_res_string, is_vertical) -######################################################################################################################## - -class LatentHalfMasks: - DESCRIPTION = """Split a latent into left and right half masks.""" - TITLE = "Latent Half Masks" - CATEGORY = "MXD/Latent" - - RETURN_TYPES = ("MASK", "MASK") - RETURN_NAMES = ("mask_left", "mask_right") - OUTPUT_TOOLTIPS = ( - "Mask covering the left half of the latent.", - "Mask covering the right half of the latent.", - ) - FUNCTION = "make_masks" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "latent": ("LATENT",), - } - } - - RETURN_TYPES = ("MASK", "MASK") - RETURN_NAMES = ("mask_left", "mask_right") - FUNCTION = "make_masks" - CATEGORY = "MXD/latent" - - def make_masks(self, latent): - # Infer width/height from latent (assumes 8x scale) - samples = latent.get("samples", None) - if samples is None or not isinstance(samples, torch.Tensor): - raise ValueError("LatentHalfMasks: invalid latent or missing 'samples' tensor.") - h_lat, w_lat = samples.shape[-2], samples.shape[-1] - w, h = int(w_lat * 8), int(h_lat * 8) - - # Always vertical, center split, no feather, no swap - split_px = w // 2 - left = torch.zeros((h, w), dtype=torch.float32) - right = torch.zeros((h, w), dtype=torch.float32) - left[:, :split_px] = 1.0 - right[:, split_px:] = 1.0 - - return left, right - -######################################################################################################################## - -# Get Latent Size -class GetLatentSizeMXD: - DESCRIPTION = """Get image width/height from a latent.""" - TITLE = "Get Latent Size" - CATEGORY = "MXD/Latent" - - RETURN_TYPES = ("INT", "INT") - RETURN_NAMES = ("width", "height") - OUTPUT_TOOLTIPS = ("Latent-derived image width in pixels.", "Latent-derived image height in pixels.") - FUNCTION = "get_size" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "latent": ("LATENT",), - } - } - - def get_size(self, latent): - if isinstance(latent, dict): - width = latent.get("width") - height = latent.get("height") - if width is not None and height is not None: - try: - return (int(width), int(height)) - except Exception: - pass - - samples = latent.get("samples") - else: - samples = None - - if samples is None or not isinstance(samples, torch.Tensor): - raise ValueError("GetLatentSizeMXD: invalid latent or missing 'samples' tensor.") - - channels = samples.shape[1] if samples.dim() >= 2 else 0 - scale = 16 if channels >= 64 else 8 - - h_lat, w_lat = samples.shape[-2], samples.shape[-1] - return (int(w_lat * scale), int(h_lat * scale)) - -######################################################################################################################## - -# --- Helper function to find the bounding box of a mask --- -def get_bounding_box(mask_tensor): - """ - Finds the bounding box of a non-zero region in a mask tensor. - The mask is expected to be a 2D tensor (H, W). - Returns a tuple (x_min, y_min, x_max, y_max) or None if the mask is empty. - """ - # Get non-zero coordinates from the mask - non_zero_coords = torch.nonzero(mask_tensor, as_tuple=False) - - # If the mask is empty, there is no bounding box - if non_zero_coords.numel() == 0: - return None - - # Find the min and max coordinates for y (dim 0) and x (dim 1) - min_y = non_zero_coords[:, 0].min().item() - max_y = non_zero_coords[:, 0].max().item() - min_x = non_zero_coords[:, 1].min().item() - max_x = non_zero_coords[:, 1].max().item() - - # The bounding box for PIL needs (left, upper, right, lower). - # We add +1 to the max values because the upper bound is exclusive. - return (min_x, min_y, max_x + 1, max_y + 1) - -# --- Tensor to PIL and PIL to Tensor conversion helpers --- -def tensor_to_pil(tensor): - """Converts a torch tensor (B, H, W, C) to a list of PIL Images.""" - if tensor is None: - return [] - - # Handle different tensor dimensions - if tensor.dim() == 4: # Batch of images - images = [] - for i in range(tensor.shape[0]): - img_np = 255. * tensor[i].cpu().numpy() - images.append(Image.fromarray(np.clip(img_np, 0, 255).astype(np.uint8))) - return images - elif tensor.dim() == 3: # Single image - img_np = 255. * tensor.cpu().numpy() - return [Image.fromarray(np.clip(img_np, 0, 255).astype(np.uint8))] - else: - raise ValueError(f"Unsupported tensor dimension: {tensor.dim()}") - -def pil_to_tensor(pil_images): - """Converts a list of PIL Images back to a torch tensor (B, H, W, C).""" - if not isinstance(pil_images, list): - pil_images = [pil_images] - - tensors = [] - for img in pil_images: - # Convert to RGB, then to a numpy array, normalize, and create a tensor - img_np = np.array(img.convert("RGB")).astype(np.float32) / 255.0 - tensors.append(torch.from_numpy(img_np).unsqueeze(0)) - - # Stack all tensors into a single batch tensor - return torch.cat(tensors, dim=0) - -# -------------------------------------------------------------------- -# ✨ The Main Node Class ✨ -# -------------------------------------------------------------------- -class PlaceImageByMask: - Description = """Place an overlay image inside the mask bounds on a base image.""" - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "base_image": ("IMAGE",), - "mask": ("MASK",), - "overlay_image": ("IMAGE",), - }, - "optional": { - "maintain_aspect_ratio": ("BOOLEAN", {"default": True}), - } - } - - RETURN_TYPES = ("IMAGE",) - FUNCTION = "place_image" - CATEGORY = "MXD/Image" - - def place_image(self, base_image, overlay_image, mask, maintain_aspect_ratio=True): - # Convert input tensors to lists of PIL Images - base_pils = tensor_to_pil(base_image) - overlay_pils = tensor_to_pil(overlay_image) - - processed_images = [] - - # Process each image in the batch - for i, base_pil in enumerate(base_pils): - # Work with an RGBA version of the base image for clean pasting - composited_image = base_pil.convert("RGBA") - - # Select the corresponding overlay and mask for the current base image - # Clamping the index prevents errors if batch sizes are mismatched - overlay_pil = overlay_pils[min(i, len(overlay_pils) - 1)].convert("RGBA") - current_mask = mask[min(i, mask.shape[0] - 1)] - - # Find the bounding box from the mask - bbox = get_bounding_box(current_mask) - - # If no mask is found, just use the original base image and skip to the next - if not bbox: - raise ValueError("The base image must be masked where you want the overlay to appear.") - - x_min, y_min, x_max, y_max = bbox - box_width = x_max - x_min - box_height = y_max - y_min - - # If the bounding box has no area, skip to the next image - if box_width <= 0 or box_height <= 0: - processed_images.append(base_pil) - continue - - # --- Resize the overlay image using the specified method --- - if maintain_aspect_ratio: - # Resize to fit *within* the box, preserving aspect ratio (like a thumbnail) - resized_overlay = overlay_pil.copy() - resized_overlay.thumbnail((box_width, box_height), Image.Resampling.LANCZOS) - - # Calculate position to center the resized overlay within the bounding box - paste_x = x_min + (box_width - resized_overlay.width) // 2 - paste_y = y_min + (box_height - resized_overlay.height) // 2 - paste_pos = (paste_x, paste_y) - else: - # As originally requested: stretch to fill the bounding box exactly - resized_overlay = overlay_pil.resize((box_width, box_height), resample=Image.Resampling.LANCZOS) - paste_pos = (x_min, y_min) - - # --- Paste the resized overlay onto the base image --- - # The alpha channel of the overlay itself is used as the mask for pasting. - # This ensures transparent areas of the overlay are handled correctly. - composited_image.paste(resized_overlay, paste_pos, resized_overlay) - - processed_images.append(composited_image) - - # Convert the list of processed PIL images back to a single batch tensor for output - output_tensor = pil_to_tensor(processed_images) - return (output_tensor,) - -###################################################################################################################################### - -class CropImageByMask: - DESCRIPTION = """Crop images to the mask bounds when a mask is provided.""" - @classmethod - def INPUT_TYPES(s): - return { - "required": { - "image": ("IMAGE", ), - }, - "optional": { - "mask": ("MASK", ), - } - } - - RETURN_TYPES = ("IMAGE", ) - RETURN_NAMES = ("image", ) - FUNCTION = "crop" - CATEGORY = "MXD/image" - - def crop(self, image, mask=None): - # If no mask is provided or the mask is completely empty, return the original image - if mask is None or not torch.any(mask > 0): - return (image, ) - - B, H, W, C = image.shape - mask = mask.round() - - # Find bounding box for each batch - crops = [] - - for b in range(B): - current_mask = mask[min(b, mask.shape[0]-1)] - - # Check if the mask for this specific image is empty. - if not torch.any(current_mask > 0): - # If a specific mask in a batch is empty, we can't crop. - # To prevent errors with torch.cat later due to different sizes, - # we'll skip cropping for the whole batch and return the original. - # This ensures the output is always a valid tensor. - print("Warning: An empty mask was found in a batch. Returning original images.") - return (image, ) - - # Get coordinates of non-zero elements - rows = torch.any(current_mask > 0, dim=1) - cols = torch.any(current_mask > 0, dim=0) - - # Find boundaries - y_min, y_max = torch.where(rows)[0][[0, -1]] - x_min, x_max = torch.where(cols)[0][[0, -1]] - - # Crop image - crop = image[b:b+1, y_min:y_max+1, x_min:x_max+1, :] - crops.append(crop) - - # Note: This will raise an error if the crops have different sizes. - # The original code had this limitation. - cropped_images = torch.cat(crops, dim=0) - - return (cropped_images, ) - -######################################################################################################################## -# ---------- Helpers (copied from latent loader style) ---------- -def _safe_json_loads(s): - if s is None: - return None - if isinstance(s, bytes): - try: - s = s.decode("utf-8", "ignore") - except Exception: - return None - if not isinstance(s, str): - return None - try: - return json.loads(s) - except Exception: - try: - return json.loads(json.loads(s)) - except Exception: - return None - - -def _extract_params_from_prompt_json(prompt_json: dict): - """ - Returns (positive, negative) from saved Comfy prompt graph. - """ - pos = "" - neg = "" - if not isinstance(prompt_json, dict): - return pos, neg - - # unwrap if saved as {"prompt": {...}} - graph = prompt_json.get("prompt", prompt_json) - if not isinstance(graph, dict): - return pos, neg - - # try to find KSampler/KSamplerAdvanced node - ks = None - for _, v in graph.items(): - if "KSampler" in v.get("class_type", ""): - ks = v - break - if not ks: - return pos, neg - - kin = ks.get("inputs", {}) - - def _as_node_id(x): - return str(x[0]) if isinstance(x, (list, tuple)) and x else None - - def _text_from_clip(node_id): - n = graph.get(str(node_id), {}) - if n.get("class_type") == "CLIPTextEncode": - return str(n.get("inputs", {}).get("text", "")).strip() - return "" - - pos = _text_from_clip(_as_node_id(kin.get("positive"))) - neg = _text_from_clip(_as_node_id(kin.get("negative"))) - - return pos, neg - -def _strip_counter(name: str) -> str: - # Only strip the trailing pattern we generate when saving: "_<5digits>_" - # Preserve numeric-only base names like "96". - stem, _ = os.path.splitext(name) - m = re.match(r"^(.*?)(?:_\d{5}_)$", stem) - return m.group(1) if m else stem - -# ---------- Node ---------- -def _indent_paths(paths): - indented = [] - for path in paths: - if not path: - indented.append("") - continue - clean_path = path.lstrip("\u00a0 ") - depth = clean_path.count("/") - indent = "\u00a0" * (depth * 4) - indented.append(indent + clean_path) - return indented - - -def _scan_subdir_mtimes(root: str, subdirs: set, branch_latest: dict, exts: tuple = None): - """ - Walk `root`, adding every subfolder's relative path to `subdirs` and - bubbling the mtime of its most recently modified file up to every - ancestor branch (including "" for the root) in `branch_latest`. - - When `exts` is given, a folder (and its ancestors) is only added if it - directly or recursively contains at least one file matching `exts` -- - so folders with no relevant content don't show up as pickable at all. - """ - try: - for dirpath, dirnames, filenames in os.walk(root): - # Exclude hidden folders (e.g. .git, .github) and __pycache__ - dirnames[:] = [d for d in dirnames if not d.startswith('.') and d != '__pycache__'] - rel_path = os.path.relpath(dirpath, root) - rel_path = "" if rel_path == "." else rel_path.replace(os.path.sep, "/") - - latest = 0.0 - has_match = exts is None - for f in filenames: - if exts and not f.lower().endswith(exts): - continue - has_match = True - try: - m = os.path.getmtime(os.path.join(dirpath, f)) - except OSError: - continue - if m > latest: - latest = m - - if not has_match: - continue - - if rel_path: - subdirs.add(rel_path) - - parts = [p for p in rel_path.split("/") if p] - for i in range(len(parts) + 1): - branch = "/".join(parts[:i]) - if latest > branch_latest.get(branch, -1.0): - branch_latest[branch] = latest - if i > 0: - subdirs.add(branch) - except OSError: - pass - - -def _list_image_batch_subdirs(root: str, exts: tuple = None): - """ - Recursive subfolders under `root`, newest first. Each folder is ordered by - the mtime of the most recently modified file anywhere inside it (so a - folder that just received a new file jumps back to the top). '' = the - root itself, always first. If `exts` is given, only folders that - directly or recursively contain a matching file are included. - """ - subdirs = set() - branch_latest = {} - _scan_subdir_mtimes(root, subdirs, branch_latest, exts) - ordered = sorted(subdirs, key=lambda d: (-branch_latest.get(d, -1.0), d.lower())) - return [""] + ordered - - -def _list_image_batch_subdirs_union(output_root: str, input_root: str, exts: tuple = None): - """ - Union of recursive subfolders from both roots, newest first. A folder - present under both roots is ranked by whichever side has the more - recent file, so it doesn't matter which source the user has selected. - """ - subdirs = set() - branch_latest = {} - _scan_subdir_mtimes(output_root, subdirs, branch_latest, exts) - _scan_subdir_mtimes(input_root, subdirs, branch_latest, exts) - ordered = sorted(subdirs, key=lambda d: (-branch_latest.get(d, -1.0), d.lower())) - return [""] + ordered - - -IMAGE_BATCH_EXTS = (".png", ".jpg", ".jpeg", ".webp") -VIDEO_BATCH_EXTS = (".mp4",) - - -def _sort_paths_newest_first(paths): - """Sort file paths by mtime desc (newest first), stable by normalized path.""" - def _mtime(path): - try: - return os.path.getmtime(path) - except OSError: - return 0.0 - - return sorted(paths, key=lambda p: (-_mtime(p), p.replace("\\", "/").lower())) - - -def _list_files_recursive(root: str, exts: tuple): - """Recursively list files under `root` matching `exts`, newest first, as relpaths.""" - try: - files = [] - for dirpath, dirnames, filenames in os.walk(root): - dirnames[:] = [d for d in dirnames if not d.startswith('.') and d != '__pycache__'] - for f in filenames: - if f.lower().endswith(exts): - files.append(os.path.join(dirpath, f)) - files = _sort_paths_newest_first(files) - return [os.path.relpath(f, root).replace(os.sep, "/") for f in files] - except OSError: - return [] - - -def _list_files_recursive_union(output_root: str, input_root: str, exts: tuple): - """ - Union of recursive files from both roots, newest first. A relative path - present under both roots is ranked by whichever side's file is more - recent, so it doesn't matter which source the user has selected. - """ - mtimes = {} - - def scan(root): - try: - for dirpath, dirnames, filenames in os.walk(root): - dirnames[:] = [d for d in dirnames if not d.startswith('.') and d != '__pycache__'] - for f in filenames: - if not f.lower().endswith(exts): - continue - full = os.path.join(dirpath, f) - rel = os.path.relpath(full, root).replace(os.sep, "/") - try: - m = os.path.getmtime(full) - except OSError: - m = 0.0 - if m > mtimes.get(rel, -1.0): - mtimes[rel] = m - except OSError: - pass - - scan(output_root) - scan(input_root) - return sorted(mtimes, key=lambda p: (-mtimes[p], p.lower())) or [""] - - -# Server routes so the frontend can swap folder/file dropdowns between -# inputs/outputs without reloading the page. -try: - from server import PromptServer as _MXD_PromptServer - from aiohttp import web as _mxd_web - - @_MXD_PromptServer.instance.routes.get("/mxd/image_batch/folders") - async def _mxd_list_image_batch_folders(request): - return _mxd_web.json_response({ - "outputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_output_directory())), - "inputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_input_directory())), - }) - - @_MXD_PromptServer.instance.routes.get("/mxd/video_batch/folders") - async def _mxd_list_video_batch_folders(request): - return _mxd_web.json_response({ - "outputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_output_directory(), VIDEO_BATCH_EXTS)), - "inputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_input_directory(), VIDEO_BATCH_EXTS)), - }) - - @_MXD_PromptServer.instance.routes.get("/mxd/single_loader/files") - async def _mxd_list_single_loader_files(request): - kind = request.query.get("kind", "image") - exts = VIDEO_BATCH_EXTS if kind == "video" else IMAGE_BATCH_EXTS - return _mxd_web.json_response({ - "outputs": _list_files_recursive(folder_paths.get_output_directory(), exts), - "inputs": _list_files_recursive(folder_paths.get_input_directory(), exts), - }) -except Exception as _e: - print(f"[LoadImageBatchMXD] Could not register folders route: {_e}") - - -class LoadImageBatchMXD: - DESCRIPTION = """Load images from an inputs or outputs folder, make masks from alpha, and read prompts.""" - TITLE = "Load Image Batch (Inputs/Outputs + Prompts)" - CATEGORY = "MXD/Image" - - RETURN_TYPES = ("IMAGE", "MASK", "STRING", "STRING") - RETURN_NAMES = ("IMAGE", "MASK", "positive", "negative") - OUTPUT_IS_LIST = (True, True, True, True) - FUNCTION = "load_batch" - - @classmethod - def INPUT_TYPES(cls): - # Provide the union of inputs + outputs subfolders so any saved value - # validates regardless of which source it belongs to. The frontend - # filters the visible list down to the selected source on the fly. - union = _indent_paths(_list_image_batch_subdirs_union( - folder_paths.get_output_directory(), folder_paths.get_input_directory() - )) - return { - "required": { - "source": (("outputs", "inputs"), {"default": "outputs"}), - "folder": (tuple(union), {"default": ""}), - } - } - - def _extract_prompts(self, image: Image.Image): - pos, neg = "", "" - try: - raw = image.info.get("prompt") - if raw: - prompt_json = _safe_json_loads(raw) - if prompt_json: - pos, neg = _extract_params_from_prompt_json(prompt_json) - else: - pos = raw - except Exception as e: - print(f"[LoadImageBatchMXD] Prompt parse failed: {e}") - return pos, neg - - def load_batch(self, folder: str, source: str = "outputs"): - folder = folder.lstrip("\u00a0 ") - root = ( - folder_paths.get_input_directory() - if source == "inputs" - else folder_paths.get_output_directory() - ) - folder_path = os.path.normpath(os.path.join(root, folder)) if folder else root - - if not os.path.isdir(folder_path): - raise FileNotFoundError(f"No such folder: {folder_path}") - - valid_exts = IMAGE_BATCH_EXTS - - # Recursively find all matching files - files = [] - for dirpath, dirnames, filenames in os.walk(folder_path): - dirnames.sort() - for f in sorted(filenames): - if f.lower().endswith(valid_exts): - files.append(os.path.join(dirpath, f)) - - if not files: - raise FileNotFoundError(f"No valid images found in folder '{folder_path}' (including subfolders)") - - images, masks, positives, negatives, prefixes = [], [], [], [], [] - - for path in files: - i = Image.open(path) - i = ImageOps.exif_transpose(i) - - pos, neg = self._extract_prompts(i) - positives.append(pos) - negatives.append(neg) - - rgb = i.convert("RGB") - arr = np.array(rgb).astype(np.float32) / 255.0 - img_t = torch.from_numpy(arr)[None, ...] - - if 'A' in i.getbands(): - mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 - mask_t = 1.0 - torch.from_numpy(mask).unsqueeze(0) - else: - h, w = arr.shape[:2] - mask_t = torch.zeros((1, h, w), dtype=torch.float32) - - images.append(img_t) - masks.append(mask_t) - - return (images, masks, positives, negatives) - - -class LoadVideoBatchMXD: - DESCRIPTION = """Load videos from an inputs or outputs folder as a batch.""" - TITLE = "Load Video Batch (Inputs/Outputs)" - CATEGORY = "MXD/Video" - - RETURN_TYPES = ("VIDEO",) - RETURN_NAMES = ("VIDEO",) - OUTPUT_IS_LIST = (True,) - FUNCTION = "load_batch" - - @classmethod - def INPUT_TYPES(cls): - # Same union-of-sources pattern as LoadImageBatchMXD; reuses that - # node's folder listing helper, filtered to folders that actually - # contain a video so empty/irrelevant folders don't show up. - union = _indent_paths(_list_image_batch_subdirs_union( - folder_paths.get_output_directory(), folder_paths.get_input_directory(), VIDEO_BATCH_EXTS - )) - return { - "required": { - "source": (("outputs", "inputs"), {"default": "outputs"}), - "folder": (tuple(union), {"default": ""}), - } - } - - def load_batch(self, folder: str, source: str = "outputs"): - if not HAVE_COMFY_API_VIDEO: - raise RuntimeError( - "[LoadVideoBatchMXD] Video output requires a newer ComfyUI core with " - "comfy_api.latest / comfy_api.input_impl support. Please update ComfyUI." - ) - - folder = folder.lstrip("\u00a0 ") - root = ( - folder_paths.get_input_directory() - if source == "inputs" - else folder_paths.get_output_directory() - ) - folder_path = os.path.normpath(os.path.join(root, folder)) if folder else root - - if not os.path.isdir(folder_path): - raise FileNotFoundError(f"No such folder: {folder_path}") - - valid_exts = VIDEO_BATCH_EXTS - - # Recursively find all matching files - files = [] - for dirpath, dirnames, filenames in os.walk(folder_path): - dirnames.sort() - for f in sorted(filenames): - if f.lower().endswith(valid_exts): - files.append(os.path.join(dirpath, f)) - - if not files: - raise FileNotFoundError(f"No valid videos found in folder '{folder_path}' (including subfolders)") - - videos = [VideoFromFile(path) for path in files] - - return (videos,) - - -class LoadImageFromFolderMXD: - DESCRIPTION = ( - "Load a single image from any inputs/outputs subfolder. Turn on run_folder " - "to auto-queue every image in that same folder, one after another." - ) - TITLE = "Load Image (From Folder) MXD" - CATEGORY = "MXD/Image" - - RETURN_TYPES = ("IMAGE", "MASK", "STRING", "STRING", "STRING") - RETURN_NAMES = ("IMAGE", "MASK", "positive", "negative", "filename") - FUNCTION = "load_image" - - @classmethod - def INPUT_TYPES(cls): - # Union of both sources so any saved value validates regardless of which - # source it belongs to; the frontend narrows the visible list to the - # selected source on the fly (mirrors LoadImageBatchMXD's folder picker). - union = _list_files_recursive_union( - folder_paths.get_output_directory(), folder_paths.get_input_directory(), IMAGE_BATCH_EXTS - ) - return { - "required": { - "source": (("outputs", "inputs"), {"default": "outputs"}), - "image": (tuple(union), ), - "run_folder": ("BOOLEAN", { - "default": False, - "tooltip": "When enabled, hitting Queue Prompt auto-queues every image in this file's folder, one after another, instead of just the selected file.", - }), - } - } - - def _extract_prompts(self, image: Image.Image): - pos, neg = "", "" - try: - raw = image.info.get("prompt") - if raw: - prompt_json = _safe_json_loads(raw) - if prompt_json: - pos, neg = _extract_params_from_prompt_json(prompt_json) - else: - pos = raw - except Exception as e: - print(f"[LoadImageFromFolderMXD] Prompt parse failed: {e}") - return pos, neg - - def load_image(self, image: str, source: str = "outputs", run_folder: bool = False): - image = image.lstrip("  ") - root = ( - folder_paths.get_input_directory() - if source == "inputs" - else folder_paths.get_output_directory() - ) - path = os.path.normpath(os.path.join(root, image)) if image else None - - if not path or not os.path.isfile(path): - raise FileNotFoundError(f"No such image: {path}") - - i = Image.open(path) - i = ImageOps.exif_transpose(i) - - pos, neg = self._extract_prompts(i) - - rgb = i.convert("RGB") - arr = np.array(rgb).astype(np.float32) / 255.0 - img_t = torch.from_numpy(arr)[None, ...] - - if 'A' in i.getbands(): - mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 - mask_t = 1.0 - torch.from_numpy(mask).unsqueeze(0) - else: - h, w = arr.shape[:2] - mask_t = torch.zeros((1, h, w), dtype=torch.float32) - - return (img_t, mask_t, pos, neg, os.path.basename(path)) - - @classmethod - def IS_CHANGED(cls, image, source="outputs", run_folder=False): - root = ( - folder_paths.get_input_directory() - if source == "inputs" - else folder_paths.get_output_directory() - ) - path = os.path.normpath(os.path.join(root, image.lstrip("  "))) if image else None - if not path or not os.path.isfile(path): - return "" - m = hashlib.sha256() - with open(path, "rb") as f: - m.update(f.read()) - return m.digest().hex() - - -class LoadVideoFromFolderMXD: - DESCRIPTION = ( - "Load a single video from any inputs/outputs subfolder. Turn on run_folder " - "to auto-queue every video in that same folder, one after another." - ) - TITLE = "Load Video (From Folder) MXD" - CATEGORY = "MXD/Video" - - RETURN_TYPES = ("VIDEO", "STRING") - RETURN_NAMES = ("VIDEO", "filename") - FUNCTION = "load_video" - - @classmethod - def INPUT_TYPES(cls): - union = _list_files_recursive_union( - folder_paths.get_output_directory(), folder_paths.get_input_directory(), VIDEO_BATCH_EXTS - ) - return { - "required": { - "source": (("outputs", "inputs"), {"default": "outputs"}), - "video": (tuple(union), ), - "run_folder": ("BOOLEAN", { - "default": False, - "tooltip": "When enabled, hitting Queue Prompt auto-queues every video in this file's folder, one after another, instead of just the selected file.", - }), - } - } - - def load_video(self, video: str, source: str = "outputs", run_folder: bool = False): - if not HAVE_COMFY_API_VIDEO: - raise RuntimeError( - "[LoadVideoFromFolderMXD] Video output requires a newer ComfyUI core with " - "comfy_api.latest / comfy_api.input_impl support. Please update ComfyUI." - ) - - video = video.lstrip("  ") - root = ( - folder_paths.get_input_directory() - if source == "inputs" - else folder_paths.get_output_directory() - ) - path = os.path.normpath(os.path.join(root, video)) if video else None - - if not path or not os.path.isfile(path): - raise FileNotFoundError(f"No such video: {path}") - - return (VideoFromFile(path), os.path.basename(path)) - - @classmethod - def IS_CHANGED(cls, video, source="outputs", run_folder=False): - root = ( - folder_paths.get_input_directory() - if source == "inputs" - else folder_paths.get_output_directory() - ) - path = os.path.normpath(os.path.join(root, video.lstrip("  "))) if video else None - if not path or not os.path.isfile(path): - return "" - try: - return str(os.path.getmtime(path)) - except OSError: - return "" - - -class LoadImageWithPromptsMXD: - DESCRIPTION = """Load one input image, create a mask from alpha, and read prompts if present.""" - CATEGORY = "image" - - RETURN_TYPES = ("IMAGE", "MASK", "STRING", "STRING") - RETURN_NAMES = ("IMAGE", "MASK", "positive", "negative") - FUNCTION = "load_image" - - @classmethod - def INPUT_TYPES(s): - input_dir = folder_paths.get_input_directory() - files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))] - files = folder_paths.filter_files_content_types(files, ["image"]) - files = _sort_paths_newest_first([os.path.join(input_dir, f) for f in files]) - files = [os.path.basename(f) for f in files] - return {"required": {"image": (files, {"image_upload": True})}} - - def _extract_prompts(self, img: Image.Image): - pos, neg = "", "" - raw = img.info.get("prompt") - if raw: - prompt_json = _safe_json_loads(raw) - if prompt_json: - pos, neg = _extract_params_from_prompt_json(prompt_json) - else: - pos = raw - return pos, neg - - def load_image(self, image): - image_path = folder_paths.get_annotated_filepath(image) - img = node_helpers.pillow(Image.open, image_path) - - output_images, output_masks = [], [] - pos, neg = "", "" - w, h = None, None - - excluded_formats = ['MPO'] - - for i in ImageSequence.Iterator(img): - i = node_helpers.pillow(ImageOps.exif_transpose, i) - - if i.mode == 'I': - i = i.point(lambda i: i * (1 / 255)) - frame = i.convert("RGB") - - if len(output_images) == 0: - w, h = frame.size - # extract prompts only once (from first frame) - pos, neg = self._extract_prompts(i) - - if frame.size != (w, h): - continue - - arr = np.array(frame).astype(np.float32) / 255.0 - tensor_img = torch.from_numpy(arr)[None, ...] - - if 'A' in i.getbands(): - mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 - mask = 1. - torch.from_numpy(mask) - elif i.mode == 'P' and 'transparency' in i.info: - mask = np.array(i.convert('RGBA').getchannel('A')).astype(np.float32) / 255.0 - mask = 1. - torch.from_numpy(mask) - else: - mask = torch.zeros((1, 64, 64), dtype=torch.float32, device="cpu") - - output_images.append(tensor_img) - output_masks.append(mask.unsqueeze(0)) - - if len(output_images) > 1 and img.format not in excluded_formats: - output_image = torch.cat(output_images, dim=0) - output_mask = torch.cat(output_masks, dim=0) - else: - output_image = output_images[0] - output_mask = output_masks[0] - - return (output_image, output_mask, pos, neg) - - @classmethod - def IS_CHANGED(s, image): - image_path = folder_paths.get_annotated_filepath(image) - m = hashlib.sha256() - with open(image_path, 'rb') as f: - m.update(f.read()) - return m.digest().hex() - - @classmethod - def VALIDATE_INPUTS(s, image): - if not folder_paths.exists_annotated_filepath(image): - return f"Invalid image file: {image}" - return True - -######################################################################################################################## - -from nodes import PreviewImage, SaveImage -class SaveImage_MXD: - TITLE = "Save Image MXD" - CATEGORY = "MXD/Image" - OUTPUT_NODE = True - FUNCTION = "save" - - DESCRIPTION = """Save images to the output folder or preview them.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "images": ("IMAGE", {"tooltip": "Images to preview and/or save."}), - "filename_prefix": ("STRING", { - "default": "ComfyUI", - "tooltip": "File name prefix. Tip: you can use a subfolder like 'tests/my_run'." - }), - "mode": ([ - "Save + Preview", - "Save Only", - "Preview only" - ], { - "default": "Save + Preview", - "tooltip": "Choose whether to write files to disk, only preview, or save quietly." - }), - }, - "optional": { - "embed_workflow": ("BOOLEAN", { - "default": True, - "tooltip": "Embed workflow metadata when saving PNG previews/files." - }), - }, - "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, - } - - RETURN_TYPES = () - OUTPUT_TOOLTIPS = ("Saves and/or previews the images.",) - - @staticmethod - def _filtered_extra_pnginfo(extra_pnginfo, embed_workflow): - if embed_workflow or not isinstance(extra_pnginfo, dict): - return extra_pnginfo - filtered = {k: v for k, v in extra_pnginfo.items() if str(k).lower() != "workflow"} - return filtered or None - - def save(self, images, filename_prefix, mode, embed_workflow=True, prompt=None, extra_pnginfo=None): - if embed_workflow: - save_prompt = prompt - save_extra_pnginfo = self._filtered_extra_pnginfo(extra_pnginfo, True) - else: - # Core SaveImage embeds the hidden `prompt` graph too. - # Drop both to truly disable workflow reconstruction from saved files. - save_prompt = None - save_extra_pnginfo = None - - if mode.startswith("Preview"): - return PreviewImage().save_images(images, filename_prefix, save_prompt, save_extra_pnginfo) - result = SaveImage().save_images(images, filename_prefix, save_prompt, save_extra_pnginfo) - if mode == "Save Only" and isinstance(result, dict): - # Strip UI previews so nothing shows up in the ComfyUI viewer. - return {k: v for k, v in result.items() if k != "ui"} - return result - -######################################################################################################################## - -class ExtractWorkflowFromImageMXD: - TITLE = "Extract Workflow From Image MXD" - CATEGORY = "MXD/Image" - OUTPUT_NODE = True - FUNCTION = "extract_and_save" - - DESCRIPTION = """Save workflow metadata to a JSON file from a wired image execution context.""" - - def __init__(self): - self.output_dir = folder_paths.get_output_directory() - self.type = "output" - self.prefix_append = "" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE", {"tooltip": "Any connected image. Used to trigger extraction/save."}), - "filename_prefix": ("STRING", { - "default": "workflow/ComfyUI", - "tooltip": "Output JSON prefix. You can include subfolders, e.g. 'workflow/my_run'.", - }), - }, - "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, - } - - RETURN_TYPES = ("STRING",) - RETURN_NAMES = ("json_path",) - OUTPUT_TOOLTIPS = ("Relative path to the saved JSON file in outputs.",) - - @staticmethod - def _decode_json_candidate(value): - if value is None: - return None - - if isinstance(value, (dict, list)): - return value - - if isinstance(value, bytes): - for enc in ("utf-8", "utf-16", "latin-1"): - try: - value = value.decode(enc) - break - except Exception: - continue - if isinstance(value, bytes): - value = value.decode("utf-8", "ignore") - - if not isinstance(value, str): - return None - - raw = value.strip() - if not raw: - return None - - if raw.lower().startswith("workflow:"): - raw = raw.split(":", 1)[1].strip() - - parsed = _safe_json_loads(raw) - if isinstance(parsed, (dict, list)): - return parsed - return None - - def _extract_workflow_from_context(self, prompt=None, extra_pnginfo=None): - if isinstance(extra_pnginfo, dict): - for key in ("workflow", "Workflow"): - parsed = self._decode_json_candidate(extra_pnginfo.get(key)) - if parsed is not None: - return parsed - - parsed_extra = self._decode_json_candidate(extra_pnginfo) - if isinstance(parsed_extra, dict): - for key in ("workflow", "Workflow"): - parsed = self._decode_json_candidate(parsed_extra.get(key)) - if parsed is not None: - return parsed - - if prompt is not None: - parsed_prompt = self._decode_json_candidate(prompt) - if parsed_prompt is not None: - return {"prompt": parsed_prompt} - if isinstance(prompt, dict): - return {"prompt": prompt} - - return None - - def extract_and_save(self, image, filename_prefix="workflow/ComfyUI", prompt=None, extra_pnginfo=None): - workflow = self._extract_workflow_from_context(prompt, extra_pnginfo) - if workflow is None: - raise ValueError( - "No workflow metadata is available in this execution context. " - "Connect generated images from the current run, or ensure workflow metadata is present." - ) - - filename_prefix += self.prefix_append - height = image[0].shape[0] - width = image[0].shape[1] - full_output_folder, filename, counter, subfolder, _ = folder_paths.get_save_image_path( - filename_prefix, self.output_dir, width, height - ) - os.makedirs(full_output_folder, exist_ok=True) - - file = f"{filename}_{counter:05}_.json" - save_path = os.path.join(full_output_folder, file) - - with open(save_path, "w", encoding="utf-8", newline="\n") as f: - json.dump(workflow, f, ensure_ascii=False, indent=2) - - rel = os.path.join(subfolder, file) if subfolder else file - rel = rel.replace("\\", "/") - return { - "ui": {"text": [f"Saved workflow JSON: {rel}"]}, - "result": (rel,), - } - -######################################################################################################################## - -class SmartCropByMaskMXD: - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE", ), - "mask": ("MASK", ), - }, - } - - RETURN_TYPES = ("IMAGE", ) - RETURN_NAMES = ("image", ) - FUNCTION = "crop" - CATEGORY = "image/transform" - DESCRIPTION = "Slides a square crop window horizontally + vertically to center on subject mask." - - def crop(self, image, mask): - B, H, W, C = image.shape - mask = mask.round() - crops = [] - - for b in range(B): - mask_b = mask[min(b, mask.shape[0]-1)] - - # Get non-zero rows and columns - rows = torch.any(mask_b > 0, dim=1) - cols = torch.any(mask_b > 0, dim=0) - - # Default to center - center_x = W // 2 - center_y = H // 2 - - # Update center_x from mask if possible - if torch.any(cols): - x_min, x_max = torch.where(cols)[0][[0, -1]] - center_x = (x_min + x_max) // 2 - - # Update center_y from mask if possible - if torch.any(rows): - y_min, y_max = torch.where(rows)[0][[0, -1]] - center_y = (y_min + y_max) // 2 - - # Compute square crop box - side = min(H, W) - half = side // 2 - - left = max(0, center_x - half) - right = min(W, left + side) - left = right - side # clamp again - - top = max(0, center_y - half) - bottom = min(H, top + side) - top = bottom - side # clamp again - - # Final crop: safe slicing - crop = image[b:b+1, top:bottom, left:right, :] - crops.append(crop) - - return (torch.cat(crops, dim=0), ) - -######################################################################################################################## - -class BboxDetectorCombinedBatchMXD: - DESCRIPTION = "Run an Impact Pack BBOX_DETECTOR combined mask over each image in a batch." - CATEGORY = "MXD/Detector" - RETURN_TYPES = ("MASK",) - RETURN_NAMES = ("mask",) - FUNCTION = "detect" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "bbox_detector": ("BBOX_DETECTOR",), - "images": ("IMAGE",), - "threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), - "dilation": ("INT", {"default": 4, "min": -512, "max": 512, "step": 1}), - } - } - - def detect(self, bbox_detector, images, threshold=0.5, dilation=4): - if images.ndim == 3: - images = images.unsqueeze(0) - if images.ndim != 4: - raise ValueError(f"[BboxDetectorCombinedBatchMXD] Expected IMAGE tensor [B,H,W,C], got shape {tuple(images.shape)}") - - masks = [] - frame_count, height, width, _ = images.shape - pbar = comfy.utils.ProgressBar(frame_count) - - for i in range(frame_count): - frame = images[i:i + 1] - mask = bbox_detector.detect_combined(frame, threshold, dilation) - if mask is None: - mask = torch.zeros((height, width), dtype=torch.float32, device="cpu") - elif torch.is_tensor(mask): - mask = mask.detach().to(dtype=torch.float32, device="cpu") - else: - mask = torch.as_tensor(mask, dtype=torch.float32, device="cpu") - - if mask.ndim == 3 and mask.shape[0] == 1: - mask = mask.squeeze(0) - if mask.ndim != 2: - raise ValueError(f"[BboxDetectorCombinedBatchMXD] Detector returned unexpected mask shape {tuple(mask.shape)} for frame {i}.") - - masks.append(mask.unsqueeze(0)) - pbar.update(1) - - return (torch.cat(masks, dim=0),) - -######################################################################################################################## - -def _parse_mxd_mask_color(color_string): - if color_string is None: - return [255, 255, 255] - - text = str(color_string).strip() - color = [255, 255, 255] - - if "," in text: - try: - values = [float(channel.strip()) for channel in text.split(",")] - if all(0.0 <= value <= 1.0 for value in values): - color = [int(value * 255) for value in values] - else: - color = [int(value) for value in values] - except Exception: - color = [255, 255, 255] - else: - try: - color = list(ImageColor.getrgb(text)) - except Exception: - try: - value = float(text) - value = int(value * 255) if 0.0 <= value <= 1.0 else int(value) - color = [value, value, value] - except Exception: - color = [255, 255, 255] - - color = np.clip(color, 0, 255).astype(np.int32).tolist() - if len(color) < 3: - color = (color + [color[-1] if color else 255] * 3)[:3] - return color[:4] - - -def _mxd_image_batch(image): - if image is None: - return None - if image.ndim == 3: - image = image.unsqueeze(0) - if image.ndim != 4: - raise ValueError(f"[ImageAndMaskPreviewMXD] Expected IMAGE tensor [B,H,W,C], got shape {tuple(image.shape)}") - return image.to(dtype=torch.float32) - - -def _mxd_mask_batch(mask, height=None, width=None, batch_size=None, device=None): - if mask is None: - return None - - if mask.ndim == 2: - mask = mask.unsqueeze(0) - elif mask.ndim == 4 and mask.shape[-1] == 1: - mask = mask[..., 0] - elif mask.ndim == 4 and mask.shape[1] == 1: - mask = mask[:, 0] - - if mask.ndim != 3: - raise ValueError(f"[ImageAndMaskPreviewMXD] Expected MASK tensor [B,H,W], got shape {tuple(mask.shape)}") - - mask = mask.to(dtype=torch.float32, device=device if device is not None else mask.device).clamp(0.0, 1.0) - - if height is not None and width is not None and (mask.shape[-2] != height or mask.shape[-1] != width): - mask = F.interpolate(mask.unsqueeze(1), size=(height, width), mode="bilinear", align_corners=False).squeeze(1) - - if batch_size is not None: - mask = comfy.utils.repeat_to_batch_size(mask, batch_size) - - return mask - - -class ImageAndMaskPreviewMXD(SaveImage): - DESCRIPTION = """Return an image with a mask composited over it without creating a node preview.""" - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("composite",) - FUNCTION = "execute" - CATEGORY = "MXD/Image" - OUTPUT_NODE = False - - def __init__(self): - self.output_dir = folder_paths.get_temp_directory() - self.type = "temp" - self.prefix_append = "_temp_" + "".join(random.choice("abcdefghijklmnopqrstupvxyz") for _ in range(5)) - self.compress_level = 4 - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "mask_opacity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), - "mask_color": ("STRING", {"default": "255, 255, 255", "tooltip": "RGB/RGBA CSV, hex, or color name."}), - "pass_through": ("BOOLEAN", {"default": True, "tooltip": "Legacy option. This node now always returns the composite without creating a preview."}), - }, - "optional": { - "image": ("IMAGE",), - "mask": ("MASK",), - }, - "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, - } - - def _build_composite(self, image=None, mask=None, mask_opacity=1.0, mask_color="255, 255, 255"): - image = _mxd_image_batch(image) - - if image is None and mask is None: - raise ValueError("[ImageAndMaskPreviewMXD] Connect an image, a mask, or both.") - - if image is None: - mask = _mxd_mask_batch(mask) - return mask.unsqueeze(-1).expand(-1, -1, -1, 3).contiguous() - - if image.shape[-1] == 1: - image = image.expand(-1, -1, -1, 3).clone() - elif image.shape[-1] >= 3: - image = image[..., :3].clone() - else: - raise ValueError(f"[ImageAndMaskPreviewMXD] Expected IMAGE tensor with 1 or more channels, got shape {tuple(image.shape)}") - if mask is None: - return image - - batch_size, height, width, channels = image.shape - mask = _mxd_mask_batch(mask, height, width, batch_size, image.device) - color = _parse_mxd_mask_color(mask_color) - alpha = mask.mul(float(mask_opacity)).clamp(0.0, 1.0) - if len(color) == 4: - alpha = alpha * (color[3] / 255.0) - - rgb = torch.tensor(color[:3], dtype=image.dtype, device=image.device).view(1, 1, 1, channels) / 255.0 - alpha = alpha.unsqueeze(-1) - return (image * (1.0 - alpha) + rgb * alpha).clamp(0.0, 1.0) - - def execute(self, mask_opacity, mask_color, pass_through, filename_prefix="ComfyUI", image=None, mask=None, prompt=None, extra_pnginfo=None): - composite = self._build_composite(image=image, mask=mask, mask_opacity=mask_opacity, mask_color=mask_color) - return (composite,) - -######################################################################################################################## - -# NODE MAPPING -NODE_CLASS_MAPPINGS = { - "Flux Empty Latent Image": FluxEmptyLatentImage, - "Flux 2 Empty Latent Image": Flux2EmptyLatentImage, - "Sdxl Empty Latent Image": SdxlEmptyLatentImage, - "Flux Resolution Selector": FluxResolutionSelector, - "Image Scale To Total Pixels (SDXL Safe)": SDXLImageScaleToTotalPixelsSafe, - "Flux Image Scale To Total Pixels (Flux Safe)": FluxImageScaleToTotalPixelsSafe, - "Prompt With Guidance (Flux)": PromptWithGuidance, - "FluxResolutionMatcher": FluxResolutionMatcher, - "SDXLResolutionMatcher": SDXLResolutionMatcher, - "LatentHalfMasks": LatentHalfMasks, - "Get Latent Size": GetLatentSizeMXD, - "Place Image By Mask": PlaceImageByMask, - "Crop Image By Mask": CropImageByMask, - "Load Image Batch MXD": LoadImageBatchMXD, - "Load Video Batch MXD": LoadVideoBatchMXD, - "LoadImageFromFolderMXD": LoadImageFromFolderMXD, - "LoadVideoFromFolderMXD": LoadVideoFromFolderMXD, - "LoadImageWithPromptsMXD": LoadImageWithPromptsMXD, - "ZImageTurboEmptyLatentImage": ZImageTurboEmptyLatentImage, - "Save Image MXD": SaveImage_MXD, - "Extract Workflow From Image MXD": ExtractWorkflowFromImageMXD, - "SmartCropByMaskMXD": SmartCropByMaskMXD, - "BboxDetectorCombinedBatchMXD": BboxDetectorCombinedBatchMXD, - "ImageAndMaskPreviewMXD": ImageAndMaskPreviewMXD, - } - -if HAVE_COMFY_API: - NODE_CLASS_MAPPINGS.update({ - "QwenImageEditSingleMXD": QwenImageEditSingleMXD, - "QwenImageEditTripleMXD": QwenImageEditTripleMXD, - }) - -NODE_DISPLAY_NAME_MAPPINGS = { - "Flux Empty Latent Image": "Flux Empty Latent Image MXD", - "Flux 2 Empty Latent Image": "Flux 2 Empty Latent Image MXD", - "Sdxl Empty Latent Image": "SDXL Empty Latent Image MXD", - "Flux Resolution Selector": "Flux Resolution Selector MXD", - "Image Scale To Total Pixels (SDXL Safe)": "Scale SDXL Image MXD", - "Flux Image Scale To Total Pixels (Flux Safe)": "Scale Flux Image MXD", - "Prompt With Guidance (Flux)": "Prompt with Flux Guidance MXD", - "FluxResolutionMatcher": "Flux Resolution Matcher MXD", - "SDXLResolutionMatcher": "SDXL Resolution Matcher MXD", - "LatentHalfMasks": "Latent to L/R Masks MXD", - "Get Latent Size": "Get Latent Size MXD", - "Place Image By Mask": "Place Image by Mask MXD", - "Crop Image By Mask": "Crop Image by Mask MXD", - "Load Image Batch MXD": "Load Image Batch (Inputs/Outputs) MXD", - "Load Video Batch MXD": "Load Video Batch (Inputs/Outputs) MXD", - "LoadImageFromFolderMXD": "Load Image (From Folder) MXD", - "LoadVideoFromFolderMXD": "Load Video (From Folder) MXD", - "LoadImageWithPromptsMXD": "Load Image MXD", - "ZImageTurboEmptyLatentImage": "ZIT Empty Latent Image MXD", - "Save Image MXD": "Save Image MXD", - "Extract Workflow From Image MXD": "Extract Workflow From Image MXD", - "SmartCropByMaskMXD": "Smart Crop by Mask MXD", - "BboxDetectorCombinedBatchMXD": "BBOX Detector Combined Batch MXD", - "ImageAndMaskPreviewMXD": "Image and Mask Preview MXD", -} - -if HAVE_COMFY_API: - NODE_DISPLAY_NAME_MAPPINGS.update({ - "QwenImageEditSingleMXD": "Qwen Image Edit + Latent MXD", - "QwenImageEditTripleMXD": "Qwen Image Edit Prompt MXD (Triple)", - }) diff --git a/nodes/__init__.py b/nodes/__init__.py new file mode 100644 index 0000000..7e8e78a --- /dev/null +++ b/nodes/__init__.py @@ -0,0 +1,40 @@ +import importlib + +def _safe_import(module_name: str): + try: + return importlib.import_module(f".{module_name}", __name__) + except Exception as e: + print(f"[ComfyUI-MaxedOut] Failed to import '{module_name}': {e}") + return None + +def _get_mappings(mod): + if mod is None: + return {}, {} + class_map = getattr(mod, "NODE_CLASS_MAPPINGS", {}) or {} + display_map = getattr(mod, "NODE_DISPLAY_NAME_MAPPINGS", {}) or {} + return class_map, display_map + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +for _name in ( + "latents", + "resolution", + "prompts", + "masks", + "media_io", + "comparers", + "checkpoints", + "ffgo", + "wan22", + "ltx", +): + _mod = _safe_import(_name) + _class_map, _display_map = _get_mappings(_mod) + NODE_CLASS_MAPPINGS.update(_class_map) + NODE_DISPLAY_NAME_MAPPINGS.update(_display_map) + +__all__ = [ + "NODE_CLASS_MAPPINGS", + "NODE_DISPLAY_NAME_MAPPINGS", +] diff --git a/save_checkpoint_mxd.py b/nodes/checkpoints.py similarity index 69% rename from save_checkpoint_mxd.py rename to nodes/checkpoints.py index cb1e5b1..e87c422 100644 --- a/save_checkpoint_mxd.py +++ b/nodes/checkpoints.py @@ -1,9 +1,49 @@ +"""Checkpoint load/save nodes. + +Registered nodes: + LoadCheckpointMXD Load Checkpoint MXD (core loader + MXD info-icon UI) + SaveCheckpointMXD Save Checkpoint MXD (core saver with the FakeDevice fix) + +Import-time side effect: replaces comfy.diffusers_convert.cat_tensors with a +version that materializes lazily-cast weights first (see comment below). +""" import torch import folder_paths +import comfy.sd import comfy.diffusers_convert from comfy_extras.nodes_model_merging import save_checkpoint +class LoadCheckpointMXD: + DESCRIPTION = ( + "Loads a diffusion model checkpoint, same as the core Load Checkpoint node, " + "with the MXD info-icon UI (CivitAI lookup, cached metadata, local notes)." + ) + TITLE = "Load Checkpoint MXD" + CATEGORY = "MXD/Loaders" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "ckpt_name": (folder_paths.get_filename_list("checkpoints"),), + } + } + + RETURN_TYPES = ("MODEL", "CLIP", "VAE") + FUNCTION = "load_checkpoint" + + def load_checkpoint(self, ckpt_name): + ckpt_path = folder_paths.get_full_path_or_raise("checkpoints", ckpt_name) + out = comfy.sd.load_checkpoint_guess_config( + ckpt_path, + output_vae=True, + output_clip=True, + embedding_directory=folder_paths.get_folder_paths("embeddings"), + ) + return out[:3] + + # comfy's checkpoint saver builds the CLIP state dict via lazy "casting" params # (comfy.model_patcher.LazyCastingParam / LazyCastingParamPiece) whose .device # property returns a fake namedtuple ("FakeDevice") instead of a real torch.device, @@ -80,9 +120,11 @@ class SaveCheckpointMXD: NODE_CLASS_MAPPINGS = { + "LoadCheckpointMXD": LoadCheckpointMXD, "SaveCheckpointMXD": SaveCheckpointMXD, } NODE_DISPLAY_NAME_MAPPINGS = { + "LoadCheckpointMXD": "Load Checkpoint MXD", "SaveCheckpointMXD": "Save Checkpoint MXD", } diff --git a/mediacomparers.py b/nodes/comparers.py similarity index 98% rename from mediacomparers.py rename to nodes/comparers.py index 2010021..21d64be 100644 --- a/mediacomparers.py +++ b/nodes/comparers.py @@ -191,6 +191,3 @@ NODE_DISPLAY_NAME_MAPPINGS = { MxdImageComparerSave.NAME: "Image Comparer + Save MXD", MxdVideoComparer.NAME: "Video Comparer MXD", } - -WEB_DIRECTORY = "." -__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS", "WEB_DIRECTORY"] diff --git a/combine_materials_ffgo_mxd.py b/nodes/ffgo.py similarity index 100% rename from combine_materials_ffgo_mxd.py rename to nodes/ffgo.py diff --git a/nodes/latents.py b/nodes/latents.py new file mode 100644 index 0000000..549efb5 --- /dev/null +++ b/nodes/latents.py @@ -0,0 +1,305 @@ +from __future__ import annotations +import torch, comfy, comfy.model_management + +######################################################################################################################## +# Flux Empty Latent Image (SD3-compatible) +class FluxEmptyLatentImage: + DESCRIPTION = """Select a Flux resolution and create an empty latent batch.""" + TITLE = "Flux Empty Latent Image" + CATEGORY = "MXD/Latent" + + RESOLUTIONS = { + "— High Resolutions —": None, + "Square (1:1) 1408x1408": (1408, 1408), + "Standard (4:3) 1664x1216": (1664, 1216), + "Landscape (3:2) 1728x1152": (1728, 1152), + "Widescreen (16:9) 1920x1088": (1920, 1088), + "Ultrawide (21:9) 2176x960": (2176, 960), + + "— Standard Resolutions —": None, + "Square (1:1) 1024x1024": (1024, 1024), + "Standard (4:3) 1152x896": (1152, 896), + "Landscape (3:2) 1216x832": (1216, 832), + "Widescreen (16:9) 1344x768": (1344, 768), + "Ultrawide (21:9) 1536x640": (1536, 640), + + "— Low Resolutions —": None, + "Square (1:1) 320x320": (320, 320), + "Standard (4:3) 448x320": (448, 320), + "Landscape (3:2) 384x256": (384, 256), + "Widescreen (16:9) 448x256": (448, 256), + "Ultrawide (21:9) 576x256": (576, 256), + } + + def __init__(self): + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(cls) -> dict: + return { + "required": { + "resolution": ( + list(cls.RESOLUTIONS.keys()), + {"default": "Square (1:1) 1024x1024"} + ), + "vertical": ("BOOLEAN", {"default": False}), + "batch_size": ( + "INT", + { + "default": 1, + "min": 1, + "max": 4096, + "tooltip": "The number of latent images in the batch." + } + ) + } + } + + RETURN_TYPES = ("LATENT",) + OUTPUT_TOOLTIPS = ("The empty latent image batch.",) + FUNCTION = "generate" + + def generate(self, resolution, vertical, batch_size=1) -> tuple: + size = self.RESOLUTIONS.get(resolution) + if size is None: + raise ValueError(f"'{resolution}' is a header or invalid option.") + + width, height = size + if vertical: + width, height = height, width + + latent = torch.zeros([batch_size, 16, height // 8, width // 8], device=self.device) + return ({"samples": latent},) + +######################################################################################################################## +# Flux 2 Empty Latent Image (Flux2-compatible) +class Flux2EmptyLatentImage: + DESCRIPTION = """Select a Flux resolution and create an empty Flux 2 latent batch.""" + TITLE = "Flux 2 Empty Latent Image" + CATEGORY = "MXD/Latent" + + RESOLUTIONS = FluxEmptyLatentImage.RESOLUTIONS + + def __init__(self): + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(cls) -> dict: + return { + "required": { + "resolution": ( + list(cls.RESOLUTIONS.keys()), + {"default": "Square (1:1) 1024x1024"} + ), + "vertical": ("BOOLEAN", {"default": False}), + "batch_size": ( + "INT", + { + "default": 1, + "min": 1, + "max": 4096, + "tooltip": "The number of latent images in the batch." + } + ) + } + } + + RETURN_TYPES = ("LATENT",) + OUTPUT_TOOLTIPS = ("The empty Flux 2 latent image batch.",) + FUNCTION = "generate" + + def generate(self, resolution, vertical, batch_size=1) -> tuple: + size = self.RESOLUTIONS.get(resolution) + if size is None: + raise ValueError(f"'{resolution}' is a header or invalid option.") + + width, height = size + if vertical: + width, height = height, width + + latent = torch.zeros([batch_size, 128, height // 16, width // 16], device=self.device) + return ({"samples": latent},) + +######################################################################################################################## +# Flux Resolution Selector (for feeding into FluxEmptyLatentImage) +class FluxResolutionSelector: + DESCRIPTION = """Pick a Flux resolution string for Flux Empty Latent Image.""" + TITLE = "Flux Resolution Selector" + CATEGORY = "MXD/Latent" + + @classmethod + def INPUT_TYPES(cls) -> dict: + return { + "required": { + "resolution": ( + list(FluxEmptyLatentImage.RESOLUTIONS.keys()), # Include ALL keys including headers + {"default": "Square (1:1) 1024x1024"} + ), + } + } + + RETURN_TYPES = (list(FluxEmptyLatentImage.RESOLUTIONS.keys()),) + RETURN_NAMES = ("resolution",) + OUTPUT_TOOLTIPS = ("The selected resolution string for FluxEmptyLatentImage.",) + FUNCTION = "select_resolution" + + def select_resolution(self, resolution) -> tuple: + return (resolution,) + +######################################################################################################################## +# Sdxl Empty Latent Image +class SdxlEmptyLatentImage: + DESCRIPTION = """Select an SDXL resolution and create an empty latent batch.""" + TITLE = "Sdxl Empty Latent Image (With Resolutions)" + CATEGORY = "MXD/Latent" + + # SDXL predefined resolutions (width, height) + RESOLUTIONS = { + "Square (1:1) 1024x1024": (1024, 1024), + "Standard (4:3) 1152x896": (1152, 896), + "Landscape (3:2) 1216x832": (1216, 832), + "Widescreen (16:9) 1344x768": (1344, 768), + "Ultra-Wide (21:9) 1536x640": (1536, 640), + } + + def __init__(self): + # Retrieve the intermediate device (usually the GPU) from ComfyUI's model management. + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(cls) -> dict: + return { + "required": { + # Dropdown selection for one of the predefined SDXL resolutions. + "resolution": (list(cls.RESOLUTIONS.keys()),), + # Toggle for vertical mode (swaps width and height). + "vertical": ("BOOLEAN", {"default": False}), + # Number of latent images to create in the batch. + "batch_size": ( + "INT", + { + "default": 1, + "min": 1, + "max": 4096, + "tooltip": "The number of latent images in the batch." + } + ) + } + } + + RETURN_TYPES = ("LATENT",) + OUTPUT_TOOLTIPS = ("The empty latent image batch.",) + FUNCTION = "generate" + + def generate(self, resolution, vertical, batch_size=1) -> tuple: + # Get the selected resolution tuple (width, height) + width, height = self.RESOLUTIONS[resolution] + # If vertical mode is enabled, swap width and height. + if vertical: + width, height = height, width + + # Create an empty latent tensor. + # Typically, the latent space has 4 channels and each spatial dimension is 1/8th of the image. + latent = torch.zeros([batch_size, 4, height // 8, width // 8], device=self.device) + return ({"samples": latent},) + +######################################################################################################################## +# Z-Image Turbo Empty Latent Image (SD3-compatible) — Flux-style grouping +class ZImageTurboEmptyLatentImage: + DESCRIPTION = """Select a Z-Image Turbo resolution and create an empty latent batch.""" + TITLE = "Z-Image Turbo Empty Latent Image" + CATEGORY = "MXD/Latent" + + # Tuned for Z-Image Turbo: + # - Rule of 64: every dimension is a multiple of 64 + # - 1MP baseline: 1024x1024 in the standard tier + # - Ceiling: keep presets below 6.5MP + MAX_TOTAL_PIXELS = 6_500_000 + MIN_BLOCK = 64 + RESOLUTIONS = { + "— High Resolutions —": None, + "Square (1:1) 1536x1536": (1536, 1536), + "Photo (4:3) 1792x1344": (1792, 1344), + "Landscape (3:2) 1920x1280": (1920, 1280), + "Widescreen (16:9) 2048x1152": (2048, 1152), + "Ultrawide (21:9) 2304x1024": (2304, 1024), + + "— Standard Resolutions —": None, + "Square (1:1) 1024x1024": (1024, 1024), + "Photo (4:3) 1152x896": (1152, 896), + "Landscape (3:2) 1280x832": (1280, 832), + "Widescreen (16:9) 1344x768": (1344, 768), + "Ultrawide (21:9) 1536x640": (1536, 640), + + "— Low Resolutions —": None, + "Square (1:1) 512x512": (512, 512), + "Photo (4:3) 576x448": (576, 448), + "Landscape (3:2) 640x448": (640, 448), + "Widescreen (16:9) 704x384": (704, 384), + "Ultrawide (21:9) 768x320": (768, 320), + } + + def __init__(self): + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(cls) -> dict: + return { + "required": { + "resolution": ( + list(cls.RESOLUTIONS.keys()), + {"default": "Square (1:1) 1024x1024"} + ), + "vertical": ( + "BOOLEAN", + {"default": False, "tooltip": "Swap width and height."} + ), + "batch_size": ( + "INT", + {"default": 1, "min": 1, "max": 4096, "tooltip": "Number of latent images in the batch."} + ) + } + } + + RETURN_TYPES = ("LATENT",) + OUTPUT_TOOLTIPS = ("The empty Z-Image Turbo latent batch.",) + FUNCTION = "generate" + + def generate(self, resolution, vertical, batch_size=1) -> tuple: + size = self.RESOLUTIONS.get(resolution) + if size is None: + raise ValueError(f"'{resolution}' is a header or invalid option.") + + width, height = size + if vertical: + width, height = height, width + + if (width % self.MIN_BLOCK) != 0 or (height % self.MIN_BLOCK) != 0: + raise ValueError( + f"Invalid preset {width}x{height}. Z-Image Turbo requires multiples of {self.MIN_BLOCK}." + ) + if (width * height) > self.MAX_TOTAL_PIXELS: + raise ValueError( + f"Invalid preset {width}x{height}. Z-Image Turbo presets must stay at or below {self.MAX_TOTAL_PIXELS:,} pixels." + ) + + latent = torch.zeros([batch_size, 16, height // 8, width // 8], device=self.device) + return ({"samples": latent},) + +######################################################################################################################## + +NODE_CLASS_MAPPINGS = { + "Flux Empty Latent Image": FluxEmptyLatentImage, + "Flux 2 Empty Latent Image": Flux2EmptyLatentImage, + "Flux Resolution Selector": FluxResolutionSelector, + "Sdxl Empty Latent Image": SdxlEmptyLatentImage, + "ZImageTurboEmptyLatentImage": ZImageTurboEmptyLatentImage, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "Flux Empty Latent Image": "Flux Empty Latent Image MXD", + "Flux 2 Empty Latent Image": "Flux 2 Empty Latent Image MXD", + "Flux Resolution Selector": "Flux Resolution Selector MXD", + "Sdxl Empty Latent Image": "SDXL Empty Latent Image MXD", + "ZImageTurboEmptyLatentImage": "ZIT Empty Latent Image MXD", +} diff --git a/nodes/ltx/__init__.py b/nodes/ltx/__init__.py new file mode 100644 index 0000000..b64d26a --- /dev/null +++ b/nodes/ltx/__init__.py @@ -0,0 +1,18 @@ +"""LTX Video node package: latent sizing, two-stage samplers, taeltx live preview.""" +import importlib + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +for _name in ( + "latents", + "samplers", + "preview", +): + try: + _mod = importlib.import_module(f".{_name}", __name__) + except Exception as e: + print(f"[ComfyUI-MaxedOut] Failed to import 'nodes.ltx.{_name}': {e}") + continue + NODE_CLASS_MAPPINGS.update(getattr(_mod, "NODE_CLASS_MAPPINGS", {}) or {}) + NODE_DISPLAY_NAME_MAPPINGS.update(getattr(_mod, "NODE_DISPLAY_NAME_MAPPINGS", {}) or {}) diff --git a/nodes/ltx/latents.py b/nodes/ltx/latents.py new file mode 100644 index 0000000..7c7650f --- /dev/null +++ b/nodes/ltx/latents.py @@ -0,0 +1,266 @@ +"""LTX Video latent sizing: empty latent generator + two-stage image scaler. + +Registered nodes: + LTXVideoEmptyLatent_MXD LTX Empty Latent Video MXD + LTX_Image_Scaler_MXD LTX Video Image Scaler MXD + +Official LTX-2.3 rules (Lightricks model card + example workflows): + - Width & height must be divisible by 32; frame count must be 8n+1. + - The distilled two-stage workflow generates Stage 1 low-res, then the + ltx-2.3-spatial-upscaler-x2 doubles it (exactly 2x) for Stage 2. + - The one published two-stage resolution is Stage 1 960x544 -> 1920x1088. + +Tiers below are FINAL (Stage 2) sizes; Stage 1 is exactly half. Finals are +kept /64 so Stage 1 stays /32 (the latent constraint). Only the 1080p 16:9 +row is officially published by Lightricks; the portrait/square rows and the +720p/576p tiers are /32-aligned siblings at the same pixel budget. + +Buckets (FINAL size, all /64) -> Stage 1 (half, all /32): + 1080p: 1920x1088 / 1088x1920 / 1408x1408 (Stage 1: 960x544 / 544x960 / 704x704) + 720p: 1280x704 / 704x1280 / 960x960 (Stage 1: 640x352 / 352x640 / 480x480) + 576p: 1024x576 / 576x1024 / 768x768 (Stage 1: 512x288 / 288x512 / 384x384) + +Fit (no pad): proportional resize <= target, /64 aligned. +Crop (no pad): resize-to-cover then center-crop to exact bucket. +Square images map to each tier's square bucket. +""" +from __future__ import annotations + +import torch + +import comfy.utils +import comfy.model_management +import nodes + +from ..wan22.buckets import _is_squareish, _validate_image_batch_4d + + +######################################################################################################################## +# LTX Video Empty Latent Image +class LTXVideoEmptyLatentMXD: + DESCRIPTION = "Create an LTX Video empty latent batch from connected width/height and frame count." + TITLE = "LTX Empty Latent Video MXD" + CATEGORY = "MXD/Latent" + + # All dimensions must be multiples of 32 (LTX 32× spatial compression). + # Lengths must be 8n+1 for LTX's 8× temporal compression. + RESOLUTIONS = { + "16:9 Landscape": None, + "16:9 512×288": (512, 288), + "16:9 768×448": (768, 448), + "16:9 832×480": (832, 480), + "16:9 1024×576": (1024, 576), + "16:9 1280×736": (1280, 736), + + "9:16 Portrait": None, + "9:16 288×512": (288, 512), + "9:16 448×768": (448, 768), + "9:16 480×832": (480, 832), + "9:16 576×1024": (576, 1024), + + "4:3 Standard": None, + "4:3 512×384": (512, 384), + "4:3 768×576": (768, 576), + + "1:1 Square": None, + "1:1 512×512": (512, 512), + "1:1 768×768": (768, 768), + } + + def __init__(self): + self.device = comfy.model_management.intermediate_device() + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "length": ( + "INT", + { + "default": 97, + "min": 9, + "max": 1025, + "step": 8, + "tooltip": "Number of frames. Must be 8n+1 (e.g. 25, 49, 73, 97, 121, 201).", + }, + ), + "batch_size": ( + "INT", + { + "default": 1, + "min": 1, + "max": 4096, + "tooltip": "Number of latent videos in the batch.", + }, + ), + }, + "optional": { + "width": ("INT", { + "default": 640, "min": 32, "max": 8192, "step": 32, + "tooltip": "Stage 1 width. Connect the LTX Image Scaler stage1_width output for I2V workflows.", + }), + "height": ("INT", { + "default": 384, "min": 32, "max": 8192, "step": 32, + "tooltip": "Stage 1 height. Connect the LTX Image Scaler stage1_height output for I2V workflows.", + }), + }, + } + + RETURN_TYPES = ("LATENT", "INT") + RETURN_NAMES = ("latent", "length") + FUNCTION = "generate" + + def generate(self, length, batch_size=1, width=640, height=384): + # LTX latent: 128 channels, 32× spatial compression, 8× temporal compression + width = max(32, int(width) // 32 * 32) + height = max(32, int(height) // 32 * 32) + length = max(9, 1 + 8 * round((int(length) - 1) / 8)) + + t = ((length - 1) // 8) + 1 + h = height // 32 + w = width // 32 + latent = torch.zeros([batch_size, 128, t, h, w], device=self.device) + return ({"samples": latent}, length) + + +_LTX_BUCKETS = { + "1080p": {"landscape": (1920, 1088), "portrait": (1088, 1920), "square": (1408, 1408)}, + "720p": {"landscape": (1280, 704), "portrait": (704, 1280), "square": (960, 960)}, + "576p": {"landscape": (1024, 576), "portrait": (576, 1024), "square": (768, 768)}, +} + + +def _ceil32(x): + x = (int(x) + 31) // 32 * 32 + return max(32, x) + + +def _floor32(x): + x = int(x) // 32 * 32 + return max(32, x) + + +def _floor64(x): + x = int(x) // 64 * 64 + return max(64, x) + + +def _ltx_stage1_dims(final_w, final_h): + """Return Stage 1 dimensions that upscale exactly to the final size.""" + return max(32, int(final_w) // 2), max(32, int(final_h) // 2) + + +def _ltx_resize_fit_inside(img, out_w, out_h): + """Resize to fit inside (out_w, out_h), output /64 aligned on both sides.""" + _, ih, iw, _ = img.shape + s = min(out_w / iw, out_h / ih) + tw = _floor64(iw * s) + th = _floor64(ih * s) + tw = max(32, min(tw, nodes.MAX_RESOLUTION)) + th = max(32, min(th, nodes.MAX_RESOLUTION)) + resized = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) + return resized, tw, th + + +def _ltx_resize_then_center_crop(img, out_w, out_h): + """Resize to cover (out_w, out_h) then center-crop to exact /32 target.""" + _, ih, iw, _ = img.shape + s = max(out_w / iw, out_h / ih) + tw = _ceil32(iw * s) + th = _ceil32(ih * s) + tmp = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) + y0 = max(0, (th - out_h) // 2) + x0 = max(0, (tw - out_w) // 2) + return tmp[:, y0:y0+out_h, x0:x0+out_w, :] + + +def _ltx_pick_bucket(iw, ih, tier): + """Pick the landscape / portrait / square bucket for the given tier.""" + tier_map = _LTX_BUCKETS[tier] + if _is_squareish(iw, ih): + return tier_map["square"] + return tier_map["landscape"] if iw >= ih else tier_map["portrait"] + + +def _ltx_scale_image_core(image, tier="1080p", crop_to_fit=True): + """ + Core LTX scaler. Returns (scaled_image, final_w, final_h, stage1_w, stage1_h). + 'tier' is the FINAL (Stage 2) size budget; Stage 1 is exactly half. + """ + _, ih, iw, _ = image.shape + + bw, bh = _ltx_pick_bucket(iw, ih, tier) + + if _is_squareish(iw, ih): + crop_to_fit = False + + if crop_to_fit: + out = _ltx_resize_then_center_crop(image, bw, bh) + else: + out, bw, bh = _ltx_resize_fit_inside(image, bw, bh) + + final_w = int(out.shape[2]) + final_h = int(out.shape[1]) + stage1_w, stage1_h = _ltx_stage1_dims(final_w, final_h) + return out, final_w, final_h, stage1_w, stage1_h + + +class LTX_Image_Scaler_MXD: + """ + MXD Image Scaler for LTX Video (distilled two-stage workflow). + + 'tier' is the FINAL (Stage 2) size; Stage 1 is exactly half. Finals are /64 + so Stage 1 stays /32 (the LTX latent constraint). Wire stage1_width / + stage1_height into the empty latent for the low-res pass; the spatial + upscaler-x2 then doubles it back to the final size. + + Tiers (final / Stage 1): + 1080p 1920x1088 (official 16:9) / 1088x1920 / 1408x1408 -> half + 720p 1280x704 / 704x1280 / 960x960 -> half + 576p 1024x576 / 576x1024 / 768x768 -> half + + Modes: + Perfect Fit (Crops Edges) resize-to-cover + center-crop to exact bucket. + Closest Fit (No Crop) proportional resize, /64-aligned; may be smaller. + + Square images (within +-3% of 1:1) map to the tier's square bucket. + Outputs the scaled image at final size plus the Stage 1 dimensions. + """ + + TITLE = "LTX Video Image Scaler MXD" + CATEGORY = "image/processing" + RETURN_TYPES = ("IMAGE", "INT", "INT") + RETURN_NAMES = ("image", "width", "height") + FUNCTION = "scale" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "tier": (["1080p", "720p", "576p"], {"default": "1080p"}), + "crop_to_fit": ("BOOLEAN", { + "default": True, + "label_on": "Crop Edges", + "label_off": "Closest Fit (No Crop)", + }), + } + } + + def scale(self, image, tier="1080p", crop_to_fit=True): + image = _validate_image_batch_4d(image, "LTX_Image_Scaler_MXD", "image") + out, _final_w, _final_h, stage1_w, stage1_h = _ltx_scale_image_core( + image, tier=tier, crop_to_fit=crop_to_fit + ) + return (out, stage1_w, stage1_h) + + +NODE_CLASS_MAPPINGS = { + "LTXVideoEmptyLatent_MXD": LTXVideoEmptyLatentMXD, + "LTX_Image_Scaler_MXD": LTX_Image_Scaler_MXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LTXVideoEmptyLatent_MXD": "LTX Empty Latent Video MXD", + "LTX_Image_Scaler_MXD": "LTX Video Image Scaler MXD", +} diff --git a/nodes/ltx/preview.py b/nodes/ltx/preview.py new file mode 100644 index 0000000..fbaf1ba --- /dev/null +++ b/nodes/ltx/preview.py @@ -0,0 +1,329 @@ +"""LTX live preview via the tiny taeltx autoencoder. + +Registered node: + LTXPreview_MXD LTX Preview MXD (attach the previewer to ANY sampler's model) + +Core ComfyUI has no preview for the LTXAV format used by LTX 2.3, and the +latent2rgb approximation looks awful for video. This installs a previewer that +decodes latent frames with the tiny "taeltx" autoencoder for accurate previews. +The taeltx model is auto-discovered in the vae / vae_approx model folders. If +it isn't found, it is downloaded to the configured vae model folder. + +Sends MXD_live_preview_start / MXD_live_preview_frame / MXD_live_preview_saved +websocket events consumed by web/live_preview_panel_mxd.js. Final clips are +saved to /live_previews. + +TAE decode path borrowed from kjnodes / VideoHelperSuite. +""" +from __future__ import annotations +import os +import base64 +import time +import urllib.error +import urllib.request +from fractions import Fraction +from io import BytesIO +from PIL import Image +from threading import Lock, Thread + +import torch +import torch.nn.functional as F + +import comfy +import comfy.model_management +import comfy.patcher_extension +import comfy.utils +import server +import folder_paths as _folder_paths +from comfy_api.latest import VideoFromComponents, VideoComponents + +_serv = server.PromptServer.instance + +_TAELTX_FILENAME = "taeltx2_3.safetensors" +_TAELTX_URL = "https://huggingface.co/Kijai/LTX2.3_comfy/resolve/main/vae/taeltx2_3.safetensors?download=true" +_TAELTX_DOWNLOAD_LOCK = Lock() + + +def _find_taeltx_path(folder_paths): + for folder in ("vae", "vae_approx"): + try: + names = folder_paths.get_filename_list(folder) + except Exception: + continue + name = next((fn for fn in names if "taeltx" in fn.lower()), None) + if name is not None: + path = folder_paths.get_full_path(folder, name) + if path: + return path + return None + + +def _download_taeltx(folder_paths): + try: + vae_dirs = folder_paths.get_folder_paths("vae") + except Exception as exc: + print(f"[MXD LTX preview] cannot find ComfyUI vae model folder: {exc}") + return None + + if not vae_dirs: + print("[MXD LTX preview] cannot find ComfyUI vae model folder.") + return None + + target_dir = vae_dirs[0] + target_path = os.path.join(target_dir, _TAELTX_FILENAME) + partial_path = f"{target_path}.part" + + with _TAELTX_DOWNLOAD_LOCK: + if os.path.isfile(target_path): + return target_path + + try: + os.makedirs(target_dir, exist_ok=True) + print(f"[MXD LTX preview] downloading {_TAELTX_FILENAME} to {target_path}") + request = urllib.request.Request(_TAELTX_URL, headers={"User-Agent": "ComfyUI-MaxedOut"}) + with urllib.request.urlopen(request, timeout=120) as response, open(partial_path, "wb") as out: + while True: + chunk = response.read(1024 * 1024) + if not chunk: + break + out.write(chunk) + if not os.path.isfile(partial_path) or os.path.getsize(partial_path) == 0: + raise RuntimeError("downloaded file is empty") + os.replace(partial_path, target_path) + try: + folder_paths.get_filename_list("vae") + except Exception: + pass + print(f"[MXD LTX preview] downloaded {_TAELTX_FILENAME}") + return target_path + except (OSError, RuntimeError, urllib.error.URLError) as exc: + try: + if os.path.exists(partial_path): + os.remove(partial_path) + except OSError: + pass + print(f"[MXD LTX preview] failed to download {_TAELTX_FILENAME}: {exc}") + return None + + +def _load_taeltx(): + """Load the taeltx TAE from the vae / vae_approx model folders. Returns a VAE or None.""" + try: + import folder_paths + from comfy.sd import VAE + except Exception: + return None + + path = _find_taeltx_path(folder_paths) + if not path: + path = _download_taeltx(folder_paths) + if not path: + return None + + try: + taeltx = VAE(comfy.utils.load_torch_file(path)) + taeltx.first_stage_model.show_progress_bar = False + except Exception as exc: + print(f"[MXD LTX preview] failed to load taeltx ({path}): {exc}") + return None + return taeltx + + +class _LTXTAEPreviewer: + """Cycles through LTX video latent frames during sampling, decoding with taeltx.""" + + def __init__(self, taeltx, rate=8): + self.first_preview = True + self.last_time = 0.0 + self.c_index = 0 + self.rate = rate + self.taeltx = taeltx + + def decode_latent_to_preview_image(self, preview_format, x0): + if x0.ndim == 5: + x0 = x0.movedim(2, 1) + x0 = x0.reshape((-1,) + x0.shape[-3:]) + num_images = x0.size(0) + new_time = time.time() + num_previews = int((new_time - self.last_time) * self.rate) + self.last_time += num_previews / self.rate + if num_previews > num_images: + num_previews = num_images + elif num_previews <= 0: + return None + if self.first_preview: + self.first_preview = False + _serv.send_sync( + 'MXD_live_preview_start', + {'length': num_images, 'rate': self.rate, 'id': _serv.last_node_id}, + ) + self.last_time = new_time + 1.0 / self.rate + if self.c_index + num_previews > num_images: + frames = x0.roll(-self.c_index, 0)[:num_previews] + else: + frames = x0[self.c_index:self.c_index + num_previews] + Thread(target=self._send_frames, args=(frames, self.c_index, num_images)).run() + self.c_index = (self.c_index + num_previews) % num_images + return None + + def _send_frames(self, image_tensor, ind, leng): + max_size, min_size = 512, 256 + image_tensor = self._decode(image_tensor) + if image_tensor.size(1) < min_size or image_tensor.size(2) < min_size: + image_tensor = F.interpolate( + image_tensor.movedim(-1, 0), scale_factor=4, mode='nearest' + ).movedim(0, -1) + if image_tensor.size(1) > max_size or image_tensor.size(2) > max_size: + t = image_tensor.movedim(-1, 0) + if t.size(2) < t.size(3): + h = (max_size * t.size(2)) // t.size(3) + t = F.interpolate(t, (h, max_size), mode='nearest') + else: + w = (max_size * t.size(3)) // t.size(2) + t = F.interpolate(t, (max_size, w), mode='nearest') + image_tensor = t.movedim(0, -1) + previews = image_tensor.clamp(0, 1).mul(0xFF).to(device="cpu", dtype=torch.uint8) + node_id = _serv.last_node_id + for preview in previews: + img = Image.fromarray(preview.numpy()) + buf = BytesIO() + img.save(buf, format="JPEG", quality=90) + data_url = "data:image/jpeg;base64," + base64.b64encode(buf.getvalue()).decode("ascii") + _serv.send_sync('MXD_live_preview_frame', {'id': node_id, 'index': ind, 'length': leng, 'data': data_url}) + # taeltx expands the 8× temporal compression on decode + ind = (ind + 1) % ((leng - 1) * 8 + 1) + + def _decode(self, x0): + dev = comfy.model_management.get_torch_device() + dtype = self.taeltx.first_stage_model.decoder[1].weight.dtype + x0 = x0.unsqueeze(0).to(dtype=dtype, device=dev) + return self.taeltx.first_stage_model.decode(x0)[0].permute(1, 2, 3, 0) + + +def _save_final_ltx_preview(node_id, previewer, x0_v, rate): + """Decode the full final clip with taeltx and save it as an mp4 to output/live_previews.""" + try: + frames = x0_v.movedim(2, 1) + frames = frames.reshape((-1,) + frames.shape[-3:]) + frames = previewer._decode(frames).clamp(0, 1).to(device="cpu", dtype=torch.float32) + if frames.ndim != 4 or frames.size(0) == 0: + return + out_dir = os.path.join(_folder_paths.get_output_directory(), "live_previews") + os.makedirs(out_dir, exist_ok=True) + safe_id = str(node_id).replace(":", "_").replace("/", "_") + filename = f"{safe_id}_{int(time.time())}.mp4" + path = os.path.join(out_dir, filename) + video = VideoFromComponents(VideoComponents(images=frames, frame_rate=Fraction(max(1, round(rate))))) + video.save_to(path) + print(f"[MXD LTX preview] Saved live preview to {path}") + _serv.send_sync("MXD_live_preview_saved", { + "node_id": node_id, "filename": filename, "subfolder": "live_previews", "type": "output", + }) + except Exception as e: + print(f"[MXD LTX preview] Failed to save live preview: {e}") + + +class _LTXPreviewWrapper: + """OUTER_SAMPLE wrapper that installs the taeltx video previewer during sampling.""" + + def __init__(self, taeltx): + self.taeltx = taeltx + + def __call__(self, executor, noise, latent_image, sampler, sigmas, + denoise_mask, callback, disable_pbar, seed, latent_shapes): + guider = executor.class_obj + device = comfy.model_management.get_torch_device() + self.taeltx.first_stage_model.to(device) + + previewer = _LTXTAEPreviewer(self.taeltx, rate=8) + pbar = comfy.utils.ProgressBar(len(sigmas) - 1) + node_id = _serv.last_node_id + + # Strip I2V guide frames appended at the end of the latent before previewing. + num_keyframes = 0 + if 'positive' in guider.conds and guider.conds['positive']: + kf = guider.conds['positive'][0].get('keyframe_idxs') + if kf is not None: + num_keyframes = len(torch.unique(kf[0, 0, :, 0])) + + def ltx_callback(step, x0, x, total_steps): + x0_v = x0 + if x0_v is not None and len(latent_shapes) > 1: + # Audio+video latents are packed into [B, 1, total]; unpack and + # take the video tensor (the 5D one). Audio is a lower-rank entry. + x0_v = next( + (p for p in comfy.utils.unpack_latents(x0, latent_shapes) if p.ndim == 5), + None, + ) + if x0_v is not None and x0_v.ndim == 5 and num_keyframes > 0: + x0_v = x0_v[:, :, :-num_keyframes] + preview = ( + previewer.decode_latent_to_preview_image("JPEG", x0_v) + if x0_v is not None and x0_v.ndim == 5 else None + ) + pbar.update_absolute(step + 1, total_steps, preview) + if step + 1 >= total_steps and x0_v is not None and x0_v.ndim == 5: + _save_final_ltx_preview(node_id, previewer, x0_v, previewer.rate) + if callback is not None: + callback(step, x0, x, total_steps) + + try: + return executor( + noise, latent_image, sampler, sigmas, denoise_mask, + ltx_callback, disable_pbar, seed, latent_shapes=latent_shapes, + ) + finally: + self.taeltx.first_stage_model.to(comfy.model_management.unet_offload_device()) + + +######################################################################################################################## +# LTX Preview — attach the taeltx previewer to any model +class LTXPreviewMXD: + DESCRIPTION = ( + "Enables taeltx video previews during sampling for ANY sampler node " + "(SamplerCustomAdvanced, KSampler, etc.), not just the MXD LTX samplers. " + "LTX 2.3 (LTXAV) ships no built-in preview decoder, so core ComfyUI shows " + "nothing; this attaches a wrapper to the model that decodes latent frames " + "with the tiny taeltx autoencoder. Wire it between your model loader and " + "the sampler's model input. Downloads taeltx to your vae folder if missing." + ) + TITLE = "LTX Preview MXD" + CATEGORY = "MXD/Sampling" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "enabled": ("BOOLEAN", {"default": True, "tooltip": "Turn taeltx previews on/off without unwiring the node."}), + }, + } + + RETURN_TYPES = ("MODEL",) + RETURN_NAMES = ("model",) + FUNCTION = "apply" + OUTPUT_NODE = False + + def apply(self, model, enabled=True): + if not enabled: + return (model,) + taeltx = _load_taeltx() + if taeltx is None: + print("[MXD LTX preview] taeltx model not found in vae / vae_approx — skipping preview.") + return (model,) + model = model.clone() + model.add_wrapper_with_key( + comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, + "ltx_mxd_preview", + _LTXPreviewWrapper(taeltx), + ) + return (model,) + + +NODE_CLASS_MAPPINGS = { + "LTXPreview_MXD": LTXPreviewMXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LTXPreview_MXD": "LTX Preview MXD", +} diff --git a/nodes/ltx/samplers.py b/nodes/ltx/samplers.py new file mode 100644 index 0000000..d400c12 --- /dev/null +++ b/nodes/ltx/samplers.py @@ -0,0 +1,296 @@ +"""LTX two-stage distilled samplers. + +Registered nodes: + LTXKSampler_MXD LTX Stage 1 Sampler MXD (distilled 8-step schedule) + LTXKSampler2_MXD LTX Stage 2 Refiner MXD (official refine, start sigma 0.85) + +Sigma schedules come from the official Lightricks LTX-2.3 two-stage distilled +workflow (LTX-2.3_T2V_I2V_Two_Stage_Distilled.json). Custom Sigmas mode accepts +a manual descending schedule ending in 0.0 for experimentation. +""" +from __future__ import annotations +import re + +import torch + +import comfy +import comfy.model_management +import comfy.patcher_extension +import comfy.samplers +import comfy.sample +import comfy.utils +import latent_preview + +from .preview import _load_taeltx, _LTXPreviewWrapper + + +######################################################################################################################## +# Shared noise helper — equivalent to ComfyUI RandomNoise +class _LTXNoise: + def __init__(self, seed: int): + self.seed = seed + + def generate_noise(self, latent: dict) -> torch.Tensor: + samples = latent["samples"] + batch_inds = latent.get("batch_index", None) + return comfy.sample.prepare_noise(samples, self.seed, batch_inds) + + +######################################################################################################################## +# LTX KSampler — Stage 1 (T2V / I2V generation at base resolution) +class LTXKSamplerMXD: + DESCRIPTION = ( + "LTX-Video Stage 1 sampler for the distilled workflow. Use Distilled 8 Step " + "for the trained schedule, or Custom Sigmas when intentionally testing a " + "manual schedule." + ) + TITLE = "LTX Stage 1 Sampler MXD" + CATEGORY = "MXD/Sampling" + + MODES = ["Distilled 8 Step", "Custom Sigmas"] + _DISTILLED_SIGMAS = [1.0, 0.99375, 0.9875, 0.98125, 0.975, + 0.909375, 0.725, 0.421875, 0.0] + _CUSTOM_SIGMAS_DEFAULT = "1.0, 0.99375, 0.9875, 0.98125, 0.975, 0.909375, 0.725, 0.421875, 0.0" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "latent_image": ("LATENT",), + "mode": (cls.MODES, {"default": "Distilled 8 Step"}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "control_after_generate": True}), + "cfg": ("FLOAT", {"default": 2.0, "min": 0.0, "max": 100.0, "step": 0.1}), + "sampler_name": ( + ["euler_ancestral_cfg_pp", "euler_cfg_pp", "euler"], + {"default": "euler_ancestral_cfg_pp"}, + ), + "custom_sigmas": ( + "STRING", + { + "default": cls._CUSTOM_SIGMAS_DEFAULT, + "multiline": True, + "tooltip": "Only used when mode is Custom Sigmas. Enter comma, space, or newline separated sigma values.", + }, + ), + "ltx_preview": ("BOOLEAN", {"default": True, "tooltip": "Show LTX video previews during sampling. Downloads the taeltx VAE to your vae model folder if it is missing."}), + }, + } + + RETURN_TYPES = ("LATENT",) + RETURN_NAMES = ("latent",) + FUNCTION = "sample" + OUTPUT_NODE = False + + def sample( + self, + model, + positive, + negative, + latent_image, + mode="Distilled 8 Step", + seed=0, + cfg=2.0, + sampler_name="euler_ancestral_cfg_pp", + custom_sigmas=_CUSTOM_SIGMAS_DEFAULT, + ltx_preview=True, + ): + sigmas = _select_sigmas( + mode, + { + "Distilled 8 Step": self._DISTILLED_SIGMAS, + }, + custom_sigmas, + "LTX Stage 1 Sampler MXD", + ) + return _run_sampling(model, positive, negative, latent_image, seed, cfg, sampler_name, sigmas, ltx_preview) + + +######################################################################################################################## +# LTX KSampler 2 — Stage 2 (refinement at 2× resolution with distilled LoRA) +class LTXKSampler2MXD: + DESCRIPTION = ( + "LTX-Video Stage 2 refiner for the distilled workflow. Official Refine " + "matches the Lightricks 2.3 two-stage example (start sigma 0.85). " + "Custom Sigmas is for manual testing." + ) + TITLE = "LTX Stage 2 Refiner MXD" + CATEGORY = "MXD/Sampling" + + # Exact stage-2 refine schedule from the official Lightricks 2.3 two-stage + # workflow (LTX-2.3_T2V_I2V_Two_Stage_Distilled.json, euler_cfg_pp, cfg 1). + # Only the starting sigma (denoise strength) is meant to vary; use Custom + # Sigmas for that. + MODES = ["Official Refine", "Custom Sigmas"] + _OFFICIAL_REFINE_SIGMAS = [0.85, 0.725, 0.4219, 0.0] + _CUSTOM_SIGMAS_DEFAULT = "0.85, 0.725, 0.4219, 0.0" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "model": ("MODEL",), + "positive": ("CONDITIONING",), + "negative": ("CONDITIONING",), + "latent_image": ("LATENT",), + "mode": (cls.MODES, {"default": "Official Refine"}), + "seed": ("INT", {"default": 0, "min": 0, "max": 0xFFFFFFFFFFFFFFFF, "control_after_generate": True}), + "cfg": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1}), + "sampler_name": ( + ["euler_cfg_pp", "euler_ancestral_cfg_pp", "euler"], + {"default": "euler_cfg_pp"}, + ), + "custom_sigmas": ( + "STRING", + { + "default": cls._CUSTOM_SIGMAS_DEFAULT, + "multiline": True, + "tooltip": "Only used when mode is Custom Sigmas. Enter comma, space, or newline separated sigma values.", + }, + ), + "ltx_preview": ("BOOLEAN", {"default": True, "tooltip": "Show LTX video previews during sampling. Downloads the taeltx VAE to your vae model folder if it is missing."}), + }, + } + + RETURN_TYPES = ("LATENT",) + RETURN_NAMES = ("latent",) + FUNCTION = "sample" + OUTPUT_NODE = False + + def sample( + self, + model, + positive, + negative, + latent_image, + mode="Official Refine", + seed=0, + cfg=1.0, + sampler_name="euler_cfg_pp", + custom_sigmas=_CUSTOM_SIGMAS_DEFAULT, + ltx_preview=True, + ): + sigmas = _select_sigmas( + mode, + { + "Official Refine": self._OFFICIAL_REFINE_SIGMAS, + }, + custom_sigmas, + "LTX Stage 2 Refiner MXD", + ) + return _run_sampling(model, positive, negative, latent_image, seed, cfg, sampler_name, sigmas, ltx_preview) + + +######################################################################################################################## +# Sigma schedule helpers +_SIGMA_RE = re.compile(r"[-+]?(?:\d*\.\d+|\d+\.?)(?:[eE][-+]?\d+)?") + + +def _select_sigmas(mode, presets, custom_sigmas, node_name): + if mode == "Custom Sigmas": + values = _parse_custom_sigmas(custom_sigmas, node_name) + else: + try: + values = presets[mode] + except KeyError as exc: + allowed = ", ".join([*presets.keys(), "Custom Sigmas"]) + raise ValueError(f"{node_name}: unknown mode '{mode}'. Expected one of: {allowed}.") from exc + + return torch.tensor(values, dtype=torch.float32) + + +def _parse_custom_sigmas(custom_sigmas, node_name): + text = str(custom_sigmas or "") + values = [float(match.group(0)) for match in _SIGMA_RE.finditer(text)] + + if len(values) < 2: + raise ValueError(f"{node_name}: Custom Sigmas needs at least two sigma values, ending with 0.0.") + + for index, (left, right) in enumerate(zip(values, values[1:]), start=1): + if right > left: + raise ValueError( + f"{node_name}: Custom Sigmas must be in descending order. " + f"Value {index + 1} ({right}) is greater than value {index} ({left})." + ) + + if abs(values[-1]) > 1e-8: + raise ValueError(f"{node_name}: Custom Sigmas must end with 0.0.") + + return values + + +######################################################################################################################## +# Shared sampling logic +def _run_sampling(model, positive, negative, latent_image, seed, cfg, sampler_name, sigmas, ltx_preview=False): + taeltx = _load_taeltx() if ltx_preview else None + if ltx_preview and taeltx is None: + print("[MXD LTX preview] taeltx model not found in vae / vae_approx — skipping preview.") + + if taeltx is not None: + model = model.clone() + model.add_wrapper_with_key( + comfy.patcher_extension.WrappersMP.OUTER_SAMPLE, + "ltx_mxd_preview", + _LTXPreviewWrapper(taeltx), + ) + + guider = comfy.samplers.CFGGuider(model) + guider.set_conds(positive, negative) + guider.set_cfg(cfg) + + sampler = comfy.samplers.sampler_object(sampler_name) + + latent = latent_image.copy() + latent_samples = latent["samples"] + + try: + latent_samples = comfy.sample.fix_empty_latent_channels( + guider.model_patcher, latent_samples, + latent.get("downscale_ratio_spacial", None), + ) + except AttributeError: + pass + + latent["samples"] = latent_samples + noise_mask = latent.get("noise_mask", None) + + noise = _LTXNoise(seed) + + if taeltx is not None: + # The preview wrapper owns the progress bar / callback. + callback = None + else: + x0_output = {} + callback = latent_preview.prepare_callback(guider.model_patcher, sigmas.shape[-1] - 1, x0_output) + + disable_pbar = not comfy.utils.PROGRESS_BAR_ENABLED + + samples = guider.sample( + noise.generate_noise(latent), + latent_samples, + sampler, + sigmas, + denoise_mask=noise_mask, + callback=callback, + disable_pbar=disable_pbar, + seed=seed, + ) + samples = samples.to(comfy.model_management.intermediate_device()) + + out = latent.copy() + out.pop("downscale_ratio_spacial", None) + out["samples"] = samples + return (out,) + + +NODE_CLASS_MAPPINGS = { + "LTXKSampler_MXD": LTXKSamplerMXD, + "LTXKSampler2_MXD": LTXKSampler2MXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LTXKSampler_MXD": "LTX Stage 1 Sampler MXD", + "LTXKSampler2_MXD": "LTX Stage 2 Refiner MXD", +} diff --git a/nodes/masks.py b/nodes/masks.py new file mode 100644 index 0000000..f7e5b0c --- /dev/null +++ b/nodes/masks.py @@ -0,0 +1,563 @@ +from __future__ import annotations +import torch, comfy, comfy.utils, folder_paths, random +import torch.nn.functional as F +import numpy as np +from PIL import Image, ImageColor +from nodes import SaveImage + +######################################################################################################################## + +class LatentHalfMasks: + DESCRIPTION = """Split a latent into left and right half masks.""" + TITLE = "Latent Half Masks" + CATEGORY = "MXD/Latent" + + RETURN_TYPES = ("MASK", "MASK") + RETURN_NAMES = ("mask_left", "mask_right") + OUTPUT_TOOLTIPS = ( + "Mask covering the left half of the latent.", + "Mask covering the right half of the latent.", + ) + FUNCTION = "make_masks" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ("LATENT",), + } + } + + RETURN_TYPES = ("MASK", "MASK") + RETURN_NAMES = ("mask_left", "mask_right") + FUNCTION = "make_masks" + CATEGORY = "MXD/latent" + + def make_masks(self, latent): + # Infer width/height from latent (assumes 8x scale) + samples = latent.get("samples", None) + if samples is None or not isinstance(samples, torch.Tensor): + raise ValueError("LatentHalfMasks: invalid latent or missing 'samples' tensor.") + h_lat, w_lat = samples.shape[-2], samples.shape[-1] + w, h = int(w_lat * 8), int(h_lat * 8) + + # Always vertical, center split, no feather, no swap + split_px = w // 2 + left = torch.zeros((h, w), dtype=torch.float32) + right = torch.zeros((h, w), dtype=torch.float32) + left[:, :split_px] = 1.0 + right[:, split_px:] = 1.0 + + return left, right + +######################################################################################################################## + +# Get Latent Size +class GetLatentSizeMXD: + DESCRIPTION = """Get image width/height from a latent.""" + TITLE = "Get Latent Size" + CATEGORY = "MXD/Latent" + + RETURN_TYPES = ("INT", "INT") + RETURN_NAMES = ("width", "height") + OUTPUT_TOOLTIPS = ("Latent-derived image width in pixels.", "Latent-derived image height in pixels.") + FUNCTION = "get_size" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "latent": ("LATENT",), + } + } + + def get_size(self, latent): + if isinstance(latent, dict): + width = latent.get("width") + height = latent.get("height") + if width is not None and height is not None: + try: + return (int(width), int(height)) + except Exception: + pass + + samples = latent.get("samples") + else: + samples = None + + if samples is None or not isinstance(samples, torch.Tensor): + raise ValueError("GetLatentSizeMXD: invalid latent or missing 'samples' tensor.") + + channels = samples.shape[1] if samples.dim() >= 2 else 0 + scale = 16 if channels >= 64 else 8 + + h_lat, w_lat = samples.shape[-2], samples.shape[-1] + return (int(w_lat * scale), int(h_lat * scale)) + +######################################################################################################################## + +# --- Helper function to find the bounding box of a mask --- +def get_bounding_box(mask_tensor): + """ + Finds the bounding box of a non-zero region in a mask tensor. + The mask is expected to be a 2D tensor (H, W). + Returns a tuple (x_min, y_min, x_max, y_max) or None if the mask is empty. + """ + # Get non-zero coordinates from the mask + non_zero_coords = torch.nonzero(mask_tensor, as_tuple=False) + + # If the mask is empty, there is no bounding box + if non_zero_coords.numel() == 0: + return None + + # Find the min and max coordinates for y (dim 0) and x (dim 1) + min_y = non_zero_coords[:, 0].min().item() + max_y = non_zero_coords[:, 0].max().item() + min_x = non_zero_coords[:, 1].min().item() + max_x = non_zero_coords[:, 1].max().item() + + # The bounding box for PIL needs (left, upper, right, lower). + # We add +1 to the max values because the upper bound is exclusive. + return (min_x, min_y, max_x + 1, max_y + 1) + +# --- Tensor to PIL and PIL to Tensor conversion helpers --- +def tensor_to_pil(tensor): + """Converts a torch tensor (B, H, W, C) to a list of PIL Images.""" + if tensor is None: + return [] + + # Handle different tensor dimensions + if tensor.dim() == 4: # Batch of images + images = [] + for i in range(tensor.shape[0]): + img_np = 255. * tensor[i].cpu().numpy() + images.append(Image.fromarray(np.clip(img_np, 0, 255).astype(np.uint8))) + return images + elif tensor.dim() == 3: # Single image + img_np = 255. * tensor.cpu().numpy() + return [Image.fromarray(np.clip(img_np, 0, 255).astype(np.uint8))] + else: + raise ValueError(f"Unsupported tensor dimension: {tensor.dim()}") + +def pil_to_tensor(pil_images): + """Converts a list of PIL Images back to a torch tensor (B, H, W, C).""" + if not isinstance(pil_images, list): + pil_images = [pil_images] + + tensors = [] + for img in pil_images: + # Convert to RGB, then to a numpy array, normalize, and create a tensor + img_np = np.array(img.convert("RGB")).astype(np.float32) / 255.0 + tensors.append(torch.from_numpy(img_np).unsqueeze(0)) + + # Stack all tensors into a single batch tensor + return torch.cat(tensors, dim=0) + +# -------------------------------------------------------------------- +# ✨ The Main Node Class ✨ +# -------------------------------------------------------------------- +class PlaceImageByMask: + Description = """Place an overlay image inside the mask bounds on a base image.""" + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "base_image": ("IMAGE",), + "mask": ("MASK",), + "overlay_image": ("IMAGE",), + }, + "optional": { + "maintain_aspect_ratio": ("BOOLEAN", {"default": True}), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "place_image" + CATEGORY = "MXD/Image" + + def place_image(self, base_image, overlay_image, mask, maintain_aspect_ratio=True): + # Convert input tensors to lists of PIL Images + base_pils = tensor_to_pil(base_image) + overlay_pils = tensor_to_pil(overlay_image) + + processed_images = [] + + # Process each image in the batch + for i, base_pil in enumerate(base_pils): + # Work with an RGBA version of the base image for clean pasting + composited_image = base_pil.convert("RGBA") + + # Select the corresponding overlay and mask for the current base image + # Clamping the index prevents errors if batch sizes are mismatched + overlay_pil = overlay_pils[min(i, len(overlay_pils) - 1)].convert("RGBA") + current_mask = mask[min(i, mask.shape[0] - 1)] + + # Find the bounding box from the mask + bbox = get_bounding_box(current_mask) + + # If no mask is found, just use the original base image and skip to the next + if not bbox: + raise ValueError("The base image must be masked where you want the overlay to appear.") + + x_min, y_min, x_max, y_max = bbox + box_width = x_max - x_min + box_height = y_max - y_min + + # If the bounding box has no area, skip to the next image + if box_width <= 0 or box_height <= 0: + processed_images.append(base_pil) + continue + + # --- Resize the overlay image using the specified method --- + if maintain_aspect_ratio: + # Resize to fit *within* the box, preserving aspect ratio (like a thumbnail) + resized_overlay = overlay_pil.copy() + resized_overlay.thumbnail((box_width, box_height), Image.Resampling.LANCZOS) + + # Calculate position to center the resized overlay within the bounding box + paste_x = x_min + (box_width - resized_overlay.width) // 2 + paste_y = y_min + (box_height - resized_overlay.height) // 2 + paste_pos = (paste_x, paste_y) + else: + # As originally requested: stretch to fill the bounding box exactly + resized_overlay = overlay_pil.resize((box_width, box_height), resample=Image.Resampling.LANCZOS) + paste_pos = (x_min, y_min) + + # --- Paste the resized overlay onto the base image --- + # The alpha channel of the overlay itself is used as the mask for pasting. + # This ensures transparent areas of the overlay are handled correctly. + composited_image.paste(resized_overlay, paste_pos, resized_overlay) + + processed_images.append(composited_image) + + # Convert the list of processed PIL images back to a single batch tensor for output + output_tensor = pil_to_tensor(processed_images) + return (output_tensor,) + +###################################################################################################################################### + +class CropImageByMask: + DESCRIPTION = """Crop images to the mask bounds when a mask is provided.""" + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "image": ("IMAGE", ), + }, + "optional": { + "mask": ("MASK", ), + } + } + + RETURN_TYPES = ("IMAGE", ) + RETURN_NAMES = ("image", ) + FUNCTION = "crop" + CATEGORY = "MXD/image" + + def crop(self, image, mask=None): + # If no mask is provided or the mask is completely empty, return the original image + if mask is None or not torch.any(mask > 0): + return (image, ) + + B, H, W, C = image.shape + mask = mask.round() + + # Find bounding box for each batch + crops = [] + + for b in range(B): + current_mask = mask[min(b, mask.shape[0]-1)] + + # Check if the mask for this specific image is empty. + if not torch.any(current_mask > 0): + # If a specific mask in a batch is empty, we can't crop. + # To prevent errors with torch.cat later due to different sizes, + # we'll skip cropping for the whole batch and return the original. + # This ensures the output is always a valid tensor. + print("Warning: An empty mask was found in a batch. Returning original images.") + return (image, ) + + # Get coordinates of non-zero elements + rows = torch.any(current_mask > 0, dim=1) + cols = torch.any(current_mask > 0, dim=0) + + # Find boundaries + y_min, y_max = torch.where(rows)[0][[0, -1]] + x_min, x_max = torch.where(cols)[0][[0, -1]] + + # Crop image + crop = image[b:b+1, y_min:y_max+1, x_min:x_max+1, :] + crops.append(crop) + + # Note: This will raise an error if the crops have different sizes. + # The original code had this limitation. + cropped_images = torch.cat(crops, dim=0) + + return (cropped_images, ) + +######################################################################################################################## + +class SmartCropByMaskMXD: + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE", ), + "mask": ("MASK", ), + }, + } + + RETURN_TYPES = ("IMAGE", ) + RETURN_NAMES = ("image", ) + FUNCTION = "crop" + CATEGORY = "image/transform" + DESCRIPTION = "Slides a square crop window horizontally + vertically to center on subject mask." + + def crop(self, image, mask): + B, H, W, C = image.shape + mask = mask.round() + crops = [] + + for b in range(B): + mask_b = mask[min(b, mask.shape[0]-1)] + + # Get non-zero rows and columns + rows = torch.any(mask_b > 0, dim=1) + cols = torch.any(mask_b > 0, dim=0) + + # Default to center + center_x = W // 2 + center_y = H // 2 + + # Update center_x from mask if possible + if torch.any(cols): + x_min, x_max = torch.where(cols)[0][[0, -1]] + center_x = (x_min + x_max) // 2 + + # Update center_y from mask if possible + if torch.any(rows): + y_min, y_max = torch.where(rows)[0][[0, -1]] + center_y = (y_min + y_max) // 2 + + # Compute square crop box + side = min(H, W) + half = side // 2 + + left = max(0, center_x - half) + right = min(W, left + side) + left = right - side # clamp again + + top = max(0, center_y - half) + bottom = min(H, top + side) + top = bottom - side # clamp again + + # Final crop: safe slicing + crop = image[b:b+1, top:bottom, left:right, :] + crops.append(crop) + + return (torch.cat(crops, dim=0), ) + +######################################################################################################################## + +class BboxDetectorCombinedBatchMXD: + DESCRIPTION = "Run an Impact Pack BBOX_DETECTOR combined mask over each image in a batch." + CATEGORY = "MXD/Detector" + RETURN_TYPES = ("MASK",) + RETURN_NAMES = ("mask",) + FUNCTION = "detect" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "bbox_detector": ("BBOX_DETECTOR",), + "images": ("IMAGE",), + "threshold": ("FLOAT", {"default": 0.5, "min": 0.0, "max": 1.0, "step": 0.01}), + "dilation": ("INT", {"default": 4, "min": -512, "max": 512, "step": 1}), + } + } + + def detect(self, bbox_detector, images, threshold=0.5, dilation=4): + if images.ndim == 3: + images = images.unsqueeze(0) + if images.ndim != 4: + raise ValueError(f"[BboxDetectorCombinedBatchMXD] Expected IMAGE tensor [B,H,W,C], got shape {tuple(images.shape)}") + + masks = [] + frame_count, height, width, _ = images.shape + pbar = comfy.utils.ProgressBar(frame_count) + + for i in range(frame_count): + frame = images[i:i + 1] + mask = bbox_detector.detect_combined(frame, threshold, dilation) + if mask is None: + mask = torch.zeros((height, width), dtype=torch.float32, device="cpu") + elif torch.is_tensor(mask): + mask = mask.detach().to(dtype=torch.float32, device="cpu") + else: + mask = torch.as_tensor(mask, dtype=torch.float32, device="cpu") + + if mask.ndim == 3 and mask.shape[0] == 1: + mask = mask.squeeze(0) + if mask.ndim != 2: + raise ValueError(f"[BboxDetectorCombinedBatchMXD] Detector returned unexpected mask shape {tuple(mask.shape)} for frame {i}.") + + masks.append(mask.unsqueeze(0)) + pbar.update(1) + + return (torch.cat(masks, dim=0),) + +######################################################################################################################## + +def _parse_mxd_mask_color(color_string): + if color_string is None: + return [255, 255, 255] + + text = str(color_string).strip() + color = [255, 255, 255] + + if "," in text: + try: + values = [float(channel.strip()) for channel in text.split(",")] + if all(0.0 <= value <= 1.0 for value in values): + color = [int(value * 255) for value in values] + else: + color = [int(value) for value in values] + except Exception: + color = [255, 255, 255] + else: + try: + color = list(ImageColor.getrgb(text)) + except Exception: + try: + value = float(text) + value = int(value * 255) if 0.0 <= value <= 1.0 else int(value) + color = [value, value, value] + except Exception: + color = [255, 255, 255] + + color = np.clip(color, 0, 255).astype(np.int32).tolist() + if len(color) < 3: + color = (color + [color[-1] if color else 255] * 3)[:3] + return color[:4] + + +def _mxd_image_batch(image): + if image is None: + return None + if image.ndim == 3: + image = image.unsqueeze(0) + if image.ndim != 4: + raise ValueError(f"[ImageAndMaskPreviewMXD] Expected IMAGE tensor [B,H,W,C], got shape {tuple(image.shape)}") + return image.to(dtype=torch.float32) + + +def _mxd_mask_batch(mask, height=None, width=None, batch_size=None, device=None): + if mask is None: + return None + + if mask.ndim == 2: + mask = mask.unsqueeze(0) + elif mask.ndim == 4 and mask.shape[-1] == 1: + mask = mask[..., 0] + elif mask.ndim == 4 and mask.shape[1] == 1: + mask = mask[:, 0] + + if mask.ndim != 3: + raise ValueError(f"[ImageAndMaskPreviewMXD] Expected MASK tensor [B,H,W], got shape {tuple(mask.shape)}") + + mask = mask.to(dtype=torch.float32, device=device if device is not None else mask.device).clamp(0.0, 1.0) + + if height is not None and width is not None and (mask.shape[-2] != height or mask.shape[-1] != width): + mask = F.interpolate(mask.unsqueeze(1), size=(height, width), mode="bilinear", align_corners=False).squeeze(1) + + if batch_size is not None: + mask = comfy.utils.repeat_to_batch_size(mask, batch_size) + + return mask + + +class ImageAndMaskPreviewMXD(SaveImage): + DESCRIPTION = """Return an image with a mask composited over it without creating a node preview.""" + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("composite",) + FUNCTION = "execute" + CATEGORY = "MXD/Image" + OUTPUT_NODE = False + + def __init__(self): + self.output_dir = folder_paths.get_temp_directory() + self.type = "temp" + self.prefix_append = "_temp_" + "".join(random.choice("abcdefghijklmnopqrstupvxyz") for _ in range(5)) + self.compress_level = 4 + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "mask_opacity": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}), + "mask_color": ("STRING", {"default": "255, 255, 255", "tooltip": "RGB/RGBA CSV, hex, or color name."}), + "pass_through": ("BOOLEAN", {"default": True, "tooltip": "Legacy option. This node now always returns the composite without creating a preview."}), + }, + "optional": { + "image": ("IMAGE",), + "mask": ("MASK",), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + def _build_composite(self, image=None, mask=None, mask_opacity=1.0, mask_color="255, 255, 255"): + image = _mxd_image_batch(image) + + if image is None and mask is None: + raise ValueError("[ImageAndMaskPreviewMXD] Connect an image, a mask, or both.") + + if image is None: + mask = _mxd_mask_batch(mask) + return mask.unsqueeze(-1).expand(-1, -1, -1, 3).contiguous() + + if image.shape[-1] == 1: + image = image.expand(-1, -1, -1, 3).clone() + elif image.shape[-1] >= 3: + image = image[..., :3].clone() + else: + raise ValueError(f"[ImageAndMaskPreviewMXD] Expected IMAGE tensor with 1 or more channels, got shape {tuple(image.shape)}") + if mask is None: + return image + + batch_size, height, width, channels = image.shape + mask = _mxd_mask_batch(mask, height, width, batch_size, image.device) + color = _parse_mxd_mask_color(mask_color) + alpha = mask.mul(float(mask_opacity)).clamp(0.0, 1.0) + if len(color) == 4: + alpha = alpha * (color[3] / 255.0) + + rgb = torch.tensor(color[:3], dtype=image.dtype, device=image.device).view(1, 1, 1, channels) / 255.0 + alpha = alpha.unsqueeze(-1) + return (image * (1.0 - alpha) + rgb * alpha).clamp(0.0, 1.0) + + def execute(self, mask_opacity, mask_color, pass_through, filename_prefix="ComfyUI", image=None, mask=None, prompt=None, extra_pnginfo=None): + composite = self._build_composite(image=image, mask=mask, mask_opacity=mask_opacity, mask_color=mask_color) + return (composite,) + +######################################################################################################################## + +NODE_CLASS_MAPPINGS = { + "LatentHalfMasks": LatentHalfMasks, + "Get Latent Size": GetLatentSizeMXD, + "Place Image By Mask": PlaceImageByMask, + "Crop Image By Mask": CropImageByMask, + "SmartCropByMaskMXD": SmartCropByMaskMXD, + "BboxDetectorCombinedBatchMXD": BboxDetectorCombinedBatchMXD, + "ImageAndMaskPreviewMXD": ImageAndMaskPreviewMXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "LatentHalfMasks": "Latent to L/R Masks MXD", + "Get Latent Size": "Get Latent Size MXD", + "Place Image By Mask": "Place Image by Mask MXD", + "Crop Image By Mask": "Crop Image by Mask MXD", + "SmartCropByMaskMXD": "Smart Crop by Mask MXD", + "BboxDetectorCombinedBatchMXD": "BBOX Detector Combined Batch MXD", + "ImageAndMaskPreviewMXD": "Image and Mask Preview MXD", +} diff --git a/nodes/media_io.py b/nodes/media_io.py new file mode 100644 index 0000000..dfca96d --- /dev/null +++ b/nodes/media_io.py @@ -0,0 +1,864 @@ +from __future__ import annotations +import torch, os, folder_paths, node_helpers, json, hashlib, re +import numpy as np +from PIL import Image, ImageOps, ImageSequence +from nodes import PreviewImage, SaveImage +try: + from comfy_api.input_impl import VideoFromFile + HAVE_COMFY_API_VIDEO = True +except Exception as _e: + VideoFromFile = None + HAVE_COMFY_API_VIDEO = False + print(f"[ComfyUI-MaxedOut] comfy_api video I/O not available in media_io: {_e}") + +######################################################################################################################## +# ---------- Helpers (copied from latent loader style) ---------- +def _safe_json_loads(s): + if s is None: + return None + if isinstance(s, bytes): + try: + s = s.decode("utf-8", "ignore") + except Exception: + return None + if not isinstance(s, str): + return None + try: + return json.loads(s) + except Exception: + try: + return json.loads(json.loads(s)) + except Exception: + return None + + +def _extract_params_from_prompt_json(prompt_json: dict): + """ + Returns (positive, negative) from saved Comfy prompt graph. + """ + pos = "" + neg = "" + if not isinstance(prompt_json, dict): + return pos, neg + + # unwrap if saved as {"prompt": {...}} + graph = prompt_json.get("prompt", prompt_json) + if not isinstance(graph, dict): + return pos, neg + + # try to find KSampler/KSamplerAdvanced node + ks = None + for _, v in graph.items(): + if "KSampler" in v.get("class_type", ""): + ks = v + break + if not ks: + return pos, neg + + kin = ks.get("inputs", {}) + + def _as_node_id(x): + return str(x[0]) if isinstance(x, (list, tuple)) and x else None + + def _text_from_clip(node_id): + n = graph.get(str(node_id), {}) + if n.get("class_type") == "CLIPTextEncode": + return str(n.get("inputs", {}).get("text", "")).strip() + return "" + + pos = _text_from_clip(_as_node_id(kin.get("positive"))) + neg = _text_from_clip(_as_node_id(kin.get("negative"))) + + return pos, neg + +def _strip_counter(name: str) -> str: + # Only strip the trailing pattern we generate when saving: "_<5digits>_" + # Preserve numeric-only base names like "96". + stem, _ = os.path.splitext(name) + m = re.match(r"^(.*?)(?:_\d{5}_)$", stem) + return m.group(1) if m else stem + +# ---------- Node ---------- +def _indent_paths(paths): + indented = [] + for path in paths: + if not path: + indented.append("") + continue + clean_path = path.lstrip("  ") + depth = clean_path.count("/") + indent = " " * (depth * 4) + indented.append(indent + clean_path) + return indented + + +def _scan_subdir_mtimes(root: str, subdirs: set, branch_latest: dict, exts: tuple = None): + """ + Walk `root`, adding every subfolder's relative path to `subdirs` and + bubbling the mtime of its most recently modified file up to every + ancestor branch (including "" for the root) in `branch_latest`. + + When `exts` is given, a folder (and its ancestors) is only added if it + directly or recursively contains at least one file matching `exts` -- + so folders with no relevant content don't show up as pickable at all. + """ + try: + for dirpath, dirnames, filenames in os.walk(root): + # Exclude hidden folders (e.g. .git, .github) and __pycache__ + dirnames[:] = [d for d in dirnames if not d.startswith('.') and d != '__pycache__'] + rel_path = os.path.relpath(dirpath, root) + rel_path = "" if rel_path == "." else rel_path.replace(os.path.sep, "/") + + latest = 0.0 + has_match = exts is None + for f in filenames: + if exts and not f.lower().endswith(exts): + continue + has_match = True + try: + m = os.path.getmtime(os.path.join(dirpath, f)) + except OSError: + continue + if m > latest: + latest = m + + if not has_match: + continue + + if rel_path: + subdirs.add(rel_path) + + parts = [p for p in rel_path.split("/") if p] + for i in range(len(parts) + 1): + branch = "/".join(parts[:i]) + if latest > branch_latest.get(branch, -1.0): + branch_latest[branch] = latest + if i > 0: + subdirs.add(branch) + except OSError: + pass + + +def _list_image_batch_subdirs(root: str, exts: tuple = None): + """ + Recursive subfolders under `root`, newest first. Each folder is ordered by + the mtime of the most recently modified file anywhere inside it (so a + folder that just received a new file jumps back to the top). '' = the + root itself, always first. If `exts` is given, only folders that + directly or recursively contain a matching file are included. + """ + subdirs = set() + branch_latest = {} + _scan_subdir_mtimes(root, subdirs, branch_latest, exts) + ordered = sorted(subdirs, key=lambda d: (-branch_latest.get(d, -1.0), d.lower())) + return [""] + ordered + + +def _list_image_batch_subdirs_union(output_root: str, input_root: str, exts: tuple = None): + """ + Union of recursive subfolders from both roots, newest first. A folder + present under both roots is ranked by whichever side has the more + recent file, so it doesn't matter which source the user has selected. + """ + subdirs = set() + branch_latest = {} + _scan_subdir_mtimes(output_root, subdirs, branch_latest, exts) + _scan_subdir_mtimes(input_root, subdirs, branch_latest, exts) + ordered = sorted(subdirs, key=lambda d: (-branch_latest.get(d, -1.0), d.lower())) + return [""] + ordered + + +IMAGE_BATCH_EXTS = (".png", ".jpg", ".jpeg", ".webp") +VIDEO_BATCH_EXTS = (".mp4",) + + +def _sort_paths_newest_first(paths): + """Sort file paths by mtime desc (newest first), stable by normalized path.""" + def _mtime(path): + try: + return os.path.getmtime(path) + except OSError: + return 0.0 + + return sorted(paths, key=lambda p: (-_mtime(p), p.replace("\\", "/").lower())) + + +def _list_files_recursive(root: str, exts: tuple): + """Recursively list files under `root` matching `exts`, newest first, as relpaths.""" + try: + files = [] + for dirpath, dirnames, filenames in os.walk(root): + dirnames[:] = [d for d in dirnames if not d.startswith('.') and d != '__pycache__'] + for f in filenames: + if f.lower().endswith(exts): + files.append(os.path.join(dirpath, f)) + files = _sort_paths_newest_first(files) + return [os.path.relpath(f, root).replace(os.sep, "/") for f in files] + except OSError: + return [] + + +def _list_files_recursive_union(output_root: str, input_root: str, exts: tuple): + """ + Union of recursive files from both roots, newest first. A relative path + present under both roots is ranked by whichever side's file is more + recent, so it doesn't matter which source the user has selected. + """ + mtimes = {} + + def scan(root): + try: + for dirpath, dirnames, filenames in os.walk(root): + dirnames[:] = [d for d in dirnames if not d.startswith('.') and d != '__pycache__'] + for f in filenames: + if not f.lower().endswith(exts): + continue + full = os.path.join(dirpath, f) + rel = os.path.relpath(full, root).replace(os.sep, "/") + try: + m = os.path.getmtime(full) + except OSError: + m = 0.0 + if m > mtimes.get(rel, -1.0): + mtimes[rel] = m + except OSError: + pass + + scan(output_root) + scan(input_root) + return sorted(mtimes, key=lambda p: (-mtimes[p], p.lower())) or [""] + + +# Server routes so the frontend can swap folder/file dropdowns between +# inputs/outputs without reloading the page. +try: + from server import PromptServer as _MXD_PromptServer + from aiohttp import web as _mxd_web + + @_MXD_PromptServer.instance.routes.get("/mxd/image_batch/folders") + async def _mxd_list_image_batch_folders(request): + return _mxd_web.json_response({ + "outputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_output_directory())), + "inputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_input_directory())), + }) + + @_MXD_PromptServer.instance.routes.get("/mxd/video_batch/folders") + async def _mxd_list_video_batch_folders(request): + return _mxd_web.json_response({ + "outputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_output_directory(), VIDEO_BATCH_EXTS)), + "inputs": _indent_paths(_list_image_batch_subdirs(folder_paths.get_input_directory(), VIDEO_BATCH_EXTS)), + }) + + @_MXD_PromptServer.instance.routes.get("/mxd/single_loader/files") + async def _mxd_list_single_loader_files(request): + kind = request.query.get("kind", "image") + exts = VIDEO_BATCH_EXTS if kind == "video" else IMAGE_BATCH_EXTS + return _mxd_web.json_response({ + "outputs": _list_files_recursive(folder_paths.get_output_directory(), exts), + "inputs": _list_files_recursive(folder_paths.get_input_directory(), exts), + }) +except Exception as _e: + print(f"[LoadImageBatchMXD] Could not register folders route: {_e}") + + +class LoadImageBatchMXD: + DESCRIPTION = """Load images from an inputs or outputs folder, make masks from alpha, and read prompts.""" + TITLE = "Load Image Batch (Inputs/Outputs + Prompts)" + CATEGORY = "MXD/Image" + + RETURN_TYPES = ("IMAGE", "MASK", "STRING", "STRING") + RETURN_NAMES = ("IMAGE", "MASK", "positive", "negative") + OUTPUT_IS_LIST = (True, True, True, True) + FUNCTION = "load_batch" + + @classmethod + def INPUT_TYPES(cls): + # Provide the union of inputs + outputs subfolders so any saved value + # validates regardless of which source it belongs to. The frontend + # filters the visible list down to the selected source on the fly. + union = _indent_paths(_list_image_batch_subdirs_union( + folder_paths.get_output_directory(), folder_paths.get_input_directory() + )) + return { + "required": { + "source": (("outputs", "inputs"), {"default": "outputs"}), + "folder": (tuple(union), {"default": ""}), + } + } + + def _extract_prompts(self, image: Image.Image): + pos, neg = "", "" + try: + raw = image.info.get("prompt") + if raw: + prompt_json = _safe_json_loads(raw) + if prompt_json: + pos, neg = _extract_params_from_prompt_json(prompt_json) + else: + pos = raw + except Exception as e: + print(f"[LoadImageBatchMXD] Prompt parse failed: {e}") + return pos, neg + + def load_batch(self, folder: str, source: str = "outputs"): + folder = folder.lstrip("  ") + root = ( + folder_paths.get_input_directory() + if source == "inputs" + else folder_paths.get_output_directory() + ) + folder_path = os.path.normpath(os.path.join(root, folder)) if folder else root + + if not os.path.isdir(folder_path): + raise FileNotFoundError(f"No such folder: {folder_path}") + + valid_exts = IMAGE_BATCH_EXTS + + # Recursively find all matching files + files = [] + for dirpath, dirnames, filenames in os.walk(folder_path): + dirnames.sort() + for f in sorted(filenames): + if f.lower().endswith(valid_exts): + files.append(os.path.join(dirpath, f)) + + if not files: + raise FileNotFoundError(f"No valid images found in folder '{folder_path}' (including subfolders)") + + images, masks, positives, negatives, prefixes = [], [], [], [], [] + + for path in files: + i = Image.open(path) + i = ImageOps.exif_transpose(i) + + pos, neg = self._extract_prompts(i) + positives.append(pos) + negatives.append(neg) + + rgb = i.convert("RGB") + arr = np.array(rgb).astype(np.float32) / 255.0 + img_t = torch.from_numpy(arr)[None, ...] + + if 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask_t = 1.0 - torch.from_numpy(mask).unsqueeze(0) + else: + h, w = arr.shape[:2] + mask_t = torch.zeros((1, h, w), dtype=torch.float32) + + images.append(img_t) + masks.append(mask_t) + + return (images, masks, positives, negatives) + + +class LoadVideoBatchMXD: + DESCRIPTION = """Load videos from an inputs or outputs folder as a batch.""" + TITLE = "Load Video Batch (Inputs/Outputs)" + CATEGORY = "MXD/Video" + + RETURN_TYPES = ("VIDEO",) + RETURN_NAMES = ("VIDEO",) + OUTPUT_IS_LIST = (True,) + FUNCTION = "load_batch" + + @classmethod + def INPUT_TYPES(cls): + # Same union-of-sources pattern as LoadImageBatchMXD; reuses that + # node's folder listing helper, filtered to folders that actually + # contain a video so empty/irrelevant folders don't show up. + union = _indent_paths(_list_image_batch_subdirs_union( + folder_paths.get_output_directory(), folder_paths.get_input_directory(), VIDEO_BATCH_EXTS + )) + return { + "required": { + "source": (("outputs", "inputs"), {"default": "outputs"}), + "folder": (tuple(union), {"default": ""}), + } + } + + def load_batch(self, folder: str, source: str = "outputs"): + if not HAVE_COMFY_API_VIDEO: + raise RuntimeError( + "[LoadVideoBatchMXD] Video output requires a newer ComfyUI core with " + "comfy_api.latest / comfy_api.input_impl support. Please update ComfyUI." + ) + + folder = folder.lstrip("  ") + root = ( + folder_paths.get_input_directory() + if source == "inputs" + else folder_paths.get_output_directory() + ) + folder_path = os.path.normpath(os.path.join(root, folder)) if folder else root + + if not os.path.isdir(folder_path): + raise FileNotFoundError(f"No such folder: {folder_path}") + + valid_exts = VIDEO_BATCH_EXTS + + # Recursively find all matching files + files = [] + for dirpath, dirnames, filenames in os.walk(folder_path): + dirnames.sort() + for f in sorted(filenames): + if f.lower().endswith(valid_exts): + files.append(os.path.join(dirpath, f)) + + if not files: + raise FileNotFoundError(f"No valid videos found in folder '{folder_path}' (including subfolders)") + + videos = [VideoFromFile(path) for path in files] + + return (videos,) + + +class LoadImageFromFolderMXD: + DESCRIPTION = ( + "Load a single image from any inputs/outputs subfolder. Turn on run_folder " + "to auto-queue every image in that same folder, one after another." + ) + TITLE = "Load Image (From Folder) MXD" + CATEGORY = "MXD/Image" + + RETURN_TYPES = ("IMAGE", "MASK", "STRING", "STRING", "STRING") + RETURN_NAMES = ("IMAGE", "MASK", "positive", "negative", "filename") + FUNCTION = "load_image" + + @classmethod + def INPUT_TYPES(cls): + # Union of both sources so any saved value validates regardless of which + # source it belongs to; the frontend narrows the visible list to the + # selected source on the fly (mirrors LoadImageBatchMXD's folder picker). + union = _list_files_recursive_union( + folder_paths.get_output_directory(), folder_paths.get_input_directory(), IMAGE_BATCH_EXTS + ) + return { + "required": { + "source": (("outputs", "inputs"), {"default": "outputs"}), + "image": (tuple(union), ), + "run_folder": ("BOOLEAN", { + "default": False, + "tooltip": "When enabled, hitting Queue Prompt auto-queues every image in this file's folder, one after another, instead of just the selected file.", + }), + } + } + + def _extract_prompts(self, image: Image.Image): + pos, neg = "", "" + try: + raw = image.info.get("prompt") + if raw: + prompt_json = _safe_json_loads(raw) + if prompt_json: + pos, neg = _extract_params_from_prompt_json(prompt_json) + else: + pos = raw + except Exception as e: + print(f"[LoadImageFromFolderMXD] Prompt parse failed: {e}") + return pos, neg + + def load_image(self, image: str, source: str = "outputs", run_folder: bool = False): + image = image.lstrip("  ") + root = ( + folder_paths.get_input_directory() + if source == "inputs" + else folder_paths.get_output_directory() + ) + path = os.path.normpath(os.path.join(root, image)) if image else None + + if not path or not os.path.isfile(path): + raise FileNotFoundError(f"No such image: {path}") + + i = Image.open(path) + i = ImageOps.exif_transpose(i) + + pos, neg = self._extract_prompts(i) + + rgb = i.convert("RGB") + arr = np.array(rgb).astype(np.float32) / 255.0 + img_t = torch.from_numpy(arr)[None, ...] + + if 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask_t = 1.0 - torch.from_numpy(mask).unsqueeze(0) + else: + h, w = arr.shape[:2] + mask_t = torch.zeros((1, h, w), dtype=torch.float32) + + return (img_t, mask_t, pos, neg, os.path.basename(path)) + + @classmethod + def IS_CHANGED(cls, image, source="outputs", run_folder=False): + root = ( + folder_paths.get_input_directory() + if source == "inputs" + else folder_paths.get_output_directory() + ) + path = os.path.normpath(os.path.join(root, image.lstrip("  "))) if image else None + if not path or not os.path.isfile(path): + return "" + m = hashlib.sha256() + with open(path, "rb") as f: + m.update(f.read()) + return m.digest().hex() + + +class LoadVideoFromFolderMXD: + DESCRIPTION = ( + "Load a single video from any inputs/outputs subfolder. Turn on run_folder " + "to auto-queue every video in that same folder, one after another." + ) + TITLE = "Load Video (From Folder) MXD" + CATEGORY = "MXD/Video" + + RETURN_TYPES = ("VIDEO", "STRING") + RETURN_NAMES = ("VIDEO", "filename") + FUNCTION = "load_video" + + @classmethod + def INPUT_TYPES(cls): + union = _list_files_recursive_union( + folder_paths.get_output_directory(), folder_paths.get_input_directory(), VIDEO_BATCH_EXTS + ) + return { + "required": { + "source": (("outputs", "inputs"), {"default": "outputs"}), + "video": (tuple(union), ), + "run_folder": ("BOOLEAN", { + "default": False, + "tooltip": "When enabled, hitting Queue Prompt auto-queues every video in this file's folder, one after another, instead of just the selected file.", + }), + } + } + + def load_video(self, video: str, source: str = "outputs", run_folder: bool = False): + if not HAVE_COMFY_API_VIDEO: + raise RuntimeError( + "[LoadVideoFromFolderMXD] Video output requires a newer ComfyUI core with " + "comfy_api.latest / comfy_api.input_impl support. Please update ComfyUI." + ) + + video = video.lstrip("  ") + root = ( + folder_paths.get_input_directory() + if source == "inputs" + else folder_paths.get_output_directory() + ) + path = os.path.normpath(os.path.join(root, video)) if video else None + + if not path or not os.path.isfile(path): + raise FileNotFoundError(f"No such video: {path}") + + return (VideoFromFile(path), os.path.basename(path)) + + @classmethod + def IS_CHANGED(cls, video, source="outputs", run_folder=False): + root = ( + folder_paths.get_input_directory() + if source == "inputs" + else folder_paths.get_output_directory() + ) + path = os.path.normpath(os.path.join(root, video.lstrip("  "))) if video else None + if not path or not os.path.isfile(path): + return "" + try: + return str(os.path.getmtime(path)) + except OSError: + return "" + + +class LoadImageWithPromptsMXD: + DESCRIPTION = """Load one input image, create a mask from alpha, and read prompts if present.""" + CATEGORY = "image" + + RETURN_TYPES = ("IMAGE", "MASK", "STRING", "STRING") + RETURN_NAMES = ("IMAGE", "MASK", "positive", "negative") + FUNCTION = "load_image" + + @classmethod + def INPUT_TYPES(s): + input_dir = folder_paths.get_input_directory() + files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))] + files = folder_paths.filter_files_content_types(files, ["image"]) + files = _sort_paths_newest_first([os.path.join(input_dir, f) for f in files]) + files = [os.path.basename(f) for f in files] + return {"required": {"image": (files, {"image_upload": True})}} + + def _extract_prompts(self, img: Image.Image): + pos, neg = "", "" + raw = img.info.get("prompt") + if raw: + prompt_json = _safe_json_loads(raw) + if prompt_json: + pos, neg = _extract_params_from_prompt_json(prompt_json) + else: + pos = raw + return pos, neg + + def load_image(self, image): + image_path = folder_paths.get_annotated_filepath(image) + img = node_helpers.pillow(Image.open, image_path) + + output_images, output_masks = [], [] + pos, neg = "", "" + w, h = None, None + + excluded_formats = ['MPO'] + + for i in ImageSequence.Iterator(img): + i = node_helpers.pillow(ImageOps.exif_transpose, i) + + if i.mode == 'I': + i = i.point(lambda i: i * (1 / 255)) + frame = i.convert("RGB") + + if len(output_images) == 0: + w, h = frame.size + # extract prompts only once (from first frame) + pos, neg = self._extract_prompts(i) + + if frame.size != (w, h): + continue + + arr = np.array(frame).astype(np.float32) / 255.0 + tensor_img = torch.from_numpy(arr)[None, ...] + + if 'A' in i.getbands(): + mask = np.array(i.getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + elif i.mode == 'P' and 'transparency' in i.info: + mask = np.array(i.convert('RGBA').getchannel('A')).astype(np.float32) / 255.0 + mask = 1. - torch.from_numpy(mask) + else: + mask = torch.zeros((1, 64, 64), dtype=torch.float32, device="cpu") + + output_images.append(tensor_img) + output_masks.append(mask.unsqueeze(0)) + + if len(output_images) > 1 and img.format not in excluded_formats: + output_image = torch.cat(output_images, dim=0) + output_mask = torch.cat(output_masks, dim=0) + else: + output_image = output_images[0] + output_mask = output_masks[0] + + return (output_image, output_mask, pos, neg) + + @classmethod + def IS_CHANGED(s, image): + image_path = folder_paths.get_annotated_filepath(image) + m = hashlib.sha256() + with open(image_path, 'rb') as f: + m.update(f.read()) + return m.digest().hex() + + @classmethod + def VALIDATE_INPUTS(s, image): + if not folder_paths.exists_annotated_filepath(image): + return f"Invalid image file: {image}" + return True + +######################################################################################################################## + +class SaveImage_MXD: + TITLE = "Save Image MXD" + CATEGORY = "MXD/Image" + OUTPUT_NODE = True + FUNCTION = "save" + + DESCRIPTION = """Save images to the output folder or preview them.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "images": ("IMAGE", {"tooltip": "Images to preview and/or save."}), + "filename_prefix": ("STRING", { + "default": "ComfyUI", + "tooltip": "File name prefix. Tip: you can use a subfolder like 'tests/my_run'." + }), + "mode": ([ + "Save + Preview", + "Save Only", + "Preview only" + ], { + "default": "Save + Preview", + "tooltip": "Choose whether to write files to disk, only preview, or save quietly." + }), + }, + "optional": { + "embed_workflow": ("BOOLEAN", { + "default": True, + "tooltip": "Embed workflow metadata when saving PNG previews/files." + }), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + RETURN_TYPES = () + OUTPUT_TOOLTIPS = ("Saves and/or previews the images.",) + + @staticmethod + def _filtered_extra_pnginfo(extra_pnginfo, embed_workflow): + if embed_workflow or not isinstance(extra_pnginfo, dict): + return extra_pnginfo + filtered = {k: v for k, v in extra_pnginfo.items() if str(k).lower() != "workflow"} + return filtered or None + + def save(self, images, filename_prefix, mode, embed_workflow=True, prompt=None, extra_pnginfo=None): + if embed_workflow: + save_prompt = prompt + save_extra_pnginfo = self._filtered_extra_pnginfo(extra_pnginfo, True) + else: + # Core SaveImage embeds the hidden `prompt` graph too. + # Drop both to truly disable workflow reconstruction from saved files. + save_prompt = None + save_extra_pnginfo = None + + if mode.startswith("Preview"): + return PreviewImage().save_images(images, filename_prefix, save_prompt, save_extra_pnginfo) + result = SaveImage().save_images(images, filename_prefix, save_prompt, save_extra_pnginfo) + if mode == "Save Only" and isinstance(result, dict): + # Strip UI previews so nothing shows up in the ComfyUI viewer. + return {k: v for k, v in result.items() if k != "ui"} + return result + +######################################################################################################################## + +class ExtractWorkflowFromImageMXD: + TITLE = "Extract Workflow From Image MXD" + CATEGORY = "MXD/Image" + OUTPUT_NODE = True + FUNCTION = "extract_and_save" + + DESCRIPTION = """Save workflow metadata to a JSON file from a wired image execution context.""" + + def __init__(self): + self.output_dir = folder_paths.get_output_directory() + self.type = "output" + self.prefix_append = "" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE", {"tooltip": "Any connected image. Used to trigger extraction/save."}), + "filename_prefix": ("STRING", { + "default": "workflow/ComfyUI", + "tooltip": "Output JSON prefix. You can include subfolders, e.g. 'workflow/my_run'.", + }), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("json_path",) + OUTPUT_TOOLTIPS = ("Relative path to the saved JSON file in outputs.",) + + @staticmethod + def _decode_json_candidate(value): + if value is None: + return None + + if isinstance(value, (dict, list)): + return value + + if isinstance(value, bytes): + for enc in ("utf-8", "utf-16", "latin-1"): + try: + value = value.decode(enc) + break + except Exception: + continue + if isinstance(value, bytes): + value = value.decode("utf-8", "ignore") + + if not isinstance(value, str): + return None + + raw = value.strip() + if not raw: + return None + + if raw.lower().startswith("workflow:"): + raw = raw.split(":", 1)[1].strip() + + parsed = _safe_json_loads(raw) + if isinstance(parsed, (dict, list)): + return parsed + return None + + def _extract_workflow_from_context(self, prompt=None, extra_pnginfo=None): + if isinstance(extra_pnginfo, dict): + for key in ("workflow", "Workflow"): + parsed = self._decode_json_candidate(extra_pnginfo.get(key)) + if parsed is not None: + return parsed + + parsed_extra = self._decode_json_candidate(extra_pnginfo) + if isinstance(parsed_extra, dict): + for key in ("workflow", "Workflow"): + parsed = self._decode_json_candidate(parsed_extra.get(key)) + if parsed is not None: + return parsed + + if prompt is not None: + parsed_prompt = self._decode_json_candidate(prompt) + if parsed_prompt is not None: + return {"prompt": parsed_prompt} + if isinstance(prompt, dict): + return {"prompt": prompt} + + return None + + def extract_and_save(self, image, filename_prefix="workflow/ComfyUI", prompt=None, extra_pnginfo=None): + workflow = self._extract_workflow_from_context(prompt, extra_pnginfo) + if workflow is None: + raise ValueError( + "No workflow metadata is available in this execution context. " + "Connect generated images from the current run, or ensure workflow metadata is present." + ) + + filename_prefix += self.prefix_append + height = image[0].shape[0] + width = image[0].shape[1] + full_output_folder, filename, counter, subfolder, _ = folder_paths.get_save_image_path( + filename_prefix, self.output_dir, width, height + ) + os.makedirs(full_output_folder, exist_ok=True) + + file = f"{filename}_{counter:05}_.json" + save_path = os.path.join(full_output_folder, file) + + with open(save_path, "w", encoding="utf-8", newline="\n") as f: + json.dump(workflow, f, ensure_ascii=False, indent=2) + + rel = os.path.join(subfolder, file) if subfolder else file + rel = rel.replace("\\", "/") + return { + "ui": {"text": [f"Saved workflow JSON: {rel}"]}, + "result": (rel,), + } + +######################################################################################################################## + +NODE_CLASS_MAPPINGS = { + "Load Image Batch MXD": LoadImageBatchMXD, + "Load Video Batch MXD": LoadVideoBatchMXD, + "LoadImageFromFolderMXD": LoadImageFromFolderMXD, + "LoadVideoFromFolderMXD": LoadVideoFromFolderMXD, + "LoadImageWithPromptsMXD": LoadImageWithPromptsMXD, + "Save Image MXD": SaveImage_MXD, + "Extract Workflow From Image MXD": ExtractWorkflowFromImageMXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "Load Image Batch MXD": "Load Image Batch (Inputs/Outputs) MXD", + "Load Video Batch MXD": "Load Video Batch (Inputs/Outputs) MXD", + "LoadImageFromFolderMXD": "Load Image (From Folder) MXD", + "LoadVideoFromFolderMXD": "Load Video (From Folder) MXD", + "LoadImageWithPromptsMXD": "Load Image MXD", + "Save Image MXD": "Save Image MXD", + "Extract Workflow From Image MXD": "Extract Workflow From Image MXD", +} diff --git a/nodes/prompts.py b/nodes/prompts.py new file mode 100644 index 0000000..b861b64 --- /dev/null +++ b/nodes/prompts.py @@ -0,0 +1,217 @@ +from __future__ import annotations +import torch, comfy, math, node_helpers, comfy.model_management, comfy.utils +from comfy.comfy_types import IO, ComfyNodeABC, InputTypeDict +try: + from comfy_api.latest import io + HAVE_COMFY_API = True +except Exception as _e: + io = None + HAVE_COMFY_API = False + print(f"[ComfyUI-MaxedOut] comfy_api not available in prompts: {_e}") + +######################################################################################################################## +# Prompt with Guidance (Flux) +class PromptWithGuidance(ComfyNodeABC): + DESCRIPTION = """Encode text and apply Flux guidance in one node.""" + @classmethod + def INPUT_TYPES(cls) -> InputTypeDict: + return { + "required": { + "text": (IO.STRING, {"multiline": True, "dynamicPrompts": True}), + "clip": (IO.CLIP, {"tooltip": "The CLIP model used for encoding the text."}), + "guidance": ("FLOAT", {"default": 3.5, "min": 0.0, "max": 100.0, "step": 0.1}) + } + } + + RETURN_TYPES = (IO.CONDITIONING,) + FUNCTION = "encode_and_guide" + CATEGORY = "MXD/conditioning" + + def encode_and_guide(self, text, clip, guidance): + if clip is None: + raise RuntimeError("CLIP model is None. Your checkpoint may not contain a text encoder.") + + tokens = clip.tokenize(text) + conditioning = clip.encode_from_tokens_scheduled(tokens) + conditioning = node_helpers.conditioning_set_values(conditioning, {"guidance": guidance}) + return (conditioning,) + +######################################################################################################################## +if HAVE_COMFY_API: + class QwenImageEditSingleMXD(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="QwenImageEditSingleMXD", + display_name="Qwen Image Edit + Latent MXD", + category="MXD/conditioning", + description="Encode prompt/image and output a matching empty latent.", + inputs=[ + io.Clip.Input("clip"), + io.String.Input("prompt", multiline=True, dynamic_prompts=True), + io.Vae.Input("vae", optional=True), + io.Image.Input("image", optional=True), + io.Int.Input("batch_size", default=1, min=1, max=4096), + ], + outputs=[ + io.Conditioning.Output(), + io.Latent.Output(), # New Output + ], + ) + + @classmethod + def execute(cls, clip, prompt, vae=None, image=None, batch_size=1) -> io.NodeOutput: + ref_latents = [] + images_vl = [] + llama_template = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" + image_prompt = "" + + # Default fallback size if no image is provided (1024x1024) + final_width, final_height = 1024, 1024 + + if image is not None: + samples = image.movedim(-1, 1) + + # --- VISION SCALING (384px area) --- + total_vl = int(384 * 384) + scale_vl = math.sqrt(total_vl / (samples.shape[3] * samples.shape[2])) + width_vl = round(samples.shape[3] * scale_vl) + height_vl = round(samples.shape[2] * scale_vl) + + s_vl = comfy.utils.common_upscale(samples, width_vl, height_vl, "area", "disabled") + images_vl.append(s_vl.movedim(1, -1)) + + # --- LATENT/VAE SCALING (1024px area) --- + total_lat = int(1024 * 1024) + scale_lat = math.sqrt(total_lat / (samples.shape[3] * samples.shape[2])) + # Calculate final dimensions to be multiples of 8 + final_width = round(samples.shape[3] * scale_lat / 8.0) * 8 + final_height = round(samples.shape[2] * scale_lat / 8.0) * 8 + + if vae is not None: + s_lat = comfy.utils.common_upscale(samples, final_width, final_height, "area", "disabled") + ref_latents.append(vae.encode(s_lat.movedim(1, -1)[:, :, :, :3])) + + image_prompt += "Picture 1: <|vision_start|><|image_pad|><|vision_end|>" + + # 1. Generate the Empty Latent (SD3 Style: 16 channels, 1/8th resolution) + # This replaces the need for the separate EmptySD3LatentImage node + latent_tensor = torch.zeros( + [batch_size, 16, final_height // 8, final_width // 8], + device=comfy.model_management.intermediate_device() + ) + latent_output = {"samples": latent_tensor} + + # 2. Process Conditioning + tokens = clip.tokenize(image_prompt + prompt, images=images_vl, llama_template=llama_template) + conditioning = clip.encode_from_tokens_scheduled(tokens) + + if len(ref_latents) > 0: + conditioning = node_helpers.conditioning_set_values( + conditioning, + {"reference_latents": ref_latents}, + append=True, + ) + + return io.NodeOutput(conditioning, latent_output) + + ######################################################################################################################## + class QwenImageEditTripleMXD(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="QwenImageEditTripleMXD", + display_name="Qwen Image Edit Prompt MXD (Triple)", + category="advanced/conditioning", + inputs=[ + io.Clip.Input("clip"), + io.String.Input("prompt", multiline=True, dynamic_prompts=True), + io.Vae.Input("vae", optional=True), + io.Image.Input("image1", optional=True), + io.Image.Input("image2", optional=True), + io.Image.Input("image3", optional=True), + io.Int.Input("batch_size", default=1, min=1, max=4096), + ], + outputs=[ + io.Conditioning.Output(), + io.Latent.Output(), + ], + ) + + @classmethod + def execute(cls, clip, prompt, vae=None, image1=None, image2=None, image3=None, batch_size=1) -> io.NodeOutput: + ref_latents = [] + images = [image1, image2, image3] + images_vl = [] + llama_template = "<|im_start|>system\nDescribe the key features of the input image (color, shape, size, texture, objects, background), then explain how the user's text instruction should alter or modify the image. Generate a new image that meets the user's requirements while maintaining consistency with the original input where appropriate.<|im_end|>\n<|im_start|>user\n{}<|im_end|>\n<|im_start|>assistant\n" + image_prompt = "" + + # Default fallback + latent_width = 1024 + latent_height = 1024 + + for i, image in enumerate(images): + if image is not None: + samples = image.movedim(-1, 1) + + # 1. VL Model Scaling (LLM Vision) + total_vl = int(384 * 384) + scale_by_vl = math.sqrt(total_vl / (samples.shape[3] * samples.shape[2])) + width_vl = round(samples.shape[3] * scale_by_vl) + height_vl = round(samples.shape[2] * scale_by_vl) + s_vl = comfy.utils.common_upscale(samples, width_vl, height_vl, "area", "disabled") + images_vl.append(s_vl.movedim(1, -1)) + + # 2. VAE Scaling (Synchronized to 16-step for SD3 compatibility) + if vae is not None: + total_ref = int(1024 * 1024) + scale_by_ref = math.sqrt(total_ref / (samples.shape[3] * samples.shape[2])) + + # Pixels as multiple of 16 ensures Latent (Pixels/8) is always even + width_ref = round(samples.shape[3] * scale_by_ref / 16.0) * 16 + height_ref = round(samples.shape[2] * scale_by_ref / 16.0) * 16 + + if i == 0: + latent_width = width_ref + latent_height = height_ref + + s_ref = comfy.utils.common_upscale(samples, width_ref, height_ref, "area", "disabled") + ref_latents.append(vae.encode(s_ref.movedim(1, -1)[:, :, :, :3])) + + image_prompt += "Picture {}: <|vision_start|><|image_pad|><|vision_end|>".format(i + 1) + + # Process tokens and conditioning + tokens = clip.tokenize(image_prompt + prompt, images=images_vl, llama_template=llama_template) + conditioning = clip.encode_from_tokens_scheduled(tokens) + + if len(ref_latents) > 0: + conditioning = node_helpers.conditioning_set_values(conditioning, {"reference_latents": ref_latents}, append=True) + + # Create Output Latent + latent = torch.zeros([batch_size, 16, latent_height // 8, latent_width // 8], device=comfy.model_management.intermediate_device()) + + # FIXED: Return outputs positionally to match the schema defined above + # Output 1: Conditioning, Output 2: Latent Dictionary + return io.NodeOutput(conditioning, {"samples": latent}) + +######################################################################################################################## + +NODE_CLASS_MAPPINGS = { + "Prompt With Guidance (Flux)": PromptWithGuidance, +} + +if HAVE_COMFY_API: + NODE_CLASS_MAPPINGS.update({ + "QwenImageEditSingleMXD": QwenImageEditSingleMXD, + "QwenImageEditTripleMXD": QwenImageEditTripleMXD, + }) + +NODE_DISPLAY_NAME_MAPPINGS = { + "Prompt With Guidance (Flux)": "Prompt with Flux Guidance MXD", +} + +if HAVE_COMFY_API: + NODE_DISPLAY_NAME_MAPPINGS.update({ + "QwenImageEditSingleMXD": "Qwen Image Edit + Latent MXD", + "QwenImageEditTripleMXD": "Qwen Image Edit Prompt MXD (Triple)", + }) diff --git a/nodes/resolution.py b/nodes/resolution.py new file mode 100644 index 0000000..fe040e8 --- /dev/null +++ b/nodes/resolution.py @@ -0,0 +1,288 @@ +from __future__ import annotations +import math, comfy, comfy.utils, torch +from .latents import SdxlEmptyLatentImage + +######################################################################################################################## +# Image Scale To Total Pixels (SDXL Safe) +class SDXLImageScaleToTotalPixelsSafe: + DESCRIPTION = """Scale to a target megapixel count and keep aspect ratio. Skips SDXL-safe sizes.""" + upscale_methods = ["bilinear", "bicubic", "lanczos", "nearest-exact", "area"] + + # SDXL-safe resolutions (width, height) – store one orientation only, + # the code will check both (w, h) and (h, w) + SDXL_SAFE_RESOLUTIONS = [ + (1024, 1024), + (1152, 896), + (1216, 832), + (1344, 768), + (1536, 640), + ] + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "upscale_method": (cls.upscale_methods, {"default": "bilinear"}), + "total_megapixels": ( + "FLOAT", + { + "default": 1.0, + "min": 0.01, + "max": 128.0, + "step": 0.01, + "tooltip": "Set the total megapixels (e.g., 1.0 = 1 MP)", + }, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "upscale" + CATEGORY = "MXD/Upscaling" + + def upscale(self, image, upscale_method, total_megapixels): + if upscale_method in ["nearest-exact", "area"]: + raise Exception( + f"❌ '{upscale_method}' gives poor results.\n\n" + f"👉 Go to the Scale SDXL Image MXD node and switch to another like 'lanczos'.\n\n" + f"Node may be hidden behind KSampler." + ) + + b, h, w, c = image.shape + + # Skip scaling if the image already matches an SDXL-safe resolution + if (w, h) in self.SDXL_SAFE_RESOLUTIONS or (h, w) in self.SDXL_SAFE_RESOLUTIONS: + return (image,) + + # ComfyUI-native megapixel math + samples = image.movedim(-1, 1) + orig_h, orig_w = samples.shape[2], samples.shape[3] + + target_pixels = int(round(total_megapixels * 1024 * 1024)) + scale_by = math.sqrt(target_pixels / (orig_w * orig_h)) + + new_w = max(1, round(orig_w * scale_by)) + new_h = max(1, round(orig_h * scale_by)) + + scaled = comfy.utils.common_upscale(samples, new_w, new_h, upscale_method, "disabled") + scaled = scaled.movedim(1, -1) + return (scaled,) + +######################################################################################################################## +# Flux Image Scale To Total Pixels (Flux Safe) +class FluxImageScaleToTotalPixelsSafe: + DESCRIPTION = """Scale to a target megapixel count and keep aspect ratio. Skips Flux-safe sizes.""" + upscale_methods = ["bilinear", "bicubic", "lanczos", "nearest-exact", "area"] + + # Flux-safe resolutions (width, height) – stored in one orientation only + FLUX_SAFE_RESOLUTIONS = [ + (1408, 1408), + (1728, 1152), + (1664, 1216), + (1920, 1088), + (2176, 960), + (1024, 1024), + (1216, 832), + (1152, 896), + (1344, 768), + (1536, 640), + (320, 320), + (384, 256), + (448, 320), + (448, 256), + (576, 256), + ] + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "upscale_method": (cls.upscale_methods, {"default": "bilinear"}), + "total_megapixels": ( + "FLOAT", + { + "default": 1.0, + "min": 0.01, + "max": 128.0, + "step": 0.01, + "tooltip": "Set the total megapixels (e.g., 1.0 = 1 MP)", + }, + ), + } + } + + RETURN_TYPES = ("IMAGE",) + FUNCTION = "upscale" + CATEGORY = "MXD/Upscaling" + + def upscale(self, image, upscale_method, total_megapixels): + if upscale_method in ["nearest-exact", "area"]: + raise Exception( + f"❌ '{upscale_method}' gives poor results.\n\n" + f"👉 Go to the Scale Flux Image MXD node and switch to another like 'lanczos'.\n\n" + f"Node may be hidden behind KSampler." + ) + + b, h, w, c = image.shape + + # Skip scaling if image matches any Flux-safe resolution + if (w, h) in self.FLUX_SAFE_RESOLUTIONS or (h, w) in self.FLUX_SAFE_RESOLUTIONS: + return (image,) + + samples = image.movedim(-1, 1) + orig_h, orig_w = samples.shape[2], samples.shape[3] + + target_pixels = int(round(total_megapixels * 1024 * 1024)) + scale_by = math.sqrt(target_pixels / (orig_w * orig_h)) + + new_w = max(1, round(orig_w * scale_by)) + new_h = max(1, round(orig_h * scale_by)) + + scaled = comfy.utils.common_upscale(samples, new_w, new_h, upscale_method, "disabled") + scaled = scaled.movedim(1, -1) + return (scaled,) + +######################################################################################################################## +class FluxResolutionMatcher: + DESCRIPTION = """Match the closest Flux resolution and orientation for the input image.""" + CATEGORY = "MXD/Latent" + FUNCTION = "match_resolution" + RETURN_NAMES = ("resolution", "vertical") + + # Full set kept for compatibility (enum list must match FluxEmptyLatentImage) + RESOLUTIONS = { + "— High Resolutions —": None, + "Square (1:1) 1408x1408": (1408, 1408), + "Standard (4:3) 1664x1216": (1664, 1216), + "Landscape (3:2) 1728x1152": (1728, 1152), + "Widescreen (16:9) 1920x1088": (1920, 1088), + "Ultrawide (21:9) 2176x960": (2176, 960), + + "— Standard Resolutions —": None, + "Square (1:1) 1024x1024": (1024, 1024), + "Standard (4:3) 1152x896": (1152, 896), + "Landscape (3:2) 1216x832": (1216, 832), + "Widescreen (16:9) 1344x768": (1344, 768), + "Ultrawide (21:9) 1536x640": (1536, 640), + + "— Low Resolutions —": None, + "Square (1:1) 320x320": (320, 320), + "Standard (4:3) 448x320": (448, 320), + "Landscape (3:2) 384x256": (384, 256), + "Widescreen (16:9) 448x256": (448, 256), + "Ultrawide (21:9) 576x256": (576, 256), + } + + # Keep same enum type so it connects to FluxEmptyLatentImage + RETURN_TYPES = (list(RESOLUTIONS.keys()), "BOOLEAN") + + # Precompute aspect ratio groups (only for standard resolutions) + ASPECT_RATIO_GROUPS = {} + for res_str, dims in RESOLUTIONS.items(): + if dims is None: + continue + # ✅ Skip high and low groups for logic + if "High" in res_str or "Low" in res_str: + continue + group_name = " ".join(res_str.split(' ')[:-1]) + if group_name not in ASPECT_RATIO_GROUPS: + w, h = dims + ratio = w / h + ASPECT_RATIO_GROUPS[group_name] = {'ratio': ratio, 'resolutions': []} + ASPECT_RATIO_GROUPS[group_name]['resolutions'].append(res_str) + + @classmethod + def INPUT_TYPES(cls): + return {"required": {"image": ("IMAGE",)}} + + def match_resolution(self, image: torch.Tensor): + if image.dim() < 4 or image.shape[1] < 1 or image.shape[2] < 1: + print("Warning: Invalid image tensor received. Falling back to default resolution.") + return ("Square (1:1) 1024x1024", False) + + _batch, height, width, _channels = image.shape + is_vertical = height > width + img_aspect_ratio = (height / width) if is_vertical else (width / height) + img_area = height * width + + best_ar_group_name = min( + self.ASPECT_RATIO_GROUPS.keys(), + key=lambda name: abs(img_aspect_ratio - self.ASPECT_RATIO_GROUPS[name]['ratio']) + ) + + candidate_res_strings = self.ASPECT_RATIO_GROUPS[best_ar_group_name]['resolutions'] + + best_res_string = min( + candidate_res_strings, + key=lambda res_str: abs(img_area - (self.RESOLUTIONS[res_str][0] * self.RESOLUTIONS[res_str][1])) + ) + + return (best_res_string, is_vertical) +######################################################################################################################## + +class SDXLResolutionMatcher: + DESCRIPTION = """Match the closest SDXL resolution and orientation for the input image.""" + CATEGORY = "MXD/Latent" + FUNCTION = "match_resolution" + RETURN_NAMES = ("resolution", "vertical") + + # Use the exact same enum list as SdxlEmptyLatentImage + RESOLUTIONS = SdxlEmptyLatentImage.RESOLUTIONS + + RETURN_TYPES = (list(RESOLUTIONS.keys()), "BOOLEAN") + + ASPECT_RATIO_GROUPS = {} + for res_str, dims in RESOLUTIONS.items(): + if dims is None: + continue + group_name = " ".join(res_str.split(" ")[:-1]) + if group_name not in ASPECT_RATIO_GROUPS: + w, h = dims + ratio = w / h + ASPECT_RATIO_GROUPS[group_name] = {"ratio": ratio, "resolutions": []} + ASPECT_RATIO_GROUPS[group_name]["resolutions"].append(res_str) + + @classmethod + def INPUT_TYPES(cls): + return {"required": {"image": ("IMAGE",)}} + + def match_resolution(self, image: torch.Tensor): + if image.dim() < 4 or image.shape[1] < 1 or image.shape[2] < 1: + print("Warning: Invalid image tensor received. Falling back to default resolution.") + return ("Square (1:1) 1024x1024", False) + + _batch, height, width, _channels = image.shape + is_vertical = height > width + img_aspect_ratio = (height / width) if is_vertical else (width / height) + img_area = height * width + + best_ar_group_name = min( + self.ASPECT_RATIO_GROUPS.keys(), + key=lambda name: abs(img_aspect_ratio - self.ASPECT_RATIO_GROUPS[name]["ratio"]) + ) + + candidate_res_strings = self.ASPECT_RATIO_GROUPS[best_ar_group_name]["resolutions"] + + best_res_string = min( + candidate_res_strings, + key=lambda res_str: abs(img_area - (self.RESOLUTIONS[res_str][0] * self.RESOLUTIONS[res_str][1])) + ) + + return (best_res_string, is_vertical) +######################################################################################################################## + +NODE_CLASS_MAPPINGS = { + "Image Scale To Total Pixels (SDXL Safe)": SDXLImageScaleToTotalPixelsSafe, + "Flux Image Scale To Total Pixels (Flux Safe)": FluxImageScaleToTotalPixelsSafe, + "FluxResolutionMatcher": FluxResolutionMatcher, + "SDXLResolutionMatcher": SDXLResolutionMatcher, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "Image Scale To Total Pixels (SDXL Safe)": "Scale SDXL Image MXD", + "Flux Image Scale To Total Pixels (Flux Safe)": "Scale Flux Image MXD", + "FluxResolutionMatcher": "Flux Resolution Matcher MXD", + "SDXLResolutionMatcher": "SDXL Resolution Matcher MXD", +} diff --git a/nodes/wan22/__init__.py b/nodes/wan22/__init__.py new file mode 100644 index 0000000..1b4d1da --- /dev/null +++ b/nodes/wan22/__init__.py @@ -0,0 +1,19 @@ +"""WAN 2.2 node package: buckets/scalers, latent save-load, I2V conditioning, video ops.""" +import importlib + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +for _name in ( + "buckets", + "latent_io", + "i2v", + "video_ops", +): + try: + _mod = importlib.import_module(f".{_name}", __name__) + except Exception as e: + print(f"[ComfyUI-MaxedOut] Failed to import 'nodes.wan22.{_name}': {e}") + continue + NODE_CLASS_MAPPINGS.update(getattr(_mod, "NODE_CLASS_MAPPINGS", {}) or {}) + NODE_DISPLAY_NAME_MAPPINGS.update(getattr(_mod, "NODE_DISPLAY_NAME_MAPPINGS", {}) or {}) diff --git a/nodes/wan22/buckets.py b/nodes/wan22/buckets.py new file mode 100644 index 0000000..1d83097 --- /dev/null +++ b/nodes/wan22/buckets.py @@ -0,0 +1,594 @@ +"""WAN 2.2 resolution buckets: empty latents, image scaler, resolution matcher, outpaint pad. + +Registered nodes: + Wan2_2EmptyLatentImageMXD Wan 2.2 Empty Latent Image MXD + wan22EmptyHunyuanLatentVideoMXD WAN2.2 Empty Latent Video MXD + WAN22_I2V_Image_Scaler_MXD Image Scaler Wan 2.2 I2V MXD + WAN22_I2V_Match_Resolution_MXD Match Resolution Wan 2.2 I2V MXD + PadImageForOutpaintingMXD Pad Image for Outpainting MXD + +Canonical WAN 2.2 buckets: 480p tier 832x480 / 480x832 / 624x624, 720p tier +1280x720 / 720x1280 / 1024x1024. All scaling keeps dimensions 16-aligned. +""" +from __future__ import annotations +from typing import Tuple + +import torch + +import comfy.utils +import comfy.model_management +import nodes + + +# ---- Canonical WAN 2.2 buckets ---- +SQUARE_TOL = 0.03 # exact-ish square passthrough tolerance +AUTO_SQUARE_MAX_AR = 1.25 # Auto may crop to square when the source is within 25% of 1:1. + +def _ar(w, h): + return w / max(1, h) + +def _safe_hw(w, h): + w = max(16, min(w, nodes.MAX_RESOLUTION)) + h = max(16, min(h, nodes.MAX_RESOLUTION)) + return w, h + +def _floor16(x): + x = int(x) // 16 * 16 + return max(16, x) + +def _ceil16(x): + x = (int(x) + 15) // 16 * 16 + return max(16, x) + +def _is_squareish(w, h, tol=SQUARE_TOL): + r = _ar(w, h) + return abs(r - 1.0) <= tol + +def _is_auto_square_candidate(w, h): + r = _ar(w, h) + return max(r, 1.0 / max(r, 1e-9)) <= AUTO_SQUARE_MAX_AR + +def _wan22_tier_from_area(iw, ih): + area = iw * ih + area_480 = 832 * 480 + area_720 = 1280 * 720 + return "480p" if abs(area - area_480) / area_480 <= abs(area - area_720) / area_720 else "720p" + +def _wan22_square_bucket(tier, iw=None, ih=None): + if tier == "720p": + return (1024, 1024) + if tier == "480p": + return (624, 624) + return (1024, 1024) if _wan22_tier_from_area(iw, ih) == "720p" else (624, 624) + +def _wan22_oriented_bucket(tier, orientation, iw=None, ih=None): + if tier == "Auto": + tier = _wan22_tier_from_area(iw, ih) + if orientation == "Tall": + return (480, 832) if tier == "480p" else (720, 1280) + if orientation == "Wide": + return (832, 480) if tier == "480p" else (1280, 720) + return _wan22_square_bucket(tier, iw, ih) + +def _closest_bucket(img_w, img_h, bucket_list, cover=False): + """ + Pick the best (bw,bh) from bucket_list for this image. + Uses scale closeness + AR diff to rank. + """ + in_ar = _ar(img_w, img_h) + best, best_key = None, (float("inf"), 0.0) + for bw, bh in bucket_list: + s = max(bw/img_w, bh/img_h) if cover else min(bw/img_w, bh/img_h) + ar_diff = abs(_ar(bw, bh) - in_ar) + key = (abs(1.0 - s), ar_diff) + if key < best_key: + best_key, best = key, (bw, bh) + return best + +def _resize_then_center_crop(img, out_w, out_h): + """ + Resize to cover target (ensures >= target on both sides after ceil16), + then center-crop. No padding. + """ + t, ih, iw, c = img.shape + s = max(out_w / iw, out_h / ih) + tw = _ceil16(iw * s) + th = _ceil16(ih * s) + tmp = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) + y0 = max(0, (th - out_h) // 2) + x0 = max(0, (tw - out_w) // 2) + return tmp[:, y0:y0+out_h, x0:x0+out_w, :] + +def _resize_fit_inside(img, out_w, out_h): + """ + Resize to fit inside target (ensures <= target on both sides via floor16), + and return the resized tensor only. No padding. + """ + t, ih, iw, c = img.shape + s = min(out_w / iw, out_h / ih) + tw = _floor16(iw * s) + th = _floor16(ih * s) + tw, th = _safe_hw(tw, th) + resized = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) + return resized, tw, th + +def _validate_image_batch_4d(image, node_name, input_name): + if image is None: + raise ValueError(f"[{node_name}] '{input_name}' is required.") + if not torch.is_tensor(image): + raise TypeError(f"[{node_name}] '{input_name}' must be an IMAGE torch tensor, got {type(image).__name__}.") + if image.ndim != 4: + raise ValueError(f"[{node_name}] '{input_name}' must have shape [T,H,W,C], got {tuple(image.shape)}.") + if image.shape[0] <= 0: + raise ValueError(f"[{node_name}] '{input_name}' contains zero images/frames.") + if image.shape[1] <= 0 or image.shape[2] <= 0 or image.shape[3] <= 0: + raise ValueError(f"[{node_name}] '{input_name}' has invalid dimensions {tuple(image.shape)}.") + return image + +def _resize_to_explicit_resolution(img, out_w, out_h, match_mode="crop_to_match"): + """ + Resize IMAGE batch to an explicit resolution. + - crop_to_match: cover + center crop (exact output) + - fit_inside_only: preserve AR, no crop (may be smaller) + - stretch_exact: force exact output (distorts AR) + """ + out_w = int(out_w) + out_h = int(out_h) + if out_w <= 0 or out_h <= 0: + raise ValueError(f"Invalid target resolution {out_w}x{out_h}.") + + if match_mode == "crop_to_match": + return _resize_then_center_crop(img, out_w, out_h) + + if match_mode == "fit_inside_only": + _, ih, iw, _ = img.shape + s = min(out_w / max(1, iw), out_h / max(1, ih)) + tw = max(1, min(out_w, int(iw * s))) + th = max(1, min(out_h, int(ih * s))) + return comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) + + if match_mode == "stretch_exact": + return comfy.utils.common_upscale(img.movedim(-1, 1), out_w, out_h, "bilinear", "center").movedim(1, -1) + + raise ValueError( + f"Invalid match_mode '{match_mode}'. Expected one of: crop_to_match, fit_inside_only, stretch_exact." + ) + +_WAN22_VALID_RES = { + (832, 480), (480, 832), + (1280, 720), (720, 1280), + (624, 624), (1024, 1024), +} + +def _wan22_is_valid_dim(w, h): + return (w, h) in _WAN22_VALID_RES + + +def _wan22_pick_bucket(iw, ih, tier, crop_to_fit, aspect_mode="Auto"): + if tier == "Safe Auto": + tier = "Auto" + + if aspect_mode in ("Tall", "Wide", "Square"): + return _wan22_oriented_bucket(tier, aspect_mode, iw, ih) + + is_squareish = _is_squareish(iw, ih) + is_landscape = iw >= ih + + # --- Square handling --- + if is_squareish or (crop_to_fit and _is_auto_square_candidate(iw, ih)): + return _wan22_square_bucket(tier, iw, ih) + + # --- 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, aspect_mode="Auto"): + """ + 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 1024x1024\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, aspect_mode=aspect_mode) + 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 + + +# ---------- Empty latent image generator (for video nodes) ---------- +class Wan2_2EmptyLatentImageMXD: + """ + Utility node for WAN 2.2 workflows. + Generates an empty latent tensor at common video-friendly resolutions. + """ + + DESCRIPTION = """Create an empty WAN 2.2 latent at a preset resolution.""" + TITLE = "WAN2.2 Empty Latent Image" + CATEGORY = "WAN2.2/Latent" + + RESOLUTIONS = { + "— 720p —": None, + "Widescreen (16:9) 1280×720": (1280, 720), + + "— 480p —": None, + "Widescreen (16:9) 832×480": (832, 480), + "Square (1:1) 624×624": (624, 624), + } + + RETURN_TYPES = ("LATENT",) + FUNCTION = "generate" + + @classmethod + def INPUT_TYPES(cls): + options = list(cls.RESOLUTIONS.keys()) + return { + "required": { + "resolution": ( + options, + {"default": "Square (1:1) 960×960", "tooltip": "Select target resolution preset."} + ), + "vertical": ( + "BOOLEAN", + {"default": False, "label_on": "Vertical", "label_off": "Landscape", + "tooltip": "Swap width/height for vertical orientation."} + ), + "batch_size": ( + "INT", + {"default": 1, "min": 1, "max": 4096, "tooltip": "Number of latents to generate."} + ), + } + } + + def generate(self, resolution, vertical, batch_size): + size = self.RESOLUTIONS.get(resolution) + if size is None: + raise ValueError(f"'{resolution}' is a header or invalid option.") + + w, h = size + if vertical: + w, h = h, w + + # Safety: ensure divisible by 8 + if (w % 8) or (h % 8): + raise ValueError(f"Resolution must be divisible by 8. Got {w}x{h}.") + + # WAN video length always t=1 + t = 1 + + latent = torch.zeros( + [batch_size, 16, t, h // 8, w // 8], + device=comfy.model_management.intermediate_device() + ) + return ({"samples": latent},) + +# ---------- Empty latent video generator with presets (for video nodes) ---------- +class wan22EmptyHunyuanLatentVideoMXD: + """ + Exactly like core EmptyHunyuanLatentVideo, but width/height are replaced + with valid WAN 2.2 resolution presets and a vertical toggle. + """ + + RETURN_TYPES = ("LATENT",) + FUNCTION = "generate" + CATEGORY = "latent/video" + + RESOLUTIONS = { + "— 720p —": None, + "Widescreen (16:9) 1280×720": (1280, 720), + "Square (1:1) 1024×1024": (1024, 1024), + + "— 480p —": None, + "Widescreen (16:9) 832×480": (832, 480), + "Square (1:1) 624×624": (624, 624), + } + + @classmethod + def INPUT_TYPES(cls): + options = list(cls.RESOLUTIONS.keys()) + return { + "required": { + "resolution": ( + options, + {"default": "Widescreen (16:9) 832×480"} + ), + "vertical": ( + "BOOLEAN", + {"default": False, "label_on": "Vertical", "label_off": "Landscape"} + ), + "length": ( + "INT", + {"default": 81, "min": 1, "max": nodes.MAX_RESOLUTION, "step": 4} + ), + "batch_size": ( + "INT", + {"default": 1, "min": 1, "max": 4096} + ), + } + } + + def generate(self, resolution, vertical, length, batch_size): + size = self.RESOLUTIONS.get(resolution) + if size is None: + raise ValueError(f"'{resolution}' is not a selectable resolution.") + w, h = size + if vertical: + w, h = h, w + + # identical to core behavior: + t = ((length - 1) // 4) + 1 + latent = torch.zeros( + [batch_size, 16, t, h // 8, w // 8], + device=comfy.model_management.intermediate_device() + ) + return ({"samples": latent},) + + +class WAN22_I2V_Image_Scaler_MXD: + """ + MXD Image Scaler for WAN 2.2 (NO PADDING) + - Modes: Auto / 480p / 720p (legacy "Safe Auto" still accepted) + - Fit (no pad): proportional resize ≤ target; returns resized dims. + - Crop (no pad): resize-to-cover then center-crop to exact target. + - Square handling: + * Auto & 480p: ~square → 624×624 + * 720p: ~square -> 1024x1024 + - “Safe Auto”: + * If input is already a valid WAN 2.2 bucket, passthrough. + * If input is far outside 480p–720p range, error early. + * Otherwise, same logic as Auto. + * Perfect for video-extend workflows. + """ + + TITLE = "Image Bucket Scaler MXD (No Pad)" + CATEGORY = "image/processing" + RETURN_TYPES = ("IMAGE",) + FUNCTION = "scale" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "tier": (["Auto", "480p", "720p"], {"default": "Auto"}), + "crop_to_fit": ("BOOLEAN", { + "default": True, + "label_on": "Perfect Fit (Crops Edges)", + "label_off": "Closest Fit (No Crop)" + }), + "aspect_mode": (["Auto", "Tall", "Wide", "Square"], { + "default": "Auto", + "tooltip": "Auto picks wide/tall/square from the source. Use Square/Tall/Wide to force the target bucket shape." + }), + } + } + + def scale(self, image, tier="Auto", crop_to_fit=False, aspect_mode="Auto"): + # Keep legacy "Safe Auto" values from old workflows working, but expose only one Auto in UI. + internal_tier = "Safe Auto" if tier == "Auto" else tier + out, _, _, _ = _wan22_scale_image_core( + image, + tier=internal_tier, + crop_to_fit=crop_to_fit, + aspect_mode=aspect_mode, + ) + return (out,) + +class WAN22_I2V_Match_Resolution_MXD: + """ + Match a second image (or image batch) to a reference image resolution for WAN 2.2 + first/last-frame workflows. + """ + TITLE = "WAN 2.2 I2V Match Resolution" + CATEGORY = "image/processing" + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("matched_image",) + FUNCTION = "match_resolution" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "reference_image": ("IMAGE", { + "tooltip": "Reference size source (usually the first image after WAN bucket scaling)." + }), + "image_to_match": ("IMAGE", { + "tooltip": "Image or batch to resize using the reference image resolution." + }), + "match_mode": (["crop_to_match", "fit_inside_only", "stretch_exact"], { + "default": "crop_to_match", + "tooltip": "crop_to_match = exact size via cover+center crop; fit_inside_only = no crop, may be smaller; stretch_exact = exact size with distortion." + }), + "enforce_wan_bucket": ("BOOLEAN", { + "default": False, + "label_on": "Validate WAN Bucket", + "label_off": "No WAN Validation", + "tooltip": "If enabled, reference_image must already be a WAN 2.2 bucket size." + }), + } + } + + def match_resolution(self, reference_image, image_to_match, match_mode="crop_to_match", enforce_wan_bucket=False): + node_name = "WAN22_I2V_Match_Resolution_MXD" + reference_image = _validate_image_batch_4d(reference_image, node_name, "reference_image") + image_to_match = _validate_image_batch_4d(image_to_match, node_name, "image_to_match") + + _, ref_h, ref_w, _ = reference_image.shape + + if enforce_wan_bucket and not _wan22_is_valid_dim(ref_w, ref_h): + raise ValueError( + f"[{node_name}] Reference image resolution {ref_w}x{ref_h} is not a valid WAN 2.2 bucket.\n" + "Valid WAN 2.2 buckets are:\n" + " - 832x480 / 480x832\n" + " - 1280x720 / 720x1280\n" + " - 624x624 / 1024x1024\n\n" + "Recommended workflow:\n" + " 1. Scale the first image with 'Image Scaler Wan 2.2 I2V MXD'\n" + " 2. Use this node to match the second image to the scaled first image" + ) + + matched = _resize_to_explicit_resolution( + image_to_match, + out_w=ref_w, + out_h=ref_h, + match_mode=match_mode, + ) + return (matched,) + + +class PadImageForOutpaintingMXD: + SEARCH_ALIASES = ["extend canvas", "expand image", "outpaint pad"] + + RETURN_TYPES = ("IMAGE", "MASK") + FUNCTION = "expand_image" + CATEGORY = "image/transform" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "image": ("IMAGE",), + "left": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), + "top": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), + "right": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), + "bottom": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), + "round_to": (["None", "2", "8", "16", "32", "64"], {"default": "16"}), + } + } + + @staticmethod + def _nearest_multiple(value: int, multiple: int, padded: bool) -> int: + if multiple <= 1 or value % multiple == 0: + return value + lower = (value // multiple) * multiple + upper = lower + multiple + if lower <= 0: + return upper + if not padded: + return lower + return lower if value - lower <= upper - value else upper + + @staticmethod + def _axis_plan(size: int, before: int, after: int, multiple: int) -> Tuple[int, int, int, int, int]: + target = size + before + after + if multiple > 1: + target = PadImageForOutpaintingMXD._nearest_multiple(target, multiple, before + after > 0) + + delta = target - (size + before + after) + if delta < 0: + remove = -delta + from_after = min(after, remove) + after -= from_after + remove -= from_after + from_before = min(before, remove) + before -= from_before + remove -= from_before + crop_before = remove // 2 + crop_after = remove - crop_before + else: + crop_before = 0 + crop_after = 0 + if before > 0 and after > 0: + add_before = delta // 2 + before += add_before + after += delta - add_before + elif before > 0: + before += delta + else: + after += delta + + final_size = size - crop_before - crop_after + before + after + if final_size <= 0: + raise ValueError("[PadImageForOutpaintingMXD] Rounding removed the full image on one axis.") + return before, after, crop_before, crop_after, final_size + + def expand_image(self, image, left, top, right, bottom, round_to="16"): + image = _validate_image_batch_4d(image, "PadImageForOutpaintingMXD", "image") + batch, height, width, channels = image.size() + multiple = 1 if round_to == "None" else int(round_to) + + left, right, crop_left, crop_right, final_width = self._axis_plan(width, left, right, multiple) + top, bottom, crop_top, crop_bottom, final_height = self._axis_plan(height, top, bottom, multiple) + + cropped = image[:, crop_top:height - crop_bottom, crop_left:width - crop_right, :] + crop_height = cropped.shape[1] + crop_width = cropped.shape[2] + + new_image = torch.full( + (batch, final_height, final_width, channels), + 0.5, + dtype=image.dtype, + device=image.device, + ) + new_image[:, top:top + crop_height, left:left + crop_width, :] = cropped + + mask = torch.ones( + (final_height, final_width), + dtype=torch.float32, + device=image.device, + ) + mask[top:top + crop_height, left:left + crop_width] = 0.0 + + return (new_image, mask.unsqueeze(0)) + + +NODE_CLASS_MAPPINGS = { + "Wan2_2EmptyLatentImageMXD": Wan2_2EmptyLatentImageMXD, + "wan22EmptyHunyuanLatentVideoMXD": wan22EmptyHunyuanLatentVideoMXD, + "WAN22_I2V_Image_Scaler_MXD": WAN22_I2V_Image_Scaler_MXD, + "WAN22_I2V_Match_Resolution_MXD": WAN22_I2V_Match_Resolution_MXD, + "PadImageForOutpaintingMXD": PadImageForOutpaintingMXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "Wan2_2EmptyLatentImageMXD": "Wan 2.2 Empty Latent Image MXD", + "wan22EmptyHunyuanLatentVideoMXD": "WAN2.2 Empty Latent Video MXD", + "WAN22_I2V_Image_Scaler_MXD": "Image Scaler Wan 2.2 I2V MXD", + "WAN22_I2V_Match_Resolution_MXD": "Match Resolution Wan 2.2 I2V MXD", + "PadImageForOutpaintingMXD": "Pad Image for Outpainting MXD", +} diff --git a/nodes/wan22/i2v.py b/nodes/wan22/i2v.py new file mode 100644 index 0000000..c76fc34 --- /dev/null +++ b/nodes/wan22/i2v.py @@ -0,0 +1,298 @@ +"""WAN 2.2 image-to-video conditioning nodes (all require comfy_api; skipped when absent). + +Registered nodes (only when HAVE_COMFY_API): + Wan22ImageToVideoMXD Wan 2.2 Image to Video MXD + WAN22_I2V_Video_Prep_MXD WAN 2.2 Video Prep I2V MXD + Wan22FirstLastImageToVideoMXD Wan 2.2 I2V First & Last Frame MXD + +These expect pre-sized inputs (use the buckets.py scaler upstream); they do no +scaling or CLIP-vision of their own. +""" +from __future__ import annotations + +import torch + +import comfy.model_management +import node_helpers, nodes + +# Comfy API +try: + from comfy_api.latest import io + from comfy_api.input_impl import VideoFromComponents + from comfy_api.util import VideoComponents + HAVE_COMFY_API = True +except Exception as _e: + io = None + VideoFromComponents = None + VideoComponents = None + HAVE_COMFY_API = False + print(f"[ComfyUI-MaxedOut] comfy_api not available in wan22.i2v: {_e}") + +from .buckets import _wan22_scale_image_core + + +def _resample_video_frames_to_fps(frames, in_fps, out_fps): + """ + Resample a frame sequence to a target FPS using nearest-frame selection. + Preserves clip duration approximately by dropping/duplicating frames, + instead of only changing FPS metadata (which changes playback speed). + Returns (frames_out, fps_out, changed). + """ + if frames is None or frames.ndim != 4: + raise ValueError("Expected frame tensor with shape [T,H,W,C].") + + if in_fps is None: + raise ValueError("Input video FPS is missing; cannot force FPS safely.") + + in_fps = float(in_fps) + out_fps = float(out_fps) + if in_fps <= 0: + raise ValueError(f"Invalid input FPS: {in_fps}") + if out_fps <= 0: + raise ValueError(f"Invalid target FPS: {out_fps}") + + if frames.shape[0] <= 1: + return frames, float(out_fps), False + + if abs(in_fps - out_fps) < 1e-6: + return frames, float(out_fps), False + + n_in = int(frames.shape[0]) + # Match the first/last frame span, then pick nearest frames on that timeline. + n_out = max(1, int(round(((n_in - 1) * out_fps) / in_fps)) + 1) + if n_out == n_in: + # Frame count may stay the same for near-equal FPS; metadata still becomes exact. + return frames, float(out_fps), False + + idx = torch.linspace(0, n_in - 1, steps=n_out, device=frames.device) + idx = idx.round().to(dtype=torch.long) + out = frames.index_select(0, idx) + return out, float(out_fps), True + + +# ---------- WAN 2.2 Image to Video (no scaling; expects pre-sized input) ---------- +if HAVE_COMFY_API: + class Wan22ImageToVideoMXD(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="Wan22ImageToVideoMXD", + display_name="WAN 2.2 Image to Video MXD", + category="conditioning/video_models", + description="WAN 2.2 image to video without scaling or CLIP vision.", + inputs=[ + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Vae.Input("vae"), + io.Int.Input("length", default=81, min=1, max=16384, step=4), + io.Int.Input("batch_size", default=1, min=1, max=4096), + io.Image.Input("start_image", optional=False), + ], + 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) -> io.NodeOutput: + if start_image is None: + raise ValueError("start_image must be provided (already pre-sized).") + + frames_in, ih, iw, ch = start_image.shape + frames_used = min(frames_in, length) + t = ((length - 1) // 4) + 1 + + latent = torch.zeros( + [batch_size, 16, t, ih // 8, iw // 8], + device=comfy.model_management.intermediate_device() + ) + + # create placeholder image tensor + image = torch.ones( + (length, ih, iw, ch), + device=start_image.device, + dtype=start_image.dtype + ) * 0.5 + image[:frames_used] = start_image[:frames_used] + + # encode using VAE + concat_latent_image = vae.encode(image[:, :, :, :3]) + + # mask zeros out the frames used + mask = torch.ones( + (1, 1, t, concat_latent_image.shape[-2], concat_latent_image.shape[-1]), + device=image.device, + dtype=image.dtype + ) + mask[:, :, :((frames_used - 1) // 4) + 1] = 0.0 + + positive = node_helpers.conditioning_set_values( + positive, {"concat_latent_image": concat_latent_image, "concat_mask": mask} + ) + negative = node_helpers.conditioning_set_values( + negative, {"concat_latent_image": concat_latent_image, "concat_mask": mask} + ) + + out_latent = {"samples": latent} + return io.NodeOutput(positive, negative, out_latent) + + class WAN22_I2V_Video_Prep_MXD: + """ + Prepare a source video for iterative WAN 2.2 extension: + - scale entire video using WAN bucket logic + - output the scaled frame batch directly + - keep default workflow simple for common use + """ + CATEGORY = "MXD/video" + FUNCTION = "prepare" + RETURN_TYPES = ("VIDEO", "IMAGE", "FLOAT") + RETURN_NAMES = ("scaled_video", "images", "fps") + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "video": ("VIDEO",), + "tier": (["Auto", "480p", "720p"], {"default": "Auto"}), + "crop_to_fit": ("BOOLEAN", { + "default": True, + "label_on": "Perfect Fit (Crops Edges)", + "label_off": "Closest Fit (No Crop)" + }), + "force_fps": ("BOOLEAN", { + "default": False, + "label_on": "Force FPS", + "label_off": "Keep Source FPS", + "tooltip": "When enabled, resample frames (drop/duplicate) and set exact target fps." + }), + "target_fps": ("INT", { + "default": 16, + "min": 1, + "max": 1000, + "step": 1, + "tooltip": "Used when Force FPS is enabled. Output video fps will be set exactly to this value." + }), + "aspect_mode": (["Auto", "Tall", "Wide", "Square"], { + "default": "Auto", + "tooltip": "Auto picks wide/tall/square from the source. Use Square/Tall/Wide to force the target bucket shape." + }), + }, + } + + def prepare(self, video, tier="Auto", crop_to_fit=True, force_fps=False, target_fps=16, aspect_mode="Auto"): + 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.") + + out_frame_rate = float(comp.frame_rate) if comp.frame_rate is not None else None + if force_fps: + frames, out_frame_rate, _ = _resample_video_frames_to_fps( + frames, comp.frame_rate, target_fps + ) + + # "Auto" in video prep uses the safer extend-friendly behavior. + # Keep accepting legacy "Safe Auto" values from older saved workflows. + internal_tier = "Safe Auto" if tier == "Auto" else tier + scaled_frames, _, _, _ = _wan22_scale_image_core( + frames, + tier=internal_tier, + crop_to_fit=crop_to_fit, + aspect_mode=aspect_mode, + ) + + scaled_video = VideoFromComponents( + VideoComponents( + images=scaled_frames, + audio=comp.audio, + frame_rate=out_frame_rate, + ) + ) + + fps = float(out_frame_rate) if out_frame_rate is not None else 0.0 + return (scaled_video, scaled_frames, fps) + + class Wan22FirstLastImageToVideoMXD(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="Wan22FirstLastImageToVideoMXD", + display_name="WAN 2.2 First & Last I2V MXD", + category="conditioning/video_models", + inputs=[ + io.Conditioning.Input("positive"), + io.Conditioning.Input("negative"), + io.Vae.Input("vae"), + io.Int.Input("length", default=81, min=1, max=nodes.MAX_RESOLUTION, 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), + ], + 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) -> io.NodeOutput: + spacial_scale = vae.spacial_compression_encode() + + # Assume incoming images are already pre-sized by upstream nodes. + height, width = start_image.shape[1], start_image.shape[2] if start_image is not None else (vae.latent_channels * spacial_scale, vae.latent_channels * spacial_scale) + + latent = torch.zeros( + [batch_size, vae.latent_channels, ((length - 1) // 4) + 1, height // spacial_scale, width // spacial_scale], + device=comfy.model_management.intermediate_device() + ) + + image = torch.ones((length, height, width, 3)) * 0.5 + mask = torch.ones((1, 1, latent.shape[2] * 4, latent.shape[-2], latent.shape[-1])) + + if start_image is not None: + image[:start_image.shape[0]] = start_image + mask[:, :, :start_image.shape[0] + 3] = 0.0 + + if end_image is not None: + image[-end_image.shape[0]:] = end_image + mask[:, :, -end_image.shape[0]:] = 0.0 + + concat_latent_image = vae.encode(image[:, :, :, :3]) + mask = mask.view(1, mask.shape[2] // 4, 4, mask.shape[3], mask.shape[4]).transpose(1, 2) + + positive = node_helpers.conditioning_set_values(positive, {"concat_latent_image": concat_latent_image, "concat_mask": mask}) + negative = node_helpers.conditioning_set_values(negative, {"concat_latent_image": concat_latent_image, "concat_mask": mask}) + + out_latent = {"samples": latent} + return io.NodeOutput(positive, negative, out_latent) + + +NODE_CLASS_MAPPINGS = {} +NODE_DISPLAY_NAME_MAPPINGS = {} + +if HAVE_COMFY_API: + NODE_CLASS_MAPPINGS.update({ + "Wan22ImageToVideoMXD": Wan22ImageToVideoMXD, + "WAN22_I2V_Video_Prep_MXD": WAN22_I2V_Video_Prep_MXD, + "Wan22FirstLastImageToVideoMXD": Wan22FirstLastImageToVideoMXD, + }) + 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", + "Wan22FirstLastImageToVideoMXD": "Wan 2.2 I2V First & Last Frame MXD", + }) diff --git a/nodes/wan22/latent_io.py b/nodes/wan22/latent_io.py new file mode 100644 index 0000000..2e9a658 --- /dev/null +++ b/nodes/wan22/latent_io.py @@ -0,0 +1,1457 @@ +"""WAN 2.2 latent save/load: two-stage I2V handoff via .latent + .cond.pt sidecar files. + +Registered nodes: + SaveLatent_I2V_MXD Save Latent MXD + LoadLatent_I2V_MXD Load Latent MXD + LoadLatents_FromFolder_I2V_MXD Load Latent Batch MXD + LoadLatent_I2V_Pipe_MXD Load Latent Pipe MXD + LoadLatents_FromFolder_I2V_Pipe_MXD Load Latent Batch Pipe MXD + LatentPipeUnpack_MXD Unpack Latent Pipe MXD + +Route: GET /mxd/latents/files (fresh re-scan for the run_folder queuing loop). + +Latents are stored under /latents. Each .latent embeds the source +workflow + KSampler params in safetensors metadata; conditioning goes in a +.cond.pt sidecar. The loaders reconstruct sampler settings from that metadata +so the second stage can resume with matching parameters. +""" +from __future__ import annotations +import os, re, glob, json, hashlib, copy +from collections import deque +from typing import Any, Dict, Tuple, Optional, List, Union + +import torch +from safetensors import safe_open + +import folder_paths +import comfy.utils +from comfy.cli_args import args +from nodes import KSamplerAdvanced + +from server import PromptServer +from aiohttp import web + +routes = PromptServer.instance.routes + + +def _sort_paths_newest_first(paths: List[str]) -> List[str]: + """Sort file paths by mtime desc (newest first), stable by normalized path.""" + def _mtime(path: str) -> float: + try: + return os.path.getmtime(path) + except OSError: + return 0.0 + + return sorted( + paths, + key=lambda p: (-_mtime(p), p.replace("\\", "/").lower()), + ) + + +def _sort_latent_options_by_folder(options: List[str], root: str = "") -> List[str]: + """ + Order relative '.latent' option paths so that the combo's prev/next arrows + stay confined to one folder before moving on, newest first: + folderA/file1, folderA/file2, ..., folderB/file1, ... + Folders are ordered by the mtime of their most recently modified file (so + a folder that just received a new file jumps back to the top), and files + within each folder are newest first. + """ + def _mtime(rel: str) -> float: + if not root: + return 0.0 + try: + return os.path.getmtime(os.path.join(root, rel)) + except OSError: + return 0.0 + + folder_of = lambda rel: rel.replace("\\", "/").rsplit("/", 1)[0] if "/" in rel.replace("\\", "/") else "" + + folder_latest: Dict[str, float] = {} + for rel in options: + folder = folder_of(rel) + m = _mtime(rel) + if m > folder_latest.get(folder, -1.0): + folder_latest[folder] = m + + def key(rel: str): + folder = folder_of(rel) + return (-folder_latest.get(folder, 0.0), -_mtime(rel)) + + return sorted(options, key=key) + +def _list_latent_subfolders(latents_root: str) -> List[str]: + """ + List latent subfolders recursively (e.g. "a", "a/b"), newest first by + latest latent mtime in each branch. + """ + files = glob.glob(os.path.join(latents_root, "**", "*.latent"), recursive=True) + if not files: + return [] + + folder_latest_mtime: Dict[str, float] = {} + for file_path in files: + rel_dir = os.path.relpath(os.path.dirname(file_path), latents_root).replace(os.sep, "/").strip("/") + if not rel_dir or rel_dir == ".": + continue + try: + mtime = os.path.getmtime(file_path) + except OSError: + mtime = 0.0 + + # Include each ancestor so both "a" and "a/b" appear as options. + parts = [p for p in rel_dir.split("/") if p] + for i in range(1, len(parts) + 1): + branch = "/".join(parts[:i]) + prev = folder_latest_mtime.get(branch, -1.0) + if mtime > prev: + folder_latest_mtime[branch] = mtime + + return [ + folder + for folder, _ in sorted( + folder_latest_mtime.items(), + key=lambda kv: (-kv[1], kv[0].lower()), + ) + ] + + +@routes.get("/mxd/latents/files") +async def mxd_list_latent_files(request): + """ + Fresh re-scan of input/latents for .latent files. Used by the run_folder + queuing loop (refresh_before_run) to pick up files a still-running + workflow is writing concurrently, instead of relying on the dropdown + list captured whenever the node's combo was last populated. + """ + latents_root = os.path.join(folder_paths.get_input_directory(), "latents") + os.makedirs(latents_root, exist_ok=True) + files = glob.glob(os.path.join(latents_root, "**", "*.latent"), recursive=True) + options = [os.path.relpath(f, latents_root).replace(os.sep, "/") for f in files] + options = _sort_latent_options_by_folder(options, latents_root) + return web.json_response(options) + + +# ---------- SaveLatent (saves latent + conditioning + optional trim_latent) ---------- +class SaveLatent_I2V_MXD: + """ + Default latent saver, works for t2v, i2v, and VACE 2.2. Persists: + • latent tensor -> .latent + • pos/neg CONDITIONING -> .cond.pt + • optional trim_latent value (VACE 2.2) -> .cond.pt + """ + TITLE = "Save Latent MXD" + CATEGORY = "MXD/Latents" + OUTPUT_NODE = True + RETURN_TYPES = () + FUNCTION = "save_only" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "samples": ("LATENT", {"tooltip": "Latent to save."}), + "positive": ("CONDITIONING", {"tooltip": "Positive CONDITIONING to save alongside the latent."}), + "negative": ("CONDITIONING", {"tooltip": "Negative CONDITIONING to save alongside the latent."}), + "filename_prefix": ("STRING", {"default": "ComfyUI", "tooltip": "Prefix for saved files"}), + }, + "optional": { + "trim_latent": ("INT", { + "default": 0, + "min": 0, + "max": 10000, + "step": 1, + "tooltip": "VACE 2.2 trim_latent value to preserve with this latent. Usually 0 or 1. Ignored for t2v/i2v." + }), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO", "unique_id": "UNIQUE_ID"}, + } + + def save_only(self, samples, positive, negative, filename_prefix="ComfyUI", + trim_latent=0, prompt=None, extra_pnginfo=None, unique_id=None): + _save_i2v_latent_bundle( + samples=samples, + positive=positive, + negative=negative, + filename_prefix=filename_prefix, + prompt=prompt, + extra_pnginfo=extra_pnginfo, + unique_id=unique_id, + sidecar_extra={"trim_latent": _coerce_trim_latent(trim_latent)}, + ) + return {} + +# ---------- Helpers ---------- +def _load_latent_file(latent_path: str) -> Tuple[Dict[str, torch.Tensor], Dict[str, Any], List[str]]: + """ + Load safetensors latent with Comfy metadata. + Returns (samples_dict, metadata_dict, keys_list) + """ + with safe_open(latent_path, framework="pt", device="cpu") as f: + keys = list(f.keys()) + # prefer explicit key we write + if "latent_tensor" in keys: + t = f.get_tensor("latent_tensor").float().contiguous() + else: + # fall back (some variants might save using a different name) + first = keys[0] + t = f.get_tensor(first).float().contiguous() + + meta = f.metadata() or {} + + # if ancient format, rescale (match Comfy behavior) + if "latent_format_version_0" not in keys: + t = t * (1.0 / 0.18215) + + return {"samples": t}, meta, keys + + +def _safe_json_loads(s: Union[str, bytes, None]) -> Optional[Dict[str, Any]]: + if s is None: + return None + if isinstance(s, bytes): + try: + s = s.decode("utf-8", "ignore") + except Exception: + return None + if not isinstance(s, str): + return None + try: + return json.loads(s) + except Exception: + # sometimes double-encoded in metadata + try: + return json.loads(json.loads(s)) + except Exception: + return None + + +def _workflow_node_bbox(nodes: List[Dict[str, Any]]) -> Optional[Tuple[float, float, float, float]]: + """(min_x, min_y, max_x, max_y) over a litegraph 'nodes' list. None if no positions found.""" + xs0, ys0, xs1, ys1 = [], [], [], [] + for n in nodes: + if not isinstance(n, dict): + continue + pos = n.get("pos") + if isinstance(pos, list) and len(pos) >= 2: + x, y = pos[0], pos[1] + elif isinstance(pos, dict): + x, y = pos.get("0", 0), pos.get("1", 0) + else: + continue + size = n.get("size") + w = size[0] if isinstance(size, (list, tuple)) and len(size) >= 1 else 200 + h = size[1] if isinstance(size, (list, tuple)) and len(size) >= 2 else 100 + xs0.append(x); ys0.append(y); xs1.append(x + w); ys1.append(y + h) + if not xs0: + return None + return (min(xs0), min(ys0), max(xs1), max(ys1)) + + +def _offset_workflow_nodes(nodes: List[Dict[str, Any]], dx: float, dy: float) -> None: + for n in nodes: + if not isinstance(n, dict): + continue + pos = n.get("pos") + if isinstance(pos, list) and len(pos) >= 2: + pos[0] = pos[0] + dx + pos[1] = pos[1] + dy + elif isinstance(pos, dict): + if "0" in pos: pos["0"] = pos["0"] + dx + if "1" in pos: pos["1"] = pos["1"] + dy + + +def _merge_prior_workflow_into_current(prior_workflow_json: Optional[str], current_workflow: Any) -> Any: + """ + Merge a previously-saved workflow graph (embedded in a loaded .latent file) into the + workflow graph of the run that's currently saving. The prior graph's nodes/links/groups + are copied in with fresh ids and shifted to sit to the left of the current graph, wrapped + in a labelled group - so dragging the final video into ComfyUI shows both stages at once, + the same as if you'd copy/pasted the first workflow onto the second one's canvas. + + Best-effort: on any parse/shape problem, returns current_workflow untouched. + """ + if not prior_workflow_json or not isinstance(current_workflow, dict): + return current_workflow + + try: + prior = json.loads(prior_workflow_json) if isinstance(prior_workflow_json, str) else prior_workflow_json + if not isinstance(prior, dict): + return current_workflow + + prior_nodes = prior.get("nodes") + if not isinstance(prior_nodes, list) or not prior_nodes: + return current_workflow + + merged = copy.deepcopy(current_workflow) + current_nodes = merged.get("nodes") + if not isinstance(current_nodes, list): + current_nodes = [] + merged["nodes"] = current_nodes + + prior_nodes = copy.deepcopy(prior_nodes) + prior_links = copy.deepcopy(prior.get("links")) if isinstance(prior.get("links"), list) else [] + prior_groups = copy.deepcopy(prior.get("groups")) if isinstance(prior.get("groups"), list) else [] + + # ---- remap node ids so they can't collide with the current graph ---- + current_last_node_id = merged.get("last_node_id") + if not isinstance(current_last_node_id, int): + current_last_node_id = max((n.get("id", 0) for n in current_nodes if isinstance(n, dict)), default=0) + next_node_id = current_last_node_id + 1 + node_id_map: Dict[Any, int] = {} + for n in prior_nodes: + if not isinstance(n, dict) or "id" not in n: + continue + node_id_map[n["id"]] = next_node_id + n["id"] = next_node_id + next_node_id += 1 + + # ---- remap link ids the same way ---- + current_last_link_id = merged.get("last_link_id") + if not isinstance(current_last_link_id, int): + current_last_link_id = max( + (l[0] for l in (merged.get("links") or []) if isinstance(l, list) and l), default=0 + ) + next_link_id = current_last_link_id + 1 + link_id_map: Dict[Any, int] = {} + for l in prior_links: + if isinstance(l, list) and l: + link_id_map[l[0]] = next_link_id + next_link_id += 1 + + for n in prior_nodes: + if not isinstance(n, dict): + continue + for inp in (n.get("inputs") or []): + if isinstance(inp, dict) and inp.get("link") is not None: + inp["link"] = link_id_map.get(inp["link"], inp["link"]) + for out in (n.get("outputs") or []): + if isinstance(out, dict) and isinstance(out.get("links"), list): + out["links"] = [link_id_map.get(x, x) for x in out["links"]] + + remapped_links = [] + for l in prior_links: + if not isinstance(l, list) or len(l) < 5: + continue + new_l = list(l) + new_l[0] = link_id_map.get(l[0], l[0]) + new_l[1] = node_id_map.get(l[1], l[1]) + new_l[3] = node_id_map.get(l[3], l[3]) + remapped_links.append(new_l) + + # ---- shift the prior graph so it sits to the left of the current one ---- + current_bbox = _workflow_node_bbox(current_nodes) + prior_bbox = _workflow_node_bbox(prior_nodes) + margin = 400 + if current_bbox and prior_bbox: + dx = (current_bbox[0] - margin) - prior_bbox[2] + dy = current_bbox[1] - prior_bbox[1] + else: + dx, dy = 0, 0 + _offset_workflow_nodes(prior_nodes, dx, dy) + for g in prior_groups: + if not isinstance(g, dict): + continue + b = g.get("bounding") + if isinstance(b, list) and len(b) >= 2: + b[0] = b[0] + dx + b[1] = b[1] + dy + + # wrap the prior graph in a labelled group so it's obvious what it is + wrapper_group = None + prior_bbox_shifted = _workflow_node_bbox(prior_nodes) + if prior_bbox_shifted: + pad = 60 + wrapper_group = { + "title": "Prior stage (loaded latent's source workflow)", + "bounding": [ + prior_bbox_shifted[0] - pad, + prior_bbox_shifted[1] - pad - 40, + (prior_bbox_shifted[2] - prior_bbox_shifted[0]) + pad * 2, + (prior_bbox_shifted[3] - prior_bbox_shifted[1]) + pad * 2 + 40, + ], + "color": "#3f789e", + "font_size": 24, + } + + merged["nodes"] = current_nodes + prior_nodes + merged["links"] = (merged.get("links") or []) + remapped_links + groups = list(merged.get("groups") or []) + prior_groups + if wrapper_group: + groups.append(wrapper_group) + merged["groups"] = groups + merged["last_node_id"] = next_node_id - 1 + merged["last_link_id"] = next_link_id - 1 + return merged + except Exception as e: + print(f"[SaveVideoMXD] Could not merge prior stage workflow into embedded metadata: {e}") + return current_workflow + + +def _node_sort_key(node_id: str) -> Tuple[int, Union[int, str]]: + s = str(node_id) + try: + return (0, int(s)) + except Exception: + return (1, s) + + +def _normalize_prompt_graph(prompt_json: Any) -> Dict[str, Any]: + if not isinstance(prompt_json, dict): + return {} + graph = prompt_json.get("prompt", prompt_json) + return graph if isinstance(graph, dict) else {} + + +def _get_graph_node(graph: Dict[str, Any], node_id: Any) -> Optional[Dict[str, Any]]: + if node_id is None or not isinstance(graph, dict): + return None + node = graph.get(str(node_id)) + return node if isinstance(node, dict) else None + + +def _iter_graph_nodes_sorted(graph: Dict[str, Any]) -> List[Tuple[str, Dict[str, Any]]]: + nodes: List[Tuple[str, Dict[str, Any]]] = [] + for node_id, node in graph.items(): + if isinstance(node, dict): + nodes.append((str(node_id), node)) + nodes.sort(key=lambda pair: _node_sort_key(pair[0])) + return nodes + + +def _linked_node_id(value: Any) -> Optional[str]: + if isinstance(value, (list, tuple)) and len(value) >= 1: + return str(value[0]) + return None + + +def _is_ksampler_node(node: Any) -> bool: + if not isinstance(node, dict): + return False + return "KSampler" in str(node.get("class_type", "")) + + +def _collect_upstream_linked_node_ids(node: Dict[str, Any]) -> List[str]: + inputs = node.get("inputs", {}) + if not isinstance(inputs, dict): + return [] + + seen = set() + ordered = [] + + # Prefer latent-carrying links first. + for key in ("samples", "latent", "latent_image"): + linked = _linked_node_id(inputs.get(key)) + if linked is not None and linked not in seen: + seen.add(linked) + ordered.append(linked) + + # Then search all other connected inputs in stable order. + for key, value in inputs.items(): + if key in ("samples", "latent", "latent_image"): + continue + linked = _linked_node_id(value) + if linked is not None and linked not in seen: + seen.add(linked) + ordered.append(linked) + + return ordered + + +def _find_upstream_ksampler_node_id(graph: Dict[str, Any], start_node_id: Any) -> Optional[str]: + if not isinstance(graph, dict) or start_node_id is None: + return None + + queue: deque[str] = deque([str(start_node_id)]) + visited = set() + + while queue: + node_id = queue.popleft() + if node_id in visited: + continue + visited.add(node_id) + + node = _get_graph_node(graph, node_id) + if not node: + continue + if _is_ksampler_node(node): + return node_id + + for upstream_id in _collect_upstream_linked_node_ids(node): + if upstream_id not in visited: + queue.append(upstream_id) + + return None + + +def _extract_ksampler_params(node: Dict[str, Any]) -> Dict[str, Any]: + inputs = node.get("inputs", {}) if isinstance(node, dict) else {} + if not isinstance(inputs, dict): + inputs = {} + + out: Dict[str, Any] = {} + + def set_int(key: str): + if key in inputs: + try: + out[key] = int(inputs[key]) + except Exception: + pass + + def set_float(key: str): + if key in inputs: + try: + out[key] = float(inputs[key]) + except Exception: + pass + + def set_str(key: str): + if key in inputs and not isinstance(inputs[key], (list, tuple, dict)): + try: + out[key] = str(inputs[key]).strip() + except Exception: + pass + + set_int("steps") + set_float("cfg") + set_str("sampler_name") + set_str("scheduler") + set_int("start_at_step") + set_int("end_at_step") + + return out + + +def _attach_source_ksampler_metadata(meta: Dict[str, Any], prompt: Any, unique_id: Any) -> None: + if not isinstance(meta, dict): + return + + graph = _normalize_prompt_graph(prompt) + if not graph: + return + + save_node_id = str(unique_id) if unique_id is not None else "" + if not save_node_id: + return + + save_node = _get_graph_node(graph, save_node_id) + if not save_node: + return + + source_candidates = _collect_upstream_linked_node_ids(save_node) + if not source_candidates: + return + + source_ksampler_id = None + for start_id in source_candidates: + source_ksampler_id = _find_upstream_ksampler_node_id(graph, start_id) + if source_ksampler_id: + break + + if not source_ksampler_id: + return + + source_node = _get_graph_node(graph, source_ksampler_id) + if not source_node: + return + + meta["mxd_source_save_node_id"] = save_node_id + meta["mxd_source_ksampler_node_id"] = source_ksampler_id + try: + meta["mxd_source_ksampler_params"] = json.dumps(_extract_ksampler_params(source_node)) + except Exception: + pass + + +def _build_latent_metadata(prompt=None, extra_pnginfo=None, unique_id=None, extra_meta=None): + if args.disable_metadata: + return None + + meta = {} + if prompt is not None: + try: + meta["prompt"] = json.dumps(prompt) + except Exception: + pass + if extra_pnginfo is not None: + for k, v in extra_pnginfo.items(): + try: + meta[k] = json.dumps(v) + except Exception: + pass + if isinstance(extra_meta, dict): + for k, v in extra_meta.items(): + try: + meta[str(k)] = json.dumps(v) + except Exception: + pass + _attach_source_ksampler_metadata(meta, prompt, unique_id) + return meta + + +def _save_i2v_latent_bundle( + samples, + positive, + negative, + filename_prefix="I2V", + prompt=None, + extra_pnginfo=None, + unique_id=None, + sidecar_extra=None, +): + latents_dir = os.path.join(folder_paths.get_input_directory(), "latents") + os.makedirs(latents_dir, exist_ok=True) + + full_output_folder, filename, counter, _subfolder, _filename_prefix = folder_paths.get_save_image_path( + filename_prefix, latents_dir + ) + + extra_meta = sidecar_extra if isinstance(sidecar_extra, dict) else None + meta = _build_latent_metadata( + prompt=prompt, + extra_pnginfo=extra_pnginfo, + unique_id=unique_id, + extra_meta=extra_meta, + ) + + latent_path = os.path.join(full_output_folder, f"{filename}_{counter:05}_.latent") + payload = { + "latent_tensor": samples["samples"].contiguous(), + "latent_format_version_0": torch.tensor([]), + } + comfy.utils.save_torch_file(payload, latent_path, metadata=meta) + + sidecar = {"positive": positive, "negative": negative} + if isinstance(sidecar_extra, dict): + sidecar.update(sidecar_extra) + torch.save(sidecar, latent_path.replace(".latent", ".cond.pt")) + return latent_path + + +def _load_i2v_conditioning_sidecar(latent_path): + cond_path = latent_path.replace(".latent", ".cond.pt") + if not os.path.exists(cond_path): + return [], [], {} + + try: + data = torch.load(cond_path, map_location="cpu") + except Exception: + return [], [], {} + + if not isinstance(data, dict): + return [], [], {} + + return data.get("positive", []), data.get("negative", []), data + + +def _coerce_trim_latent(value, default=0): + try: + if isinstance(value, str): + parsed = _safe_json_loads(value) + value = parsed if parsed is not None else value + return int(value) + except Exception: + return int(default) + + +def _extract_prompt_text_from_ksampler(graph: Dict[str, Any], ks_node: Dict[str, Any]) -> Tuple[str, str]: + pos = "" + neg = "" + + inputs = ks_node.get("inputs", {}) if isinstance(ks_node, dict) else {} + if not isinstance(inputs, dict): + return pos, neg + + def _text_from_clip(link_value: Any) -> str: + node_id = _linked_node_id(link_value) + if node_id is None: + return "" + node = _get_graph_node(graph, node_id) or {} + if node.get("class_type") == "CLIPTextEncode": + return str(node.get("inputs", {}).get("text", "")).strip() + return "" + + pos = _text_from_clip(inputs.get("positive")) + neg = _text_from_clip(inputs.get("negative")) + return pos, neg + + +def _extract_params_from_prompt_json( + prompt_json: Dict[str, Any], + meta: Optional[Dict[str, Any]] = None, +) -> Tuple[str, str, int, float, str, str, int]: + """ + Returns: (positive, negative, steps, cfg, sampler_name, scheduler, end_at_step) + parsed from the saved Comfy prompt graph with deterministic KSampler selection. + """ + pos = "" + neg = "" + steps = 20 + cfg = 8.0 + sampler_name = "" + scheduler = "" + end_at_step = 0 + + graph = _normalize_prompt_graph(prompt_json) + if not isinstance(graph, dict): + return pos, neg, steps, cfg, sampler_name, scheduler, end_at_step + + ks_node = None + extracted_params: Dict[str, Any] = {} + + # 1) Source KSampler id saved directly in latent metadata. + if isinstance(meta, dict): + raw_ks = meta.get("mxd_source_ksampler_node_id") + if raw_ks is not None: + candidate = _get_graph_node(graph, str(raw_ks)) + if candidate and _is_ksampler_node(candidate): + ks_node = candidate + + # 2) Source save node id -> trace upstream to nearest KSampler. + if ks_node is None and isinstance(meta, dict): + raw_save = meta.get("mxd_source_save_node_id") + if raw_save is not None: + save_node = _get_graph_node(graph, str(raw_save)) + if save_node: + for start_id in _collect_upstream_linked_node_ids(save_node): + trace_id = _find_upstream_ksampler_node_id(graph, start_id) + if trace_id: + candidate = _get_graph_node(graph, trace_id) + if candidate and _is_ksampler_node(candidate): + ks_node = candidate + break + + # 3) Legacy fallback: last KSampler node in graph. + if ks_node is None: + for _, node in _iter_graph_nodes_sorted(graph): + if _is_ksampler_node(node): + ks_node = node + + if not ks_node: + return pos, neg, steps, cfg, sampler_name, scheduler, end_at_step + + pos, neg = _extract_prompt_text_from_ksampler(graph, ks_node) + extracted_params = _extract_ksampler_params(ks_node) + + if "steps" in extracted_params: + steps = int(extracted_params["steps"]) + if "cfg" in extracted_params: + cfg = float(extracted_params["cfg"]) + if "end_at_step" in extracted_params: + end_at_step = int(extracted_params["end_at_step"]) + if "sampler_name" in extracted_params: + sampler_name = str(extracted_params["sampler_name"]).strip() + if "scheduler" in extracted_params: + scheduler = str(extracted_params["scheduler"]).strip() + + # Fallback to saved parameter snapshot if graph parse is incomplete. + if isinstance(meta, dict): + saved_params = _safe_json_loads(meta.get("mxd_source_ksampler_params")) + if isinstance(saved_params, dict): + if "steps" in saved_params and "steps" not in extracted_params: + try: + steps = int(saved_params["steps"]) + except Exception: + pass + if "cfg" in saved_params and "cfg" not in extracted_params: + try: + cfg = float(saved_params["cfg"]) + except Exception: + pass + if "end_at_step" in saved_params and "end_at_step" not in extracted_params: + try: + end_at_step = int(saved_params["end_at_step"]) + except Exception: + pass + if "sampler_name" in saved_params and "sampler_name" not in extracted_params: + try: + sampler_name = str(saved_params["sampler_name"]).strip() + except Exception: + pass + if "scheduler" in saved_params and "scheduler" not in extracted_params: + try: + scheduler = str(saved_params["scheduler"]).strip() + except Exception: + pass + + return pos, neg, steps, cfg, sampler_name, scheduler, end_at_step + +# ---------- Load one latent (conditioning + sampler params + optional trim_latent) ---------- +class LoadLatent_I2V_MXD: + """ + Default single-latent loader: sampler settings, CONDITIONING (positive/negative), and + an optional trim_latent value, all read from the .latent file and its .cond.pt sidecar. + """ + DESCRIPTION = """Load one latent and return conditioning and sampler settings.""" + TITLE = "Load Latent MXD" + CATEGORY = "MXD/Latents" + FUNCTION = "load" + + RETURN_TYPES = ( + "FLOAT", # shift + "CONDITIONING", # positive conditioning + "CONDITIONING", # negative conditioning + "LATENT", + "INT", + "FLOAT", + "STRING", + "STRING", + "INT", + "STRING", + "INT", # trim_latent + "STRING", # high_workflow + ) + RETURN_NAMES = ( + "shift", + "positive", + "negative", + "samples", + "steps", + "cfg", + "sampler_name", + "scheduler", + "end_at_step", + "filename_prefix", + "trim_latent", + "high_workflow", + ) + + @classmethod + def INPUT_TYPES(s): + latents_root = os.path.join(folder_paths.get_input_directory(), "latents") + os.makedirs(latents_root, exist_ok=True) + files = glob.glob(os.path.join(latents_root, "**", "*.latent"), recursive=True) + # Clean dropdown display (no "latents/" prefix) + options = [os.path.relpath(f, latents_root).replace(os.sep, "/") for f in files] + # Group by folder so the combo's prev/next arrows walk one folder at a time. + options = _sort_latent_options_by_folder(options, latents_root) + + ks_inputs = KSamplerAdvanced.INPUT_TYPES().get("required", {}) + samplers_enum = ks_inputs.get("sampler_name", ("STRING",))[0] + schedulers_enum = ks_inputs.get("scheduler", ("STRING",))[0] + + s.RETURN_TYPES = ( + "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", + "INT", "FLOAT", samplers_enum, schedulers_enum, + "INT", "STRING", "INT", "STRING", + ) + s._SAMPLERS_ENUM = samplers_enum + s._SCHEDULERS_ENUM = schedulers_enum + + return { + "required": { + "latent": (options, ), + "run_folder": ("BOOLEAN", { + "default": False, + "tooltip": "When enabled, hitting Queue Prompt auto-queues every latent in this file's folder, one after another, instead of just the selected file.", + }), + "refresh_before_run": ("BOOLEAN", { + "default": False, + "tooltip": "When run_folder is on, re-scan the latents folder for new files right before the queuing loop starts, instead of using the dropdown list as of whenever it was last populated. Use this when another workflow is still writing latents into this folder as you queue this one.", + }), + } + } + + def _coerce_enum(self, value, enum_values): + try: + return value if (enum_values and value in enum_values) else (enum_values[0] if enum_values else value) + except Exception: + return value + + def _strip_counter(self, name: str) -> str: + # Only strip the trailing pattern we generate when saving: "_<5digits>_" + # Preserve numeric-only base names like "96". + stem, _ = os.path.splitext(name) + m = re.match(r"^(.*?)(?:_\d{5}_)$", stem) + return m.group(1) if m else stem + + def _extract_sd3_shift(self, meta: dict, prompt_json: dict | None) -> float: + """ + Find SD3 'shift' in several places: + 1) flat meta["shift"] + 2) nested in prompt/workflow JSON: + - nodes[].{type|class_type} == "ModelSamplingSD3" -> inputs.shift or widgets_values[0] + - runtime-style prompt dict mapping IDs -> {..., class_type: "ModelSamplingSD3"} + Falls back to 5.0 if not found. + """ + def try_float(x): + try: + return float(x) + except Exception: + return None + + # 1) flat meta + if isinstance(meta, dict): + v = try_float(meta.get("shift")) + if v is not None: + return v + + # parse any JSON-like strings present in meta + def safe_load(x): + try: + return _safe_json_loads(x) if isinstance(x, str) else x + except Exception: + return None + + # Search helper over various JSON shapes + def search_container(obj): + # Direct dict containing shift + if isinstance(obj, dict): + if "shift" in obj: + v = try_float(obj.get("shift")) + if v is not None: + return v + + # Comfy "nodes": [ {...}, ... ] + nodes = obj.get("nodes") + if isinstance(nodes, list): + # take the last SD3 node (most recent in graph) + ms_nodes = [n for n in nodes if isinstance(n, dict) and ( + n.get("type") == "ModelSamplingSD3" or + n.get("class_type") == "ModelSamplingSD3" or + (isinstance(n.get("properties"), dict) and n["properties"].get("Node name for S&R") == "ModelSamplingSD3") + )] + if ms_nodes: + nd = ms_nodes[-1] + # Prefer explicit inputs.shift if present and literal + inp = nd.get("inputs") + if isinstance(inp, dict) and "shift" in inp: + vv = inp["shift"] + # ignore connection like [node_id, idx] + if not isinstance(vv, (list, tuple)): + v2 = try_float(vv) + if v2 is not None: + return v2 + # Fallback: first widget is shift for SD3 (as seen in your JSON) + w = nd.get("widgets_values") + if isinstance(w, list) and len(w) >= 1: + v2 = try_float(w[0]) + if v2 is not None: + return v2 + + # Runtime prompt map: {"42": {"class_type":"ModelSamplingSD3", "inputs":{...}, "widgets_values":[...]}, ...} + # Heuristic: values that are dicts with class_type keys + has_ct = [v for v in obj.values() if isinstance(v, dict) and "class_type" in v] + if has_ct: + for nd in has_ct: + if nd.get("class_type") == "ModelSamplingSD3": + inp = nd.get("inputs", {}) + if isinstance(inp, dict) and "shift" in inp: + vv = inp["shift"] + if not isinstance(vv, (list, tuple)): + v2 = try_float(vv) + if v2 is not None: + return v2 + w = nd.get("widgets_values") + if isinstance(w, list) and len(w) >= 1: + v2 = try_float(w[0]) + if v2 is not None: + return v2 + + # Lists / nested + if isinstance(obj, list): + for it in obj: + v = search_container(it) + if v is not None: + return v + return None + + # 2) Look in provided prompt_json + v = search_container(prompt_json) + if v is not None: + return v + + # Also look in common meta fields that can hold the full workflow/prompt + for key in ("workflow", "prompt", "extra_pnginfo"): + candidate = meta.get(key) + cand_obj = safe_load(candidate) + if isinstance(cand_obj, dict) or isinstance(cand_obj, list): + v = search_container(cand_obj) + if v is not None: + return v + # extra_pnginfo can nest "workflow"/"prompt" again + if isinstance(cand_obj, dict): + for subkey in ("workflow", "prompt"): + sub = safe_load(cand_obj.get(subkey)) + if isinstance(sub, dict) or isinstance(sub, list): + v = search_container(sub) + if v is not None: + return v + + # default + return 5.0 + + @classmethod + def IS_CHANGED(s, latent): + # Fix path lookup (add "latents/" prefix back) + p = folder_paths.get_annotated_filepath(f"latents/{latent}") + m = hashlib.sha256() + with open(p, "rb") as f: + m.update(f.read()) + side = p.replace(".latent", ".cond.pt") + if os.path.exists(side): + with open(side, "rb") as f: + m.update(f.read()) + return m.digest().hex() + + @classmethod + def VALIDATE_INPUTS(s, latent): + check_path = latent if latent.startswith("latents/") else f"latents/{latent}" + try: + folder_paths.get_annotated_filepath(check_path) + except Exception: + return f"Invalid latent file: {latent}" + return True + + def load(self, latent, run_folder=False, refresh_before_run=False): + # Ensure we prepend "latents/" if missing, but don't duplicate it + if not latent.startswith("latents/"): + latent_path = folder_paths.get_annotated_filepath(f"latents/{latent}") + else: + latent_path = folder_paths.get_annotated_filepath(latent) + + sample_dict, meta, _ = _load_latent_file(latent_path) + t = sample_dict["samples"] + + if isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) > 1: + samples = {"samples": t[0:1].contiguous()} + elif isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) == 1: + samples = {"samples": t} + else: + samples = {"samples": t.unsqueeze(0)} + + prompt_json = _safe_json_loads(meta.get("prompt")) + _pos_text, _neg_text, steps, cfg, sampler_name, scheduler, end_at_step = \ + _extract_params_from_prompt_json(prompt_json or {}, meta) + + # SD3 shift (not in KSamplerAdvanced, but we want it) + shift = self._extract_sd3_shift(meta, prompt_json) + + sampler_name = self._coerce_enum(sampler_name, getattr(self.__class__, "_SAMPLERS_ENUM", ())) + scheduler = self._coerce_enum(scheduler, getattr(self.__class__, "_SCHEDULERS_ENUM", ())) + + def normalize_folder(part: str) -> str: + part = part.replace("\\", "/").strip("/") + if not part: + return "" + segments = [seg for seg in part.split("/") if seg] + if segments and segments[0].lower() == "latents": + segments = segments[1:] + return "/".join(segments) + + folder_part = normalize_folder(os.path.dirname(latent)) + base_name = os.path.basename(latent_path) + clean_stem = self._strip_counter(base_name) + prefix = f"{folder_part}/{clean_stem}" if folder_part else clean_stem + + # Raw workflow JSON embedded when this latent was saved (empty string if none). + source_workflow = meta.get("workflow") or "" + + positive_conditioning, negative_conditioning, sidecar = _load_i2v_conditioning_sidecar(latent_path) + trim_latent = _coerce_trim_latent(sidecar.get("trim_latent", meta.get("trim_latent", 0))) + + return ( + float(shift), + positive_conditioning, + negative_conditioning, + samples, + int(steps), + float(cfg), + sampler_name, + scheduler, + int(end_at_step), + prefix, + trim_latent, + source_workflow, + ) + +# ---------- Load multiple latents from a folder (conditioning + sampler params + optional trim_latent) ---------- +class LoadLatents_FromFolder_I2V_MXD: + """ + Default folder/batch loader: same outputs as LoadLatent_I2V_MXD, one set per latent + found in the folder. + """ + DESCRIPTION = """Load all latents in a folder with conditioning and sampler settings.""" + TITLE = "Load Latent Batch MXD" + CATEGORY = "MXD/Latents" + FUNCTION = "load_batch_i2v" + + RETURN_TYPES = ( + "FLOAT", # shift + "CONDITIONING", # positive conditioning + "CONDITIONING", # negative conditioning + "LATENT", + "INT", + "FLOAT", + "STRING", # will be replaced with sampler enum in INPUT_TYPES + "STRING", # will be replaced with scheduler enum in INPUT_TYPES + "INT", + "STRING", + "INT", # trim_latent + ) + RETURN_NAMES = ( + "shift", + "positive", + "negative", + "samples", + "steps", + "cfg", + "sampler_name", + "scheduler", + "end_at_step", + "filename_prefix", + "trim_latent", + ) + + OUTPUT_IS_LIST = (True,) * 11 + + @classmethod + def INPUT_TYPES(s): + # Same folder logic as the single loader + latents_root = os.path.join(folder_paths.get_input_directory(), "latents") + os.makedirs(latents_root, exist_ok=True) + subs = [""] + _list_latent_subfolders(latents_root) + + # Pull live enums from KSamplerAdvanced so sampler/scheduler wire cleanly + from nodes import KSamplerAdvanced + ks_inputs = KSamplerAdvanced.INPUT_TYPES().get("required", {}) + samplers_enum = ks_inputs.get("sampler_name", ("STRING",))[0] + schedulers_enum = ks_inputs.get("scheduler", ("STRING",))[0] + + s.RETURN_TYPES = ( + "FLOAT", # shift + "CONDITIONING", # positive conditioning + "CONDITIONING", # negative conditioning + "LATENT", + "INT", + "FLOAT", + samplers_enum, # enum type for sampler_name + schedulers_enum, # enum type for scheduler + "INT", + "STRING", + "INT", + ) + s._SAMPLERS_ENUM = samplers_enum + s._SCHEDULERS_ENUM = schedulers_enum + + return {"required": {"subfolder": (subs, )}} + + def _coerce_enum(self, value, enum_values): + try: + return value if (enum_values and value in enum_values) else (enum_values[0] if enum_values else value) + except Exception: + return value + + def _strip_counter(self, name: str) -> str: + stem, _ = os.path.splitext(name) + m = re.match(r"^(.*?)(?:_\d{5}_)$", stem) + return m.group(1) if m else stem + + def _extract_sd3_shift(self, meta: dict, prompt_json: dict | None) -> float: + def try_float(x): + try: return float(x) + except Exception: return None + + if isinstance(meta, dict): + v = try_float(meta.get("shift")) + if v is not None: return v + + def safe_load(x): + try: return _safe_json_loads(x) if isinstance(x, str) else x + except Exception: return None + + def search_container(obj): + if isinstance(obj, dict): + if "shift" in obj: + v = try_float(obj.get("shift")) + if v is not None: return v + nodes = obj.get("nodes") + if isinstance(nodes, list): + ms_nodes = [n for n in nodes if isinstance(n, dict) and ( + n.get("type") == "ModelSamplingSD3" or + n.get("class_type") == "ModelSamplingSD3" or + (isinstance(n.get("properties"), dict) and n["properties"].get("Node name for S&R") == "ModelSamplingSD3") + )] + if ms_nodes: + nd = ms_nodes[-1] + inp = nd.get("inputs") + if isinstance(inp, dict) and "shift" in inp: + vv = inp["shift"] + if not isinstance(vv, (list, tuple)): + v2 = try_float(vv) + if v2 is not None: return v2 + w = nd.get("widgets_values") + if isinstance(w, list) and len(w) >= 1: + v2 = try_float(w[0]) + if v2 is not None: return v2 + has_ct = [v for v in obj.values() if isinstance(v, dict) and "class_type" in v] + for nd in has_ct: + if nd.get("class_type") == "ModelSamplingSD3": + inp = nd.get("inputs", {}) + if isinstance(inp, dict) and "shift" in inp: + vv = inp["shift"] + if not isinstance(vv, (list, tuple)): + v2 = try_float(vv) + if v2 is not None: return v2 + w = nd.get("widgets_values") + if isinstance(w, list) and len(w) >= 1: + v2 = try_float(w[0]) + if v2 is not None: return v2 + if isinstance(obj, list): + for it in obj: + v = search_container(it) + if v is not None: return v + return None + + v = search_container(prompt_json) + if v is not None: return v + + for key in ("workflow", "prompt", "extra_pnginfo"): + candidate = meta.get(key) + cand_obj = safe_load(candidate) + if isinstance(cand_obj, (dict, list)): + v = search_container(cand_obj) + if v is not None: return v + if isinstance(cand_obj, dict): + for subkey in ("workflow", "prompt"): + sub = safe_load(cand_obj.get(subkey)) + if isinstance(sub, (dict, list)): + v = search_container(sub) + if v is not None: return v + + return 5.0 + + def load_batch_i2v(self, subfolder): + latents_root = os.path.join(folder_paths.get_input_directory(), "latents") + base = os.path.join(latents_root, subfolder) if subfolder else latents_root + files = glob.glob(os.path.join(base, "**", "*.latent"), recursive=True) + files = _sort_paths_newest_first(files) + if not files: + raise RuntimeError(f"[LoadLatents_FromFolder_I2V_MXD] No .latent files found in '{base}'.") + + shifts, samples_list = [], [] + positives, negatives = [], [] + steps_list, cfgs, samplers, schedulers, end_steps = [], [], [], [], [] + filename_prefixes, trims = [], [] + + for path in files: + sample_dict, meta, _ = _load_latent_file(path) + t = sample_dict["samples"] + + if isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) > 1: + slices = [t[i:i+1].contiguous() for i in range(t.size(0))] + else: + slices = [t if (isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) == 1) + else t.unsqueeze(0)] + + prompt_json = _safe_json_loads(meta.get("prompt")) + _pos_text, _neg_text, n_steps, cfg, sampler_name, scheduler, end_at_step = \ + _extract_params_from_prompt_json(prompt_json or {}, meta) + + sampler_name = self._coerce_enum(sampler_name, getattr(self.__class__, "_SAMPLERS_ENUM", ())) + scheduler = self._coerce_enum(scheduler, getattr(self.__class__, "_SCHEDULERS_ENUM", ())) + shift_val = self._extract_sd3_shift(meta, prompt_json) + + positive_conditioning, negative_conditioning, sidecar = _load_i2v_conditioning_sidecar(path) + trim_latent = _coerce_trim_latent(sidecar.get("trim_latent", meta.get("trim_latent", 0))) + + folder_part = subfolder if subfolder else "" + clean_stem = self._strip_counter(os.path.basename(path)) + prefix = os.path.join(folder_part, clean_stem) if folder_part else clean_stem + + for sl in slices: + shifts.append(float(shift_val)) + positives.append(positive_conditioning) + negatives.append(negative_conditioning) + samples_list.append({"samples": sl}) + steps_list.append(int(n_steps)) + cfgs.append(float(cfg)) + samplers.append(sampler_name) + schedulers.append(scheduler) + end_steps.append(int(end_at_step)) + filename_prefixes.append(prefix) + trims.append(trim_latent) + + return ( + shifts, + positives, + negatives, + samples_list, + steps_list, + cfgs, + samplers, + schedulers, + end_steps, + filename_prefixes, + trims, + ) + +# ---------- Pipe variants: bundle all the loader outputs into one wire ---------- +class LoadLatent_I2V_Pipe_MXD(LoadLatent_I2V_MXD): + """ + Same loading logic as LoadLatent_I2V_MXD, but bundles every value into a single + MXD_LATENT_PIPE output so switching between the single/batch loaders is a one-wire swap. + Unpack with LatentPipeUnpack_MXD. + """ + TITLE = "Load Latent Pipe MXD" + CATEGORY = "MXD/Latents" + FUNCTION = "load_pipe" + + RETURN_TYPES = ("MXD_LATENT_PIPE",) + RETURN_NAMES = ("latent_pipe",) + + @classmethod + def INPUT_TYPES(s): + inputs = LoadLatent_I2V_MXD.INPUT_TYPES.__func__(s) + s.RETURN_TYPES = ("MXD_LATENT_PIPE",) + return inputs + + def load_pipe(self, latent, run_folder=False, refresh_before_run=False): + ( + shift, positive, negative, samples, + steps, cfg, sampler_name, scheduler, + end_at_step, prefix, trim_latent, source_workflow, + ) = self.load(latent, run_folder, refresh_before_run) + + pipe = { + "shift": shift, + "positive": positive, + "negative": negative, + "samples": samples, + "steps": steps, + "cfg": cfg, + "sampler_name": sampler_name, + "scheduler": scheduler, + "end_at_step": end_at_step, + "filename_prefix": prefix, + "trim_latent": trim_latent, + "high_workflow": source_workflow, + } + return (pipe,) + + +class LoadLatents_FromFolder_I2V_Pipe_MXD(LoadLatents_FromFolder_I2V_MXD): + """ + Same loading logic as LoadLatents_FromFolder_I2V_MXD, but bundles every value into a + single MXD_LATENT_PIPE output per item. Unpack with LatentPipeUnpack_MXD. + """ + TITLE = "Load Latent Batch Pipe MXD" + CATEGORY = "MXD/Latents" + FUNCTION = "load_batch_pipe" + + RETURN_TYPES = ("MXD_LATENT_PIPE",) + RETURN_NAMES = ("latent_pipe",) + OUTPUT_IS_LIST = (True,) + + @classmethod + def INPUT_TYPES(s): + inputs = LoadLatents_FromFolder_I2V_MXD.INPUT_TYPES.__func__(s) + s.RETURN_TYPES = ("MXD_LATENT_PIPE",) + return inputs + + def load_batch_pipe(self, subfolder): + ( + shifts, positives, negatives, samples_list, + steps_list, cfgs, samplers, schedulers, + end_steps, filename_prefixes, trims, + ) = self.load_batch_i2v(subfolder) + + pipes = [] + for i in range(len(samples_list)): + pipes.append({ + "shift": shifts[i], + "positive": positives[i], + "negative": negatives[i], + "samples": samples_list[i], + "steps": steps_list[i], + "cfg": cfgs[i], + "sampler_name": samplers[i], + "scheduler": schedulers[i], + "end_at_step": end_steps[i], + "filename_prefix": filename_prefixes[i], + "trim_latent": trims[i], + }) + return (pipes,) + + +class LatentPipeUnpack_MXD: + """ + Splits an MXD_LATENT_PIPE back into shift, conditioning, samples, and sampler settings. + Works with any MXD latent pipe loader (single or batch, I2V or VACE 2.2) - missing + fields like trim_latent just fall back to a safe default. + """ + DESCRIPTION = """Split a latent pipe back into shift, positive, negative, samples, and sampler settings.""" + TITLE = "Unpack Latent Pipe MXD" + CATEGORY = "MXD/Latents" + FUNCTION = "unpack" + + RETURN_TYPES = ( + "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", + "INT", "FLOAT", "STRING", "STRING", "INT", "STRING", "INT", "STRING", + ) + RETURN_NAMES = ( + "shift", "positive", "negative", "samples", + "steps", "cfg", "sampler_name", "scheduler", + "end_at_step", "filename_prefix", "trim_latent", "high_workflow", + ) + + @classmethod + def INPUT_TYPES(s): + from nodes import KSamplerAdvanced + ks_inputs = KSamplerAdvanced.INPUT_TYPES().get("required", {}) + samplers_enum = ks_inputs.get("sampler_name", ("STRING",))[0] + schedulers_enum = ks_inputs.get("scheduler", ("STRING",))[0] + + s.RETURN_TYPES = ( + "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", + "INT", "FLOAT", samplers_enum, schedulers_enum, + "INT", "STRING", "INT", "STRING", + ) + s._SAMPLERS_ENUM = samplers_enum + s._SCHEDULERS_ENUM = schedulers_enum + + return {"required": {"latent_pipe": ("MXD_LATENT_PIPE",)}} + + def _coerce_enum(self, value, enum_values): + try: + return value if (enum_values and value in enum_values) else (enum_values[0] if enum_values else value) + except Exception: + return value + + def unpack(self, latent_pipe): + sampler_name = self._coerce_enum(latent_pipe.get("sampler_name"), getattr(self.__class__, "_SAMPLERS_ENUM", ())) + scheduler = self._coerce_enum(latent_pipe.get("scheduler"), getattr(self.__class__, "_SCHEDULERS_ENUM", ())) + + return ( + latent_pipe.get("shift", 0.0), + latent_pipe.get("positive", []), + latent_pipe.get("negative", []), + latent_pipe.get("samples"), + latent_pipe.get("steps", 0), + latent_pipe.get("cfg", 0.0), + sampler_name, + scheduler, + latent_pipe.get("end_at_step", 0), + latent_pipe.get("filename_prefix", ""), + latent_pipe.get("trim_latent", 0), + latent_pipe.get("high_workflow", ""), + ) + + +NODE_CLASS_MAPPINGS = { + "SaveLatent_I2V_MXD": SaveLatent_I2V_MXD, + "LoadLatent_I2V_MXD": LoadLatent_I2V_MXD, + "LoadLatents_FromFolder_I2V_MXD": LoadLatents_FromFolder_I2V_MXD, + "LoadLatent_I2V_Pipe_MXD": LoadLatent_I2V_Pipe_MXD, + "LoadLatents_FromFolder_I2V_Pipe_MXD": LoadLatents_FromFolder_I2V_Pipe_MXD, + "LatentPipeUnpack_MXD": LatentPipeUnpack_MXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "SaveLatent_I2V_MXD": "Save Latent MXD", + "LoadLatent_I2V_MXD": "Load Latent MXD", + "LoadLatents_FromFolder_I2V_MXD": "Load Latent Batch MXD", + "LoadLatent_I2V_Pipe_MXD": "Load Latent Pipe MXD", + "LoadLatents_FromFolder_I2V_Pipe_MXD": "Load Latent Batch Pipe MXD", + "LatentPipeUnpack_MXD": "Unpack Latent Pipe MXD", +} diff --git a/nodes/wan22/video_ops.py b/nodes/wan22/video_ops.py new file mode 100644 index 0000000..e8b5144 --- /dev/null +++ b/nodes/wan22/video_ops.py @@ -0,0 +1,588 @@ +"""Video frame utilities and video I/O nodes. + +Registered nodes (always): + Frames_Select_StartEnd_MXD Select Frames MXD + Frames_Remove_From_Start_MXD Remove Frames From Start MXD + GroupVideoFramesMXD Group Video Frames MXD + +Registered nodes (only when HAVE_COMFY_API): + CombineVideos_MXD Combine Videos MXD + LoadVideoMXD Load Video MXD + SaveVideoMXD Save Video MXD (merges a prior stage's workflow + into the embedded metadata via latent_io helpers) + PreviewVideoMXD Preview Video MXD + +Route: GET /mxd/videos/input (video-only file list for LoadVideoMXD's combo). +""" +from __future__ import annotations +import os + +import torch + +import folder_paths +import comfy.model_management +from comfy.cli_args import args + +# Comfy API +try: + from comfy_api.latest import io, ui + from comfy_api.input import VideoInput + from comfy_api.input_impl import VideoFromFile, VideoFromComponents + from comfy_api.util import VideoComponents, VideoContainer, VideoCodec + HAVE_COMFY_API = True +except Exception as _e: + io = None + ui = None + VideoInput = None + VideoFromFile = None + VideoFromComponents = None + VideoComponents = None + VideoContainer = None + VideoCodec = None + HAVE_COMFY_API = False + print(f"[ComfyUI-MaxedOut] comfy_api not available in wan22.video_ops: {_e}") + +from server import PromptServer +from aiohttp import web + +from .latent_io import _merge_prior_workflow_into_current + +VIDEO_EXTS = {".mp4", ".mov", ".mkv", ".webm", ".avi"} + +routes = PromptServer.instance.routes + + +@routes.get("/mxd/videos/input") +async def mxd_list_input_videos(request): + """ + Return a JSON list of *video* files under the input folder (relative paths), + sorted by last modified time (newest first) so the combo's 'first' entry + is always the latest render. + """ + input_dir = folder_paths.get_input_directory() + entries = [] + + for root, _, filenames in os.walk(input_dir): + for name in filenames: + ext = os.path.splitext(name)[1].lower() + if ext in VIDEO_EXTS: + full = os.path.join(root, name) + rel = os.path.relpath(full, input_dir).replace("\\", "/") + try: + mtime = os.path.getmtime(full) + except OSError: + mtime = 0 + entries.append((mtime, rel)) + + # Sort newest -> oldest, to match Comfy's internal behavior + entries.sort(key=lambda x: x[0], reverse=True) + + files = [rel for _, rel in entries] + return web.json_response(files) + + +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 + + +# ---------- MXD Frames Select Start/End (from start or end of sequence) ---------- +class Frames_Select_StartEnd_MXD: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "frames": ("IMAGE",), + "count": ("INT", { + "default": 1, + "min": 1, + "max": 10000, + "tooltip": "Number of frames to select" + }), + "offset": ("INT", { + "default": 1, + "min": 1, + "max": 10000, + "tooltip": "How far into the video to start selection (from start or end)" + }), + "mode": (["start", "end"], { + "default": "end", + "tooltip": "Select frames from the start or end of the sequence" + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + FUNCTION = "main" + 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,) + + +# ---------- MXD Frames Remove From Start ---------- +class Frames_Remove_From_Start_MXD: + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "frames": ("IMAGE",), + "count": ("INT", { + "default": 10, + "min": 1, + "max": 10000, + "tooltip": "Number of frames to remove from the start" + }), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + FUNCTION = "main" + CATEGORY = "MXD/images" + + def main(self, frames=None, count=10): + # Skip the first `count` frames instead of keeping them + frames_after = frames[count:].clone() + return (frames_after,) + + +class GroupVideoFramesMXD: + CATEGORY = "MXD/Video" + TITLE = "Group Video Frames (MXD)" + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("IMAGE_GROUPS",) + OUTPUT_IS_LIST = (True,) + FUNCTION = "group_frames" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "frames": ("IMAGE",), + "group_size": ("INT", {"default": 81, "min": 1, "max": 5000, "step": 1}), + } + } + + def group_frames(self, frames, group_size): + import math, torch + + all_frames = list(frames) + total = len(all_frames) + num_groups = math.ceil(total / group_size) + grouped_tensors = [] + + for i in range(num_groups): + start = i * group_size + end = min(start + group_size, total) + group = all_frames[start:end] + + clean = [] + for f in group: + # drop redundant singleton batch dim if present + if f.ndim == 4 and f.shape[0] == 1: + f = f.squeeze(0) # (H,W,C) + # ensure shape (H,W,C) + if f.ndim != 3: + print(f"[GroupVideoFramesMXD] weird frame shape {f.shape}") + continue + clean.append(f) + + # stack back to (N,H,W,C) + if len(clean) == 0: + continue + stacked = torch.stack(clean, dim=0) + grouped_tensors.append(stacked) + + print(f"[GroupVideoFramesMXD] Split {total} frames into {len(grouped_tensors)} groups of up to {group_size}.") + return (grouped_tensors,) + + +if HAVE_COMFY_API: + class CombineVideos_MXD: + """ + Combine two VIDEO inputs end-to-end (sequentially). + """ + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "front_video": ("VIDEO", {"tooltip": "The first video (plays first)"}), + "back_video": ("VIDEO", {"tooltip": "The second video (plays after the first)"}), + }, + } + + RETURN_TYPES = ("VIDEO",) + RETURN_NAMES = ("video",) + FUNCTION = "combine" + CATEGORY = "MXD/video" + + def combine(self, front_video, back_video): + comp_a = front_video.get_components() + comp_b = back_video.get_components() + + # Check frame rate consistency + if comp_a.frame_rate != comp_b.frame_rate: + raise ValueError(f"FPS mismatch: {comp_a.frame_rate} vs {comp_b.frame_rate}") + + # 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: + 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 + + combined_video = VideoFromComponents( + VideoComponents( + images=combined_images, + audio=combined_audio, + frame_rate=comp_a.frame_rate, + ) + ) + + return (combined_video,) + + # ---------- Load Video MXD (video-only picker with refresh) ---------- + class LoadVideoMXD: + """Load a video from /input with a refresh button (videos only).""" + + CATEGORY = "image/video" + FUNCTION = "load" + RETURN_TYPES = ("VIDEO", "STRING") + RETURN_NAMES = ("video", "video_path") + TITLE = "Load Video MXD" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "file": ("COMBO", { + # Only allow video uploads in the picker + "video_upload": True, + # Custom route that returns ONLY videos in /input + "remote": { + "route": "/mxd/videos/input", + "refresh_button": True, + "control_after_refresh": "first", + }, + }), + } + } + + # --- helpers -------------------------------------------------------------- + + @staticmethod + def _resolve_video_path(file: str) -> str: + """ + Try to resolve `file` in a backwards-compatible way: + 1. If it's an annotated path, let folder_paths handle it. + 2. Otherwise treat it as relative to the input directory. + """ + # 1) Try annotated style (old workflows / uploads) + try: + return folder_paths.get_annotated_filepath(file) + except Exception: + pass + + # 2) Fall back to /input relative + base = folder_paths.get_input_directory() + candidate = os.path.join(base, file) + if os.path.isfile(candidate): + return candidate + + # If all else fails, just return what we got (will error later) + return candidate + + @staticmethod + def _is_video_file(path: str) -> bool: + _, ext = os.path.splitext(path) + return ext.lower() in VIDEO_EXTS + + # --- main function -------------------------------------------------------- + + def load(self, file: str): + video_path = self._resolve_video_path(file) + + if not os.path.isfile(video_path): + raise FileNotFoundError(f"[LoadVideoMXD] File not found: {video_path}") + + if not self._is_video_file(video_path): + raise ValueError(f"[LoadVideoMXD] Not a video file: {video_path}") + + print(f"[LoadVideoMXD] Loaded exactly: {video_path}") + return (VideoFromFile(video_path), video_path) + + # --- nice-to-haves -------------------------------------------------------- + + @classmethod + def IS_CHANGED(cls, file: str): + try: + p = cls._resolve_video_path(file) + return os.path.getmtime(p) + except Exception: + return 0 + + @classmethod + def VALIDATE_INPUTS(cls, file: str): + # First, try the annotated path (for backwards compat) + if folder_paths.exists_annotated_filepath(file): + resolved = folder_paths.get_annotated_filepath(file) + if not cls._is_video_file(resolved): + return f"This node only accepts video files ({', '.join(sorted(VIDEO_EXTS))})." + return True + + # Then, try treating it as /input-relative + base = folder_paths.get_input_directory() + candidate = os.path.join(base, file) + if os.path.isfile(candidate): + if not cls._is_video_file(candidate): + return f"This node only accepts video files ({', '.join(sorted(VIDEO_EXTS))})." + return True + + return f"Invalid video file: {file}" + + # ---------- Save Video MXD ---------- + class SaveVideoMXD(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="SaveVideoMXD", + display_name="Save Video MXD", + category="image/video", + description="Saves the input video to your ComfyUI output directory.", + inputs=[ + io.Video.Input("video", tooltip="The video to save."), + io.String.Input("filename_prefix", default="video/ComfyUI", tooltip="The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes."), + io.Combo.Input("format", options=VideoContainer.as_input(), default="auto", tooltip="The format to save the video as."), + io.Combo.Input("codec", options=VideoCodec.as_input(), default="auto", tooltip="The codec to use for the video."), + io.Boolean.Input( + "embed_workflow", + default=True, + label_on="embed", + label_off="skip", + tooltip="When high_workflow is connected, merge it into this video's embedded workflow " + "so dragging the final video into ComfyUI shows both the high-noise stage and " + "this stage together.", + ), + io.String.Input( + "high_workflow", + optional=True, + force_input=True, + tooltip="Connect a Load Latent node's 'high_workflow' output here to carry the " + "high-noise stage's workflow into this video's metadata.", + ), + ], + hidden=[io.Hidden.prompt, io.Hidden.extra_pnginfo], + is_output_node=True, + ) + + @classmethod + def execute(cls, video: VideoInput, filename_prefix: str, format: str, codec: str, + embed_workflow: bool = True, high_workflow: str = "") -> io.NodeOutput: + width, height = video.get_dimensions() + full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path( + filename_prefix, + folder_paths.get_output_directory(), + width, + height + ) + + saved_metadata = None + if not args.disable_metadata: + metadata = {} + if cls.hidden.extra_pnginfo is not None: + metadata.update(cls.hidden.extra_pnginfo) + if cls.hidden.prompt is not None: + metadata["prompt"] = cls.hidden.prompt + if embed_workflow and high_workflow: + current_workflow = metadata.get("workflow") + merged_workflow = _merge_prior_workflow_into_current(high_workflow, current_workflow) + if merged_workflow is not current_workflow: + metadata["workflow"] = merged_workflow + if len(metadata) > 0: + saved_metadata = metadata + + file = f"{filename}_{counter:05}_.{VideoContainer.get_extension(format)}" + video.save_to( + os.path.join(full_output_folder, file), + format=VideoContainer(format), + codec=codec, + metadata=saved_metadata + ) + + return io.NodeOutput(ui=ui.PreviewVideo([ui.SavedResult(file, subfolder, io.FolderType.output)])) + + class PreviewVideoMXD(io.ComfyNode): + @classmethod + def define_schema(cls): + return io.Schema( + node_id="PreviewVideoMXD", + display_name="Preview Video MXD", + category="image/video", + 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 + def execute(cls, input_video: VideoInput): + # Save a temporary H264 file so ComfyUI has something to preview + out_dir = os.path.join(folder_paths.get_output_directory(), "previews") + os.makedirs(out_dir, exist_ok=True) + + preview_path = os.path.join(out_dir, "preview_temp.mp4") + input_video.save_to(preview_path, format="mp4", codec="h264") + + # Return the raw video object (not a tuple) + return io.NodeOutput( + input_video, + ui=ui.PreviewVideo([ + ui.SavedResult("preview_temp.mp4", "previews", io.FolderType.output) + ]) + ) + + +NODE_CLASS_MAPPINGS = { + "Frames_Remove_From_Start_MXD": Frames_Remove_From_Start_MXD, + "GroupVideoFramesMXD": GroupVideoFramesMXD, + "Frames_Select_StartEnd_MXD": Frames_Select_StartEnd_MXD, +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "Frames_Remove_From_Start_MXD": "Remove Frames From Start MXD", + "GroupVideoFramesMXD": "Group Video Frames MXD", + "Frames_Select_StartEnd_MXD": "Select Frames MXD", +} + +if HAVE_COMFY_API: + NODE_CLASS_MAPPINGS.update({ + "CombineVideos_MXD": CombineVideos_MXD, + "LoadVideoMXD": LoadVideoMXD, + "SaveVideoMXD": SaveVideoMXD, + "PreviewVideoMXD": PreviewVideoMXD, + }) + NODE_DISPLAY_NAME_MAPPINGS.update({ + "CombineVideos_MXD": "Combine Videos MXD", + "LoadVideoMXD": "Load Video MXD", + "SaveVideoMXD": "Save Video MXD", + "PreviewVideoMXD": "Preview Video MXD", + }) diff --git a/system/__init__.py b/system/__init__.py new file mode 100644 index 0000000..fd33a29 --- /dev/null +++ b/system/__init__.py @@ -0,0 +1,5 @@ +"""Import-time side-effect modules (no nodes registered here). + + live_preview.py monkeypatches latent_preview for streaming video previews + model_paths.py registers the user's external model storage folders +""" diff --git a/video_preview_mxd.py b/system/live_preview.py similarity index 100% rename from video_preview_mxd.py rename to system/live_preview.py diff --git a/model_paths_autoregister_mxd.py b/system/model_paths.py similarity index 88% rename from model_paths_autoregister_mxd.py rename to system/model_paths.py index 140daf9..d11bff2 100644 --- a/model_paths_autoregister_mxd.py +++ b/system/model_paths.py @@ -3,7 +3,7 @@ with folder_paths, the same way ComfyUI/models/ works. The root is resolved in this order (first hit wins): 1. MAXEDOUT_MODEL_STORAGE environment variable - 2. model_storage_config.json next to this file (gitignored -- copy + 2. model_storage_config.json at the repo root (gitignored -- copy model_storage_config.json.example to create your own, it never gets committed) 3. The "MXD > Model Storage > Root Folder" setting in the ComfyUI @@ -23,8 +23,10 @@ try: except ImportError: folder_paths = None -_THIS_DIR = os.path.dirname(os.path.abspath(__file__)) -_CONFIG_PATH = os.path.join(_THIS_DIR, "model_storage_config.json") +# The config lives at the REPO ROOT (one level above this system/ package), +# where users have always placed it — keep that path stable across refactors. +_REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +_CONFIG_PATH = os.path.join(_REPO_ROOT, "model_storage_config.json") def _root_from_env(): diff --git a/wan22nodes.py b/wan22nodes.py deleted file mode 100644 index 7187334..0000000 --- a/wan22nodes.py +++ /dev/null @@ -1,3093 +0,0 @@ -from __future__ import annotations -import os, re, glob, json, hashlib, copy -from collections import deque -from typing import Any, Dict, Tuple, Optional, List, Union - -import torch -from safetensors import safe_open - -import folder_paths -import comfy.utils -import comfy.model_management -from comfy.cli_args import args -from nodes import KSamplerAdvanced -import node_helpers, nodes - -# Comfy API -try: - from comfy_api.latest import io, ui - from comfy_api.input import VideoInput - from comfy_api.input_impl import VideoFromFile, VideoFromComponents - from comfy_api.util import VideoComponents, VideoContainer, VideoCodec - HAVE_COMFY_API = True -except Exception as _e: - io = None - ui = None - VideoInput = None - VideoFromFile = None - VideoFromComponents = None - VideoComponents = None - VideoContainer = None - VideoCodec = None - HAVE_COMFY_API = False - print(f"[ComfyUI-MaxedOut] comfy_api not available in wan22nodes: {_e}") - -from server import PromptServer -from aiohttp import web - -VIDEO_EXTS = {".mp4", ".mov", ".mkv", ".webm", ".avi"} - -routes = PromptServer.instance.routes - -def _sort_paths_newest_first(paths: List[str]) -> List[str]: - """Sort file paths by mtime desc (newest first), stable by normalized path.""" - def _mtime(path: str) -> float: - try: - return os.path.getmtime(path) - except OSError: - return 0.0 - - return sorted( - paths, - key=lambda p: (-_mtime(p), p.replace("\\", "/").lower()), - ) - - -def _sort_latent_options_by_folder(options: List[str], root: str = "") -> List[str]: - """ - Order relative '.latent' option paths so that the combo's prev/next arrows - stay confined to one folder before moving on, newest first: - folderA/file1, folderA/file2, ..., folderB/file1, ... - Folders are ordered by the mtime of their most recently modified file (so - a folder that just received a new file jumps back to the top), and files - within each folder are newest first. - """ - def _mtime(rel: str) -> float: - if not root: - return 0.0 - try: - return os.path.getmtime(os.path.join(root, rel)) - except OSError: - return 0.0 - - folder_of = lambda rel: rel.replace("\\", "/").rsplit("/", 1)[0] if "/" in rel.replace("\\", "/") else "" - - folder_latest: Dict[str, float] = {} - for rel in options: - folder = folder_of(rel) - m = _mtime(rel) - if m > folder_latest.get(folder, -1.0): - folder_latest[folder] = m - - def key(rel: str): - folder = folder_of(rel) - return (-folder_latest.get(folder, 0.0), -_mtime(rel)) - - return sorted(options, key=key) - -def _list_latent_subfolders(latents_root: str) -> List[str]: - """ - List latent subfolders recursively (e.g. "a", "a/b"), newest first by - latest latent mtime in each branch. - """ - files = glob.glob(os.path.join(latents_root, "**", "*.latent"), recursive=True) - if not files: - return [] - - folder_latest_mtime: Dict[str, float] = {} - for file_path in files: - rel_dir = os.path.relpath(os.path.dirname(file_path), latents_root).replace(os.sep, "/").strip("/") - if not rel_dir or rel_dir == ".": - continue - try: - mtime = os.path.getmtime(file_path) - except OSError: - mtime = 0.0 - - # Include each ancestor so both "a" and "a/b" appear as options. - parts = [p for p in rel_dir.split("/") if p] - for i in range(1, len(parts) + 1): - branch = "/".join(parts[:i]) - prev = folder_latest_mtime.get(branch, -1.0) - if mtime > prev: - folder_latest_mtime[branch] = mtime - - return [ - folder - for folder, _ in sorted( - folder_latest_mtime.items(), - key=lambda kv: (-kv[1], kv[0].lower()), - ) - ] - -@routes.get("/mxd/videos/input") -async def mxd_list_input_videos(request): - """ - Return a JSON list of *video* files under the input folder (relative paths), - sorted by last modified time (newest first) so the combo's 'first' entry - is always the latest render. - """ - input_dir = folder_paths.get_input_directory() - entries = [] - - for root, _, filenames in os.walk(input_dir): - for name in filenames: - ext = os.path.splitext(name)[1].lower() - if ext in VIDEO_EXTS: - full = os.path.join(root, name) - rel = os.path.relpath(full, input_dir).replace("\\", "/") - try: - mtime = os.path.getmtime(full) - except OSError: - mtime = 0 - entries.append((mtime, rel)) - - # 🔁 Sort newest → oldest, to match Comfy's internal behavior - entries.sort(key=lambda x: x[0], reverse=True) - - files = [rel for _, rel in entries] - return web.json_response(files) - - -@routes.get("/mxd/latents/files") -async def mxd_list_latent_files(request): - """ - Fresh re-scan of input/latents for .latent files. Used by the run_folder - queuing loop (refresh_before_run) to pick up files a still-running - workflow is writing concurrently, instead of relying on the dropdown - list captured whenever the node's combo was last populated. - """ - latents_root = os.path.join(folder_paths.get_input_directory(), "latents") - os.makedirs(latents_root, exist_ok=True) - files = glob.glob(os.path.join(latents_root, "**", "*.latent"), recursive=True) - options = [os.path.relpath(f, latents_root).replace(os.sep, "/") for f in files] - options = _sort_latent_options_by_folder(options, latents_root) - return web.json_response(options) - - -# ---------- SaveLatent (saves latent + conditioning + optional trim_latent) ---------- -class SaveLatent_I2V_MXD: - """ - Default latent saver, works for t2v, i2v, and VACE 2.2. Persists: - • latent tensor -> .latent - • pos/neg CONDITIONING -> .cond.pt - • optional trim_latent value (VACE 2.2) -> .cond.pt - """ - TITLE = "Save Latent MXD" - CATEGORY = "MXD/Latents" - OUTPUT_NODE = True - RETURN_TYPES = () - FUNCTION = "save_only" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "samples": ("LATENT", {"tooltip": "Latent to save."}), - "positive": ("CONDITIONING", {"tooltip": "Positive CONDITIONING to save alongside the latent."}), - "negative": ("CONDITIONING", {"tooltip": "Negative CONDITIONING to save alongside the latent."}), - "filename_prefix": ("STRING", {"default": "ComfyUI", "tooltip": "Prefix for saved files"}), - }, - "optional": { - "trim_latent": ("INT", { - "default": 0, - "min": 0, - "max": 10000, - "step": 1, - "tooltip": "VACE 2.2 trim_latent value to preserve with this latent. Usually 0 or 1. Ignored for t2v/i2v." - }), - }, - "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO", "unique_id": "UNIQUE_ID"}, - } - - def save_only(self, samples, positive, negative, filename_prefix="ComfyUI", - trim_latent=0, prompt=None, extra_pnginfo=None, unique_id=None): - _save_i2v_latent_bundle( - samples=samples, - positive=positive, - negative=negative, - filename_prefix=filename_prefix, - prompt=prompt, - extra_pnginfo=extra_pnginfo, - unique_id=unique_id, - sidecar_extra={"trim_latent": _coerce_trim_latent(trim_latent)}, - ) - return {} - -# ---------- Helpers ---------- -def _load_latent_file(latent_path: str) -> Tuple[Dict[str, torch.Tensor], Dict[str, Any], List[str]]: - """ - Load safetensors latent with Comfy metadata. - Returns (samples_dict, metadata_dict, keys_list) - """ - with safe_open(latent_path, framework="pt", device="cpu") as f: - keys = list(f.keys()) - # prefer explicit key we write - if "latent_tensor" in keys: - t = f.get_tensor("latent_tensor").float().contiguous() - else: - # fall back (some variants might save using a different name) - first = keys[0] - t = f.get_tensor(first).float().contiguous() - - meta = f.metadata() or {} - - # if ancient format, rescale (match Comfy behavior) - if "latent_format_version_0" not in keys: - t = t * (1.0 / 0.18215) - - return {"samples": t}, meta, keys - - -def _safe_json_loads(s: Union[str, bytes, None]) -> Optional[Dict[str, Any]]: - if s is None: - return None - if isinstance(s, bytes): - try: - s = s.decode("utf-8", "ignore") - except Exception: - return None - if not isinstance(s, str): - return None - try: - return json.loads(s) - except Exception: - # sometimes double-encoded in metadata - try: - return json.loads(json.loads(s)) - except Exception: - return None - - -def _workflow_node_bbox(nodes: List[Dict[str, Any]]) -> Optional[Tuple[float, float, float, float]]: - """(min_x, min_y, max_x, max_y) over a litegraph 'nodes' list. None if no positions found.""" - xs0, ys0, xs1, ys1 = [], [], [], [] - for n in nodes: - if not isinstance(n, dict): - continue - pos = n.get("pos") - if isinstance(pos, list) and len(pos) >= 2: - x, y = pos[0], pos[1] - elif isinstance(pos, dict): - x, y = pos.get("0", 0), pos.get("1", 0) - else: - continue - size = n.get("size") - w = size[0] if isinstance(size, (list, tuple)) and len(size) >= 1 else 200 - h = size[1] if isinstance(size, (list, tuple)) and len(size) >= 2 else 100 - xs0.append(x); ys0.append(y); xs1.append(x + w); ys1.append(y + h) - if not xs0: - return None - return (min(xs0), min(ys0), max(xs1), max(ys1)) - - -def _offset_workflow_nodes(nodes: List[Dict[str, Any]], dx: float, dy: float) -> None: - for n in nodes: - if not isinstance(n, dict): - continue - pos = n.get("pos") - if isinstance(pos, list) and len(pos) >= 2: - pos[0] = pos[0] + dx - pos[1] = pos[1] + dy - elif isinstance(pos, dict): - if "0" in pos: pos["0"] = pos["0"] + dx - if "1" in pos: pos["1"] = pos["1"] + dy - - -def _merge_prior_workflow_into_current(prior_workflow_json: Optional[str], current_workflow: Any) -> Any: - """ - Merge a previously-saved workflow graph (embedded in a loaded .latent file) into the - workflow graph of the run that's currently saving. The prior graph's nodes/links/groups - are copied in with fresh ids and shifted to sit to the left of the current graph, wrapped - in a labelled group - so dragging the final video into ComfyUI shows both stages at once, - the same as if you'd copy/pasted the first workflow onto the second one's canvas. - - Best-effort: on any parse/shape problem, returns current_workflow untouched. - """ - if not prior_workflow_json or not isinstance(current_workflow, dict): - return current_workflow - - try: - prior = json.loads(prior_workflow_json) if isinstance(prior_workflow_json, str) else prior_workflow_json - if not isinstance(prior, dict): - return current_workflow - - prior_nodes = prior.get("nodes") - if not isinstance(prior_nodes, list) or not prior_nodes: - return current_workflow - - merged = copy.deepcopy(current_workflow) - current_nodes = merged.get("nodes") - if not isinstance(current_nodes, list): - current_nodes = [] - merged["nodes"] = current_nodes - - prior_nodes = copy.deepcopy(prior_nodes) - prior_links = copy.deepcopy(prior.get("links")) if isinstance(prior.get("links"), list) else [] - prior_groups = copy.deepcopy(prior.get("groups")) if isinstance(prior.get("groups"), list) else [] - - # ---- remap node ids so they can't collide with the current graph ---- - current_last_node_id = merged.get("last_node_id") - if not isinstance(current_last_node_id, int): - current_last_node_id = max((n.get("id", 0) for n in current_nodes if isinstance(n, dict)), default=0) - next_node_id = current_last_node_id + 1 - node_id_map: Dict[Any, int] = {} - for n in prior_nodes: - if not isinstance(n, dict) or "id" not in n: - continue - node_id_map[n["id"]] = next_node_id - n["id"] = next_node_id - next_node_id += 1 - - # ---- remap link ids the same way ---- - current_last_link_id = merged.get("last_link_id") - if not isinstance(current_last_link_id, int): - current_last_link_id = max( - (l[0] for l in (merged.get("links") or []) if isinstance(l, list) and l), default=0 - ) - next_link_id = current_last_link_id + 1 - link_id_map: Dict[Any, int] = {} - for l in prior_links: - if isinstance(l, list) and l: - link_id_map[l[0]] = next_link_id - next_link_id += 1 - - for n in prior_nodes: - if not isinstance(n, dict): - continue - for inp in (n.get("inputs") or []): - if isinstance(inp, dict) and inp.get("link") is not None: - inp["link"] = link_id_map.get(inp["link"], inp["link"]) - for out in (n.get("outputs") or []): - if isinstance(out, dict) and isinstance(out.get("links"), list): - out["links"] = [link_id_map.get(x, x) for x in out["links"]] - - remapped_links = [] - for l in prior_links: - if not isinstance(l, list) or len(l) < 5: - continue - new_l = list(l) - new_l[0] = link_id_map.get(l[0], l[0]) - new_l[1] = node_id_map.get(l[1], l[1]) - new_l[3] = node_id_map.get(l[3], l[3]) - remapped_links.append(new_l) - - # ---- shift the prior graph so it sits to the left of the current one ---- - current_bbox = _workflow_node_bbox(current_nodes) - prior_bbox = _workflow_node_bbox(prior_nodes) - margin = 400 - if current_bbox and prior_bbox: - dx = (current_bbox[0] - margin) - prior_bbox[2] - dy = current_bbox[1] - prior_bbox[1] - else: - dx, dy = 0, 0 - _offset_workflow_nodes(prior_nodes, dx, dy) - for g in prior_groups: - if not isinstance(g, dict): - continue - b = g.get("bounding") - if isinstance(b, list) and len(b) >= 2: - b[0] = b[0] + dx - b[1] = b[1] + dy - - # wrap the prior graph in a labelled group so it's obvious what it is - wrapper_group = None - prior_bbox_shifted = _workflow_node_bbox(prior_nodes) - if prior_bbox_shifted: - pad = 60 - wrapper_group = { - "title": "Prior stage (loaded latent's source workflow)", - "bounding": [ - prior_bbox_shifted[0] - pad, - prior_bbox_shifted[1] - pad - 40, - (prior_bbox_shifted[2] - prior_bbox_shifted[0]) + pad * 2, - (prior_bbox_shifted[3] - prior_bbox_shifted[1]) + pad * 2 + 40, - ], - "color": "#3f789e", - "font_size": 24, - } - - merged["nodes"] = current_nodes + prior_nodes - merged["links"] = (merged.get("links") or []) + remapped_links - groups = list(merged.get("groups") or []) + prior_groups - if wrapper_group: - groups.append(wrapper_group) - merged["groups"] = groups - merged["last_node_id"] = next_node_id - 1 - merged["last_link_id"] = next_link_id - 1 - return merged - except Exception as e: - print(f"[SaveVideoMXD] Could not merge prior stage workflow into embedded metadata: {e}") - return current_workflow - - -def _node_sort_key(node_id: str) -> Tuple[int, Union[int, str]]: - s = str(node_id) - try: - return (0, int(s)) - except Exception: - return (1, s) - - -def _normalize_prompt_graph(prompt_json: Any) -> Dict[str, Any]: - if not isinstance(prompt_json, dict): - return {} - graph = prompt_json.get("prompt", prompt_json) - return graph if isinstance(graph, dict) else {} - - -def _get_graph_node(graph: Dict[str, Any], node_id: Any) -> Optional[Dict[str, Any]]: - if node_id is None or not isinstance(graph, dict): - return None - node = graph.get(str(node_id)) - return node if isinstance(node, dict) else None - - -def _iter_graph_nodes_sorted(graph: Dict[str, Any]) -> List[Tuple[str, Dict[str, Any]]]: - nodes: List[Tuple[str, Dict[str, Any]]] = [] - for node_id, node in graph.items(): - if isinstance(node, dict): - nodes.append((str(node_id), node)) - nodes.sort(key=lambda pair: _node_sort_key(pair[0])) - return nodes - - -def _linked_node_id(value: Any) -> Optional[str]: - if isinstance(value, (list, tuple)) and len(value) >= 1: - return str(value[0]) - return None - - -def _is_ksampler_node(node: Any) -> bool: - if not isinstance(node, dict): - return False - return "KSampler" in str(node.get("class_type", "")) - - -def _collect_upstream_linked_node_ids(node: Dict[str, Any]) -> List[str]: - inputs = node.get("inputs", {}) - if not isinstance(inputs, dict): - return [] - - seen = set() - ordered = [] - - # Prefer latent-carrying links first. - for key in ("samples", "latent", "latent_image"): - linked = _linked_node_id(inputs.get(key)) - if linked is not None and linked not in seen: - seen.add(linked) - ordered.append(linked) - - # Then search all other connected inputs in stable order. - for key, value in inputs.items(): - if key in ("samples", "latent", "latent_image"): - continue - linked = _linked_node_id(value) - if linked is not None and linked not in seen: - seen.add(linked) - ordered.append(linked) - - return ordered - - -def _find_upstream_ksampler_node_id(graph: Dict[str, Any], start_node_id: Any) -> Optional[str]: - if not isinstance(graph, dict) or start_node_id is None: - return None - - queue: deque[str] = deque([str(start_node_id)]) - visited = set() - - while queue: - node_id = queue.popleft() - if node_id in visited: - continue - visited.add(node_id) - - node = _get_graph_node(graph, node_id) - if not node: - continue - if _is_ksampler_node(node): - return node_id - - for upstream_id in _collect_upstream_linked_node_ids(node): - if upstream_id not in visited: - queue.append(upstream_id) - - return None - - -def _extract_ksampler_params(node: Dict[str, Any]) -> Dict[str, Any]: - inputs = node.get("inputs", {}) if isinstance(node, dict) else {} - if not isinstance(inputs, dict): - inputs = {} - - out: Dict[str, Any] = {} - - def set_int(key: str): - if key in inputs: - try: - out[key] = int(inputs[key]) - except Exception: - pass - - def set_float(key: str): - if key in inputs: - try: - out[key] = float(inputs[key]) - except Exception: - pass - - def set_str(key: str): - if key in inputs and not isinstance(inputs[key], (list, tuple, dict)): - try: - out[key] = str(inputs[key]).strip() - except Exception: - pass - - set_int("steps") - set_float("cfg") - set_str("sampler_name") - set_str("scheduler") - set_int("start_at_step") - set_int("end_at_step") - - return out - - -def _attach_source_ksampler_metadata(meta: Dict[str, Any], prompt: Any, unique_id: Any) -> None: - if not isinstance(meta, dict): - return - - graph = _normalize_prompt_graph(prompt) - if not graph: - return - - save_node_id = str(unique_id) if unique_id is not None else "" - if not save_node_id: - return - - save_node = _get_graph_node(graph, save_node_id) - if not save_node: - return - - source_candidates = _collect_upstream_linked_node_ids(save_node) - if not source_candidates: - return - - source_ksampler_id = None - for start_id in source_candidates: - source_ksampler_id = _find_upstream_ksampler_node_id(graph, start_id) - if source_ksampler_id: - break - - if not source_ksampler_id: - return - - source_node = _get_graph_node(graph, source_ksampler_id) - if not source_node: - return - - meta["mxd_source_save_node_id"] = save_node_id - meta["mxd_source_ksampler_node_id"] = source_ksampler_id - try: - meta["mxd_source_ksampler_params"] = json.dumps(_extract_ksampler_params(source_node)) - except Exception: - pass - - -def _build_latent_metadata(prompt=None, extra_pnginfo=None, unique_id=None, extra_meta=None): - if args.disable_metadata: - return None - - meta = {} - if prompt is not None: - try: - meta["prompt"] = json.dumps(prompt) - except Exception: - pass - if extra_pnginfo is not None: - for k, v in extra_pnginfo.items(): - try: - meta[k] = json.dumps(v) - except Exception: - pass - if isinstance(extra_meta, dict): - for k, v in extra_meta.items(): - try: - meta[str(k)] = json.dumps(v) - except Exception: - pass - _attach_source_ksampler_metadata(meta, prompt, unique_id) - return meta - - -def _save_i2v_latent_bundle( - samples, - positive, - negative, - filename_prefix="I2V", - prompt=None, - extra_pnginfo=None, - unique_id=None, - sidecar_extra=None, -): - latents_dir = os.path.join(folder_paths.get_input_directory(), "latents") - os.makedirs(latents_dir, exist_ok=True) - - full_output_folder, filename, counter, _subfolder, _filename_prefix = folder_paths.get_save_image_path( - filename_prefix, latents_dir - ) - - extra_meta = sidecar_extra if isinstance(sidecar_extra, dict) else None - meta = _build_latent_metadata( - prompt=prompt, - extra_pnginfo=extra_pnginfo, - unique_id=unique_id, - extra_meta=extra_meta, - ) - - latent_path = os.path.join(full_output_folder, f"{filename}_{counter:05}_.latent") - payload = { - "latent_tensor": samples["samples"].contiguous(), - "latent_format_version_0": torch.tensor([]), - } - comfy.utils.save_torch_file(payload, latent_path, metadata=meta) - - sidecar = {"positive": positive, "negative": negative} - if isinstance(sidecar_extra, dict): - sidecar.update(sidecar_extra) - torch.save(sidecar, latent_path.replace(".latent", ".cond.pt")) - return latent_path - - -def _load_i2v_conditioning_sidecar(latent_path): - cond_path = latent_path.replace(".latent", ".cond.pt") - if not os.path.exists(cond_path): - return [], [], {} - - try: - data = torch.load(cond_path, map_location="cpu") - except Exception: - return [], [], {} - - if not isinstance(data, dict): - return [], [], {} - - return data.get("positive", []), data.get("negative", []), data - - -def _coerce_trim_latent(value, default=0): - try: - if isinstance(value, str): - parsed = _safe_json_loads(value) - value = parsed if parsed is not None else value - return int(value) - except Exception: - return int(default) - - -def _extract_prompt_text_from_ksampler(graph: Dict[str, Any], ks_node: Dict[str, Any]) -> Tuple[str, str]: - pos = "" - neg = "" - - inputs = ks_node.get("inputs", {}) if isinstance(ks_node, dict) else {} - if not isinstance(inputs, dict): - return pos, neg - - def _text_from_clip(link_value: Any) -> str: - node_id = _linked_node_id(link_value) - if node_id is None: - return "" - node = _get_graph_node(graph, node_id) or {} - if node.get("class_type") == "CLIPTextEncode": - return str(node.get("inputs", {}).get("text", "")).strip() - return "" - - pos = _text_from_clip(inputs.get("positive")) - neg = _text_from_clip(inputs.get("negative")) - return pos, neg - - -def _extract_params_from_prompt_json( - prompt_json: Dict[str, Any], - meta: Optional[Dict[str, Any]] = None, -) -> Tuple[str, str, int, float, str, str, int]: - """ - Returns: (positive, negative, steps, cfg, sampler_name, scheduler, end_at_step) - parsed from the saved Comfy prompt graph with deterministic KSampler selection. - """ - pos = "" - neg = "" - steps = 20 - cfg = 8.0 - sampler_name = "" - scheduler = "" - end_at_step = 0 - - graph = _normalize_prompt_graph(prompt_json) - if not isinstance(graph, dict): - return pos, neg, steps, cfg, sampler_name, scheduler, end_at_step - - ks_node = None - extracted_params: Dict[str, Any] = {} - - # 1) Source KSampler id saved directly in latent metadata. - if isinstance(meta, dict): - raw_ks = meta.get("mxd_source_ksampler_node_id") - if raw_ks is not None: - candidate = _get_graph_node(graph, str(raw_ks)) - if candidate and _is_ksampler_node(candidate): - ks_node = candidate - - # 2) Source save node id -> trace upstream to nearest KSampler. - if ks_node is None and isinstance(meta, dict): - raw_save = meta.get("mxd_source_save_node_id") - if raw_save is not None: - save_node = _get_graph_node(graph, str(raw_save)) - if save_node: - for start_id in _collect_upstream_linked_node_ids(save_node): - trace_id = _find_upstream_ksampler_node_id(graph, start_id) - if trace_id: - candidate = _get_graph_node(graph, trace_id) - if candidate and _is_ksampler_node(candidate): - ks_node = candidate - break - - # 3) Legacy fallback: last KSampler node in graph. - if ks_node is None: - for _, node in _iter_graph_nodes_sorted(graph): - if _is_ksampler_node(node): - ks_node = node - - if not ks_node: - return pos, neg, steps, cfg, sampler_name, scheduler, end_at_step - - pos, neg = _extract_prompt_text_from_ksampler(graph, ks_node) - extracted_params = _extract_ksampler_params(ks_node) - - if "steps" in extracted_params: - steps = int(extracted_params["steps"]) - if "cfg" in extracted_params: - cfg = float(extracted_params["cfg"]) - if "end_at_step" in extracted_params: - end_at_step = int(extracted_params["end_at_step"]) - if "sampler_name" in extracted_params: - sampler_name = str(extracted_params["sampler_name"]).strip() - if "scheduler" in extracted_params: - scheduler = str(extracted_params["scheduler"]).strip() - - # Fallback to saved parameter snapshot if graph parse is incomplete. - if isinstance(meta, dict): - saved_params = _safe_json_loads(meta.get("mxd_source_ksampler_params")) - if isinstance(saved_params, dict): - if "steps" in saved_params and "steps" not in extracted_params: - try: - steps = int(saved_params["steps"]) - except Exception: - pass - if "cfg" in saved_params and "cfg" not in extracted_params: - try: - cfg = float(saved_params["cfg"]) - except Exception: - pass - if "end_at_step" in saved_params and "end_at_step" not in extracted_params: - try: - end_at_step = int(saved_params["end_at_step"]) - except Exception: - pass - if "sampler_name" in saved_params and "sampler_name" not in extracted_params: - try: - sampler_name = str(saved_params["sampler_name"]).strip() - except Exception: - pass - if "scheduler" in saved_params and "scheduler" not in extracted_params: - try: - scheduler = str(saved_params["scheduler"]).strip() - except Exception: - pass - - return pos, neg, steps, cfg, sampler_name, scheduler, end_at_step - -# ---------- Load a single latent (WITH Comfy params, consistent with folder version) ---------- -# ---------- Load one latent (conditioning + sampler params + optional trim_latent) ---------- -class LoadLatent_I2V_MXD: - """ - Default single-latent loader: sampler settings, CONDITIONING (positive/negative), and - an optional trim_latent value, all read from the .latent file and its .cond.pt sidecar. - """ - DESCRIPTION = """Load one latent and return conditioning and sampler settings.""" - TITLE = "Load Latent MXD" - CATEGORY = "MXD/Latents" - FUNCTION = "load" - - RETURN_TYPES = ( - "FLOAT", # shift - "CONDITIONING", # positive conditioning - "CONDITIONING", # negative conditioning - "LATENT", - "INT", - "FLOAT", - "STRING", - "STRING", - "INT", - "STRING", - "INT", # trim_latent - "STRING", # high_workflow - ) - RETURN_NAMES = ( - "shift", - "positive", - "negative", - "samples", - "steps", - "cfg", - "sampler_name", - "scheduler", - "end_at_step", - "filename_prefix", - "trim_latent", - "high_workflow", - ) - - @classmethod - def INPUT_TYPES(s): - latents_root = os.path.join(folder_paths.get_input_directory(), "latents") - os.makedirs(latents_root, exist_ok=True) - files = glob.glob(os.path.join(latents_root, "**", "*.latent"), recursive=True) - # Clean dropdown display (no "latents/" prefix) - options = [os.path.relpath(f, latents_root).replace(os.sep, "/") for f in files] - # Group by folder so the combo's prev/next arrows walk one folder at a time. - options = _sort_latent_options_by_folder(options, latents_root) - - ks_inputs = KSamplerAdvanced.INPUT_TYPES().get("required", {}) - samplers_enum = ks_inputs.get("sampler_name", ("STRING",))[0] - schedulers_enum = ks_inputs.get("scheduler", ("STRING",))[0] - - s.RETURN_TYPES = ( - "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", - "INT", "FLOAT", samplers_enum, schedulers_enum, - "INT", "STRING", "INT", "STRING", - ) - s._SAMPLERS_ENUM = samplers_enum - s._SCHEDULERS_ENUM = schedulers_enum - - return { - "required": { - "latent": (options, ), - "run_folder": ("BOOLEAN", { - "default": False, - "tooltip": "When enabled, hitting Queue Prompt auto-queues every latent in this file's folder, one after another, instead of just the selected file.", - }), - "refresh_before_run": ("BOOLEAN", { - "default": False, - "tooltip": "When run_folder is on, re-scan the latents folder for new files right before the queuing loop starts, instead of using the dropdown list as of whenever it was last populated. Use this when another workflow is still writing latents into this folder as you queue this one.", - }), - } - } - - def _coerce_enum(self, value, enum_values): - try: - return value if (enum_values and value in enum_values) else (enum_values[0] if enum_values else value) - except Exception: - return value - - def _strip_counter(self, name: str) -> str: - # Only strip the trailing pattern we generate when saving: "_<5digits>_" - # Preserve numeric-only base names like "96". - stem, _ = os.path.splitext(name) - m = re.match(r"^(.*?)(?:_\d{5}_)$", stem) - return m.group(1) if m else stem - - def _extract_sd3_shift(self, meta: dict, prompt_json: dict | None) -> float: - """ - Find SD3 'shift' in several places: - 1) flat meta["shift"] - 2) nested in prompt/workflow JSON: - - nodes[].{type|class_type} == "ModelSamplingSD3" -> inputs.shift or widgets_values[0] - - runtime-style prompt dict mapping IDs -> {..., class_type: "ModelSamplingSD3"} - Falls back to 5.0 if not found. - """ - def try_float(x): - try: - return float(x) - except Exception: - return None - - # 1) flat meta - if isinstance(meta, dict): - v = try_float(meta.get("shift")) - if v is not None: - return v - - # parse any JSON-like strings present in meta - def safe_load(x): - try: - return _safe_json_loads(x) if isinstance(x, str) else x - except Exception: - return None - - # Search helper over various JSON shapes - def search_container(obj): - # Direct dict containing shift - if isinstance(obj, dict): - if "shift" in obj: - v = try_float(obj.get("shift")) - if v is not None: - return v - - # Comfy "nodes": [ {...}, ... ] - nodes = obj.get("nodes") - if isinstance(nodes, list): - # take the last SD3 node (most recent in graph) - ms_nodes = [n for n in nodes if isinstance(n, dict) and ( - n.get("type") == "ModelSamplingSD3" or - n.get("class_type") == "ModelSamplingSD3" or - (isinstance(n.get("properties"), dict) and n["properties"].get("Node name for S&R") == "ModelSamplingSD3") - )] - if ms_nodes: - nd = ms_nodes[-1] - # Prefer explicit inputs.shift if present and literal - inp = nd.get("inputs") - if isinstance(inp, dict) and "shift" in inp: - vv = inp["shift"] - # ignore connection like [node_id, idx] - if not isinstance(vv, (list, tuple)): - v2 = try_float(vv) - if v2 is not None: - return v2 - # Fallback: first widget is shift for SD3 (as seen in your JSON) - w = nd.get("widgets_values") - if isinstance(w, list) and len(w) >= 1: - v2 = try_float(w[0]) - if v2 is not None: - return v2 - - # Runtime prompt map: {"42": {"class_type":"ModelSamplingSD3", "inputs":{...}, "widgets_values":[...]}, ...} - # Heuristic: values that are dicts with class_type keys - has_ct = [v for v in obj.values() if isinstance(v, dict) and "class_type" in v] - if has_ct: - for nd in has_ct: - if nd.get("class_type") == "ModelSamplingSD3": - inp = nd.get("inputs", {}) - if isinstance(inp, dict) and "shift" in inp: - vv = inp["shift"] - if not isinstance(vv, (list, tuple)): - v2 = try_float(vv) - if v2 is not None: - return v2 - w = nd.get("widgets_values") - if isinstance(w, list) and len(w) >= 1: - v2 = try_float(w[0]) - if v2 is not None: - return v2 - - # Lists / nested - if isinstance(obj, list): - for it in obj: - v = search_container(it) - if v is not None: - return v - return None - - # 2) Look in provided prompt_json - v = search_container(prompt_json) - if v is not None: - return v - - # Also look in common meta fields that can hold the full workflow/prompt - for key in ("workflow", "prompt", "extra_pnginfo"): - candidate = meta.get(key) - cand_obj = safe_load(candidate) - if isinstance(cand_obj, dict) or isinstance(cand_obj, list): - v = search_container(cand_obj) - if v is not None: - return v - # extra_pnginfo can nest "workflow"/"prompt" again - if isinstance(cand_obj, dict): - for subkey in ("workflow", "prompt"): - sub = safe_load(cand_obj.get(subkey)) - if isinstance(sub, dict) or isinstance(sub, list): - v = search_container(sub) - if v is not None: - return v - - # default - return 5.0 - - @classmethod - def IS_CHANGED(s, latent): - # Fix path lookup (add "latents/" prefix back) - p = folder_paths.get_annotated_filepath(f"latents/{latent}") - m = hashlib.sha256() - with open(p, "rb") as f: - m.update(f.read()) - side = p.replace(".latent", ".cond.pt") - if os.path.exists(side): - with open(side, "rb") as f: - m.update(f.read()) - return m.digest().hex() - - @classmethod - def VALIDATE_INPUTS(s, latent): - check_path = latent if latent.startswith("latents/") else f"latents/{latent}" - try: - folder_paths.get_annotated_filepath(check_path) - except Exception: - return f"Invalid latent file: {latent}" - return True - - def load(self, latent, run_folder=False, refresh_before_run=False): - # Ensure we prepend "latents/" if missing, but don't duplicate it - if not latent.startswith("latents/"): - latent_path = folder_paths.get_annotated_filepath(f"latents/{latent}") - else: - latent_path = folder_paths.get_annotated_filepath(latent) - - sample_dict, meta, _ = _load_latent_file(latent_path) - t = sample_dict["samples"] - - if isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) > 1: - samples = {"samples": t[0:1].contiguous()} - elif isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) == 1: - samples = {"samples": t} - else: - samples = {"samples": t.unsqueeze(0)} - - prompt_json = _safe_json_loads(meta.get("prompt")) - _pos_text, _neg_text, steps, cfg, sampler_name, scheduler, end_at_step = \ - _extract_params_from_prompt_json(prompt_json or {}, meta) - - # SD3 shift (not in KSamplerAdvanced, but we want it) - shift = self._extract_sd3_shift(meta, prompt_json) - - sampler_name = self._coerce_enum(sampler_name, getattr(self.__class__, "_SAMPLERS_ENUM", ())) - scheduler = self._coerce_enum(scheduler, getattr(self.__class__, "_SCHEDULERS_ENUM", ())) - - def normalize_folder(part: str) -> str: - part = part.replace("\\", "/").strip("/") - if not part: - return "" - segments = [seg for seg in part.split("/") if seg] - if segments and segments[0].lower() == "latents": - segments = segments[1:] - return "/".join(segments) - - folder_part = normalize_folder(os.path.dirname(latent)) - base_name = os.path.basename(latent_path) - clean_stem = self._strip_counter(base_name) - prefix = f"{folder_part}/{clean_stem}" if folder_part else clean_stem - - # Raw workflow JSON embedded when this latent was saved (empty string if none). - source_workflow = meta.get("workflow") or "" - - positive_conditioning, negative_conditioning, sidecar = _load_i2v_conditioning_sidecar(latent_path) - trim_latent = _coerce_trim_latent(sidecar.get("trim_latent", meta.get("trim_latent", 0))) - - return ( - float(shift), - positive_conditioning, - negative_conditioning, - samples, - int(steps), - float(cfg), - sampler_name, - scheduler, - int(end_at_step), - prefix, - trim_latent, - source_workflow, - ) - -# ---------- Load multiple latents from a folder (conditioning + sampler params + optional trim_latent) ---------- -class LoadLatents_FromFolder_I2V_MXD: - """ - Default folder/batch loader: same outputs as LoadLatent_I2V_MXD, one set per latent - found in the folder. - """ - DESCRIPTION = """Load all latents in a folder with conditioning and sampler settings.""" - TITLE = "Load Latent Batch MXD" - CATEGORY = "MXD/Latents" - FUNCTION = "load_batch_i2v" - - RETURN_TYPES = ( - "FLOAT", # shift - "CONDITIONING", # positive conditioning - "CONDITIONING", # negative conditioning - "LATENT", - "INT", - "FLOAT", - "STRING", # will be replaced with sampler enum in INPUT_TYPES - "STRING", # will be replaced with scheduler enum in INPUT_TYPES - "INT", - "STRING", - "INT", # trim_latent - ) - RETURN_NAMES = ( - "shift", - "positive", - "negative", - "samples", - "steps", - "cfg", - "sampler_name", - "scheduler", - "end_at_step", - "filename_prefix", - "trim_latent", - ) - - OUTPUT_IS_LIST = (True,) * 11 - - @classmethod - def INPUT_TYPES(s): - # Same folder logic as the single loader - latents_root = os.path.join(folder_paths.get_input_directory(), "latents") - os.makedirs(latents_root, exist_ok=True) - subs = [""] + _list_latent_subfolders(latents_root) - - # Pull live enums from KSamplerAdvanced so sampler/scheduler wire cleanly - from nodes import KSamplerAdvanced - ks_inputs = KSamplerAdvanced.INPUT_TYPES().get("required", {}) - samplers_enum = ks_inputs.get("sampler_name", ("STRING",))[0] - schedulers_enum = ks_inputs.get("scheduler", ("STRING",))[0] - - s.RETURN_TYPES = ( - "FLOAT", # shift - "CONDITIONING", # positive conditioning - "CONDITIONING", # negative conditioning - "LATENT", - "INT", - "FLOAT", - samplers_enum, # enum type for sampler_name - schedulers_enum, # enum type for scheduler - "INT", - "STRING", - "INT", - ) - s._SAMPLERS_ENUM = samplers_enum - s._SCHEDULERS_ENUM = schedulers_enum - - return {"required": {"subfolder": (subs, )}} - - def _coerce_enum(self, value, enum_values): - try: - return value if (enum_values and value in enum_values) else (enum_values[0] if enum_values else value) - except Exception: - return value - - def _strip_counter(self, name: str) -> str: - stem, _ = os.path.splitext(name) - m = re.match(r"^(.*?)(?:_\d{5}_)$", stem) - return m.group(1) if m else stem - - def _extract_sd3_shift(self, meta: dict, prompt_json: dict | None) -> float: - def try_float(x): - try: return float(x) - except Exception: return None - - if isinstance(meta, dict): - v = try_float(meta.get("shift")) - if v is not None: return v - - def safe_load(x): - try: return _safe_json_loads(x) if isinstance(x, str) else x - except Exception: return None - - def search_container(obj): - if isinstance(obj, dict): - if "shift" in obj: - v = try_float(obj.get("shift")) - if v is not None: return v - nodes = obj.get("nodes") - if isinstance(nodes, list): - ms_nodes = [n for n in nodes if isinstance(n, dict) and ( - n.get("type") == "ModelSamplingSD3" or - n.get("class_type") == "ModelSamplingSD3" or - (isinstance(n.get("properties"), dict) and n["properties"].get("Node name for S&R") == "ModelSamplingSD3") - )] - if ms_nodes: - nd = ms_nodes[-1] - inp = nd.get("inputs") - if isinstance(inp, dict) and "shift" in inp: - vv = inp["shift"] - if not isinstance(vv, (list, tuple)): - v2 = try_float(vv) - if v2 is not None: return v2 - w = nd.get("widgets_values") - if isinstance(w, list) and len(w) >= 1: - v2 = try_float(w[0]) - if v2 is not None: return v2 - has_ct = [v for v in obj.values() if isinstance(v, dict) and "class_type" in v] - for nd in has_ct: - if nd.get("class_type") == "ModelSamplingSD3": - inp = nd.get("inputs", {}) - if isinstance(inp, dict) and "shift" in inp: - vv = inp["shift"] - if not isinstance(vv, (list, tuple)): - v2 = try_float(vv) - if v2 is not None: return v2 - w = nd.get("widgets_values") - if isinstance(w, list) and len(w) >= 1: - v2 = try_float(w[0]) - if v2 is not None: return v2 - if isinstance(obj, list): - for it in obj: - v = search_container(it) - if v is not None: return v - return None - - v = search_container(prompt_json) - if v is not None: return v - - for key in ("workflow", "prompt", "extra_pnginfo"): - candidate = meta.get(key) - cand_obj = safe_load(candidate) - if isinstance(cand_obj, (dict, list)): - v = search_container(cand_obj) - if v is not None: return v - if isinstance(cand_obj, dict): - for subkey in ("workflow", "prompt"): - sub = safe_load(cand_obj.get(subkey)) - if isinstance(sub, (dict, list)): - v = search_container(sub) - if v is not None: return v - - return 5.0 - - def load_batch_i2v(self, subfolder): - latents_root = os.path.join(folder_paths.get_input_directory(), "latents") - base = os.path.join(latents_root, subfolder) if subfolder else latents_root - files = glob.glob(os.path.join(base, "**", "*.latent"), recursive=True) - files = _sort_paths_newest_first(files) - if not files: - raise RuntimeError(f"[LoadLatents_FromFolder_I2V_MXD] No .latent files found in '{base}'.") - - shifts, samples_list = [], [] - positives, negatives = [], [] - steps_list, cfgs, samplers, schedulers, end_steps = [], [], [], [], [] - filename_prefixes, trims = [], [] - - for path in files: - sample_dict, meta, _ = _load_latent_file(path) - t = sample_dict["samples"] - - if isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) > 1: - slices = [t[i:i+1].contiguous() for i in range(t.size(0))] - else: - slices = [t if (isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) == 1) - else t.unsqueeze(0)] - - prompt_json = _safe_json_loads(meta.get("prompt")) - _pos_text, _neg_text, n_steps, cfg, sampler_name, scheduler, end_at_step = \ - _extract_params_from_prompt_json(prompt_json or {}, meta) - - sampler_name = self._coerce_enum(sampler_name, getattr(self.__class__, "_SAMPLERS_ENUM", ())) - scheduler = self._coerce_enum(scheduler, getattr(self.__class__, "_SCHEDULERS_ENUM", ())) - shift_val = self._extract_sd3_shift(meta, prompt_json) - - positive_conditioning, negative_conditioning, sidecar = _load_i2v_conditioning_sidecar(path) - trim_latent = _coerce_trim_latent(sidecar.get("trim_latent", meta.get("trim_latent", 0))) - - folder_part = subfolder if subfolder else "" - clean_stem = self._strip_counter(os.path.basename(path)) - prefix = os.path.join(folder_part, clean_stem) if folder_part else clean_stem - - for sl in slices: - shifts.append(float(shift_val)) - positives.append(positive_conditioning) - negatives.append(negative_conditioning) - samples_list.append({"samples": sl}) - steps_list.append(int(n_steps)) - cfgs.append(float(cfg)) - samplers.append(sampler_name) - schedulers.append(scheduler) - end_steps.append(int(end_at_step)) - filename_prefixes.append(prefix) - trims.append(trim_latent) - - return ( - shifts, - positives, - negatives, - samples_list, - steps_list, - cfgs, - samplers, - schedulers, - end_steps, - filename_prefixes, - trims, - ) - -# ---------- Pipe variants: bundle all the loader outputs into one wire ---------- -class LoadLatent_I2V_Pipe_MXD(LoadLatent_I2V_MXD): - """ - Same loading logic as LoadLatent_I2V_MXD, but bundles every value into a single - MXD_LATENT_PIPE output so switching between the single/batch loaders is a one-wire swap. - Unpack with LatentPipeUnpack_MXD. - """ - TITLE = "Load Latent Pipe MXD" - CATEGORY = "MXD/Latents" - FUNCTION = "load_pipe" - - RETURN_TYPES = ("MXD_LATENT_PIPE",) - RETURN_NAMES = ("latent_pipe",) - - @classmethod - def INPUT_TYPES(s): - inputs = LoadLatent_I2V_MXD.INPUT_TYPES.__func__(s) - s.RETURN_TYPES = ("MXD_LATENT_PIPE",) - return inputs - - def load_pipe(self, latent, run_folder=False, refresh_before_run=False): - ( - shift, positive, negative, samples, - steps, cfg, sampler_name, scheduler, - end_at_step, prefix, trim_latent, source_workflow, - ) = self.load(latent, run_folder, refresh_before_run) - - pipe = { - "shift": shift, - "positive": positive, - "negative": negative, - "samples": samples, - "steps": steps, - "cfg": cfg, - "sampler_name": sampler_name, - "scheduler": scheduler, - "end_at_step": end_at_step, - "filename_prefix": prefix, - "trim_latent": trim_latent, - "high_workflow": source_workflow, - } - return (pipe,) - - -class LoadLatents_FromFolder_I2V_Pipe_MXD(LoadLatents_FromFolder_I2V_MXD): - """ - Same loading logic as LoadLatents_FromFolder_I2V_MXD, but bundles every value into a - single MXD_LATENT_PIPE output per item. Unpack with LatentPipeUnpack_MXD. - """ - TITLE = "Load Latent Batch Pipe MXD" - CATEGORY = "MXD/Latents" - FUNCTION = "load_batch_pipe" - - RETURN_TYPES = ("MXD_LATENT_PIPE",) - RETURN_NAMES = ("latent_pipe",) - OUTPUT_IS_LIST = (True,) - - @classmethod - def INPUT_TYPES(s): - inputs = LoadLatents_FromFolder_I2V_MXD.INPUT_TYPES.__func__(s) - s.RETURN_TYPES = ("MXD_LATENT_PIPE",) - return inputs - - def load_batch_pipe(self, subfolder): - ( - shifts, positives, negatives, samples_list, - steps_list, cfgs, samplers, schedulers, - end_steps, filename_prefixes, trims, - ) = self.load_batch_i2v(subfolder) - - pipes = [] - for i in range(len(samples_list)): - pipes.append({ - "shift": shifts[i], - "positive": positives[i], - "negative": negatives[i], - "samples": samples_list[i], - "steps": steps_list[i], - "cfg": cfgs[i], - "sampler_name": samplers[i], - "scheduler": schedulers[i], - "end_at_step": end_steps[i], - "filename_prefix": filename_prefixes[i], - "trim_latent": trims[i], - }) - return (pipes,) - - -class LatentPipeUnpack_MXD: - """ - Splits an MXD_LATENT_PIPE back into shift, conditioning, samples, and sampler settings. - Works with any MXD latent pipe loader (single or batch, I2V or VACE 2.2) - missing - fields like trim_latent just fall back to a safe default. - """ - DESCRIPTION = """Split a latent pipe back into shift, positive, negative, samples, and sampler settings.""" - TITLE = "Unpack Latent Pipe MXD" - CATEGORY = "MXD/Latents" - FUNCTION = "unpack" - - RETURN_TYPES = ( - "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", - "INT", "FLOAT", "STRING", "STRING", "INT", "STRING", "INT", "STRING", - ) - RETURN_NAMES = ( - "shift", "positive", "negative", "samples", - "steps", "cfg", "sampler_name", "scheduler", - "end_at_step", "filename_prefix", "trim_latent", "high_workflow", - ) - - @classmethod - def INPUT_TYPES(s): - from nodes import KSamplerAdvanced - ks_inputs = KSamplerAdvanced.INPUT_TYPES().get("required", {}) - samplers_enum = ks_inputs.get("sampler_name", ("STRING",))[0] - schedulers_enum = ks_inputs.get("scheduler", ("STRING",))[0] - - s.RETURN_TYPES = ( - "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", - "INT", "FLOAT", samplers_enum, schedulers_enum, - "INT", "STRING", "INT", "STRING", - ) - s._SAMPLERS_ENUM = samplers_enum - s._SCHEDULERS_ENUM = schedulers_enum - - return {"required": {"latent_pipe": ("MXD_LATENT_PIPE",)}} - - def _coerce_enum(self, value, enum_values): - try: - return value if (enum_values and value in enum_values) else (enum_values[0] if enum_values else value) - except Exception: - return value - - def unpack(self, latent_pipe): - sampler_name = self._coerce_enum(latent_pipe.get("sampler_name"), getattr(self.__class__, "_SAMPLERS_ENUM", ())) - scheduler = self._coerce_enum(latent_pipe.get("scheduler"), getattr(self.__class__, "_SCHEDULERS_ENUM", ())) - - return ( - latent_pipe.get("shift", 0.0), - latent_pipe.get("positive", []), - latent_pipe.get("negative", []), - latent_pipe.get("samples"), - latent_pipe.get("steps", 0), - latent_pipe.get("cfg", 0.0), - sampler_name, - scheduler, - latent_pipe.get("end_at_step", 0), - latent_pipe.get("filename_prefix", ""), - latent_pipe.get("trim_latent", 0), - latent_pipe.get("high_workflow", ""), - ) - -# ---------- Empty latent image generator (for video nodes) ---------- -class Wan2_2EmptyLatentImageMXD: - """ - Utility node for WAN 2.2 workflows. - Generates an empty latent tensor at common video-friendly resolutions. - """ - - DESCRIPTION = """Create an empty WAN 2.2 latent at a preset resolution.""" - TITLE = "WAN2.2 Empty Latent Image" - CATEGORY = "WAN2.2/Latent" - - RESOLUTIONS = { - "— 720p —": None, - "Widescreen (16:9) 1280×720": (1280, 720), - - "— 480p —": None, - "Widescreen (16:9) 832×480": (832, 480), - "Square (1:1) 624×624": (624, 624), - } - - RETURN_TYPES = ("LATENT",) - FUNCTION = "generate" - - @classmethod - def INPUT_TYPES(cls): - options = list(cls.RESOLUTIONS.keys()) - return { - "required": { - "resolution": ( - options, - {"default": "Square (1:1) 960×960", "tooltip": "Select target resolution preset."} - ), - "vertical": ( - "BOOLEAN", - {"default": False, "label_on": "Vertical", "label_off": "Landscape", - "tooltip": "Swap width/height for vertical orientation."} - ), - "batch_size": ( - "INT", - {"default": 1, "min": 1, "max": 4096, "tooltip": "Number of latents to generate."} - ), - } - } - - def generate(self, resolution, vertical, batch_size): - size = self.RESOLUTIONS.get(resolution) - if size is None: - raise ValueError(f"'{resolution}' is a header or invalid option.") - - w, h = size - if vertical: - w, h = h, w - - # Safety: ensure divisible by 8 - if (w % 8) or (h % 8): - raise ValueError(f"Resolution must be divisible by 8. Got {w}x{h}.") - - # WAN video length always t=1 - t = 1 - - latent = torch.zeros( - [batch_size, 16, t, h // 8, w // 8], - device=comfy.model_management.intermediate_device() - ) - return ({"samples": latent},) - -# ---------- Empty latent video generator with presets (for video nodes) ---------- -class wan22EmptyHunyuanLatentVideoMXD: - """ - Exactly like core EmptyHunyuanLatentVideo, but width/height are replaced - with valid WAN 2.2 resolution presets and a vertical toggle. - """ - - RETURN_TYPES = ("LATENT",) - FUNCTION = "generate" - CATEGORY = "latent/video" - - # ✅ Cleaned, WAN 2.2–accurate presets - RESOLUTIONS = { - "— 720p —": None, - "Widescreen (16:9) 1280×720": (1280, 720), - "Square (1:1) 1024×1024": (1024, 1024), - - "— 480p —": None, - "Widescreen (16:9) 832×480": (832, 480), - "Square (1:1) 624×624": (624, 624), - } - - @classmethod - def INPUT_TYPES(cls): - options = list(cls.RESOLUTIONS.keys()) - return { - "required": { - "resolution": ( - options, - {"default": "Widescreen (16:9) 832×480"} - ), - "vertical": ( - "BOOLEAN", - {"default": False, "label_on": "Vertical", "label_off": "Landscape"} - ), - "length": ( - "INT", - {"default": 81, "min": 1, "max": nodes.MAX_RESOLUTION, "step": 4} - ), - "batch_size": ( - "INT", - {"default": 1, "min": 1, "max": 4096} - ), - } - } - - def generate(self, resolution, vertical, length, batch_size): - size = self.RESOLUTIONS.get(resolution) - if size is None: - raise ValueError(f"'{resolution}' is not a selectable resolution.") - w, h = size - if vertical: - w, h = h, w - - # identical to core behavior: - t = ((length - 1) // 4) + 1 - latent = torch.zeros( - [batch_size, 16, t, h // 8, w // 8], - device=comfy.model_management.intermediate_device() - ) - return ({"samples": latent},) -# ---------- WAN 2.2 Image to Video (no scaling; expects pre-sized input) ---------- -if HAVE_COMFY_API: - class Wan22ImageToVideoMXD(io.ComfyNode): - @classmethod - def define_schema(cls): - return io.Schema( - node_id="Wan22ImageToVideoMXD", - display_name="WAN 2.2 Image to Video MXD", - category="conditioning/video_models", - description="WAN 2.2 image to video without scaling or CLIP vision.", - inputs=[ - io.Conditioning.Input("positive"), - io.Conditioning.Input("negative"), - io.Vae.Input("vae"), - io.Int.Input("length", default=81, min=1, max=16384, step=4), - io.Int.Input("batch_size", default=1, min=1, max=4096), - io.Image.Input("start_image", optional=False), - ], - 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) -> io.NodeOutput: - if start_image is None: - raise ValueError("start_image must be provided (already pre-sized).") - - frames_in, ih, iw, ch = start_image.shape - frames_used = min(frames_in, length) - t = ((length - 1) // 4) + 1 - - latent = torch.zeros( - [batch_size, 16, t, ih // 8, iw // 8], - device=comfy.model_management.intermediate_device() - ) - - # create placeholder image tensor - image = torch.ones( - (length, ih, iw, ch), - device=start_image.device, - dtype=start_image.dtype - ) * 0.5 - image[:frames_used] = start_image[:frames_used] - - # encode using VAE - concat_latent_image = vae.encode(image[:, :, :, :3]) - - # mask zeros out the frames used - mask = torch.ones( - (1, 1, t, concat_latent_image.shape[-2], concat_latent_image.shape[-1]), - device=image.device, - dtype=image.dtype - ) - mask[:, :, :((frames_used - 1) // 4) + 1] = 0.0 - - positive = node_helpers.conditioning_set_values( - positive, {"concat_latent_image": concat_latent_image, "concat_mask": mask} - ) - negative = node_helpers.conditioning_set_values( - negative, {"concat_latent_image": concat_latent_image, "concat_mask": mask} - ) - - out_latent = {"samples": latent} - return io.NodeOutput(positive, negative, out_latent) - -# ---- Canonical WAN 2.2 buckets ---- -BUCKETS_480 = [(832,480), (480,832), (624,624)] # 16:9, 9:16, 1:1 -BUCKETS_720 = [(1280,720), (720,1280), (1024,1024)] # 16:9, 9:16, 1:1 -SQUARE_TOL = 0.03 # exact-ish square passthrough tolerance -AUTO_SQUARE_MAX_AR = 1.25 # Auto may crop to square when the source is within 25% of 1:1. - -def _ar(w, h): - return w / max(1, h) - -def _safe_hw(w, h): - w = max(16, min(w, nodes.MAX_RESOLUTION)) - h = max(16, min(h, nodes.MAX_RESOLUTION)) - return w, h - -def _floor16(x): - x = int(x) // 16 * 16 - return max(16, x) - -def _ceil16(x): - x = (int(x) + 15) // 16 * 16 - return max(16, x) - -def _is_squareish(w, h, tol=SQUARE_TOL): - r = _ar(w, h) - return abs(r - 1.0) <= tol - -def _is_auto_square_candidate(w, h): - r = _ar(w, h) - return max(r, 1.0 / max(r, 1e-9)) <= AUTO_SQUARE_MAX_AR - -def _wan22_tier_from_area(iw, ih): - area = iw * ih - area_480 = 832 * 480 - area_720 = 1280 * 720 - return "480p" if abs(area - area_480) / area_480 <= abs(area - area_720) / area_720 else "720p" - -def _wan22_square_bucket(tier, iw=None, ih=None): - if tier == "720p": - return (1024, 1024) - if tier == "480p": - return (624, 624) - return (1024, 1024) if _wan22_tier_from_area(iw, ih) == "720p" else (624, 624) - -def _wan22_oriented_bucket(tier, orientation, iw=None, ih=None): - if tier == "Auto": - tier = _wan22_tier_from_area(iw, ih) - if orientation == "Tall": - return (480, 832) if tier == "480p" else (720, 1280) - if orientation == "Wide": - return (832, 480) if tier == "480p" else (1280, 720) - return _wan22_square_bucket(tier, iw, ih) - -def _closest_bucket(img_w, img_h, bucket_list, cover=False): - """ - Pick the best (bw,bh) from bucket_list for this image. - Uses scale closeness + AR diff to rank. - """ - in_ar = _ar(img_w, img_h) - best, best_key = None, (float("inf"), 0.0) - for bw, bh in bucket_list: - s = max(bw/img_w, bh/img_h) if cover else min(bw/img_w, bh/img_h) - ar_diff = abs(_ar(bw, bh) - in_ar) - key = (abs(1.0 - s), ar_diff) - if key < best_key: - best_key, best = key, (bw, bh) - return best - -def _resize_then_center_crop(img, out_w, out_h): - """ - Resize to cover target (ensures >= target on both sides after ceil16), - then center-crop. No padding. - """ - t, ih, iw, c = img.shape - s = max(out_w / iw, out_h / ih) - tw = _ceil16(iw * s) - th = _ceil16(ih * s) - tmp = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) - y0 = max(0, (th - out_h) // 2) - x0 = max(0, (tw - out_w) // 2) - return tmp[:, y0:y0+out_h, x0:x0+out_w, :] - -def _resize_fit_inside(img, out_w, out_h): - """ - Resize to fit inside target (ensures <= target on both sides via floor16), - and return the resized tensor only. No padding. - """ - t, ih, iw, c = img.shape - s = min(out_w / iw, out_h / ih) - tw = _floor16(iw * s) - th = _floor16(ih * s) - tw, th = _safe_hw(tw, th) - resized = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) - return resized, tw, th - -def _validate_image_batch_4d(image, node_name, input_name): - if image is None: - raise ValueError(f"[{node_name}] '{input_name}' is required.") - if not torch.is_tensor(image): - raise TypeError(f"[{node_name}] '{input_name}' must be an IMAGE torch tensor, got {type(image).__name__}.") - if image.ndim != 4: - raise ValueError(f"[{node_name}] '{input_name}' must have shape [T,H,W,C], got {tuple(image.shape)}.") - if image.shape[0] <= 0: - raise ValueError(f"[{node_name}] '{input_name}' contains zero images/frames.") - if image.shape[1] <= 0 or image.shape[2] <= 0 or image.shape[3] <= 0: - raise ValueError(f"[{node_name}] '{input_name}' has invalid dimensions {tuple(image.shape)}.") - return image - -def _resize_to_explicit_resolution(img, out_w, out_h, match_mode="crop_to_match"): - """ - Resize IMAGE batch to an explicit resolution. - - crop_to_match: cover + center crop (exact output) - - fit_inside_only: preserve AR, no crop (may be smaller) - - stretch_exact: force exact output (distorts AR) - """ - out_w = int(out_w) - out_h = int(out_h) - if out_w <= 0 or out_h <= 0: - raise ValueError(f"Invalid target resolution {out_w}x{out_h}.") - - if match_mode == "crop_to_match": - return _resize_then_center_crop(img, out_w, out_h) - - if match_mode == "fit_inside_only": - _, ih, iw, _ = img.shape - s = min(out_w / max(1, iw), out_h / max(1, ih)) - tw = max(1, min(out_w, int(iw * s))) - th = max(1, min(out_h, int(ih * s))) - return comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) - - if match_mode == "stretch_exact": - return comfy.utils.common_upscale(img.movedim(-1, 1), out_w, out_h, "bilinear", "center").movedim(1, -1) - - raise ValueError( - f"Invalid match_mode '{match_mode}'. Expected one of: crop_to_match, fit_inside_only, stretch_exact." - ) - -# ---------- WAN22_I2V_Image_Scaler_MXD ---------- -# Adds a new “Safe Auto” mode for video extend workflows. -# Normal modes (Auto / 480p / 720p) behave exactly as before. -# “Safe Auto” adds passthrough + strict checks to prevent failures on WAN 2.2 extend. - -_WAN22_VALID_RES = { - (832, 480), (480, 832), - (1280, 720), (720, 1280), - (624, 624), (1024, 1024), -} - -def _wan22_is_valid_dim(w, h): - return (w, h) in _WAN22_VALID_RES - - -def _wan22_pick_bucket(iw, ih, tier, crop_to_fit, aspect_mode="Auto"): - if tier == "Safe Auto": - tier = "Auto" - - if aspect_mode in ("Tall", "Wide", "Square"): - return _wan22_oriented_bucket(tier, aspect_mode, iw, ih) - - is_squareish = _is_squareish(iw, ih) - is_landscape = iw >= ih - - # --- Square handling --- - if is_squareish or (crop_to_fit and _is_auto_square_candidate(iw, ih)): - return _wan22_square_bucket(tier, iw, ih) - - # --- 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, aspect_mode="Auto"): - """ - 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 1024x1024\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, aspect_mode=aspect_mode) - 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 _resample_video_frames_to_fps(frames, in_fps, out_fps): - """ - Resample a frame sequence to a target FPS using nearest-frame selection. - Preserves clip duration approximately by dropping/duplicating frames, - instead of only changing FPS metadata (which changes playback speed). - Returns (frames_out, fps_out, changed). - """ - if frames is None or frames.ndim != 4: - raise ValueError("Expected frame tensor with shape [T,H,W,C].") - - if in_fps is None: - raise ValueError("Input video FPS is missing; cannot force FPS safely.") - - in_fps = float(in_fps) - out_fps = float(out_fps) - if in_fps <= 0: - raise ValueError(f"Invalid input FPS: {in_fps}") - if out_fps <= 0: - raise ValueError(f"Invalid target FPS: {out_fps}") - - if frames.shape[0] <= 1: - return frames, float(out_fps), False - - if abs(in_fps - out_fps) < 1e-6: - return frames, float(out_fps), False - - n_in = int(frames.shape[0]) - # Match the first/last frame span, then pick nearest frames on that timeline. - n_out = max(1, int(round(((n_in - 1) * out_fps) / in_fps)) + 1) - if n_out == n_in: - # Frame count may stay the same for near-equal FPS; metadata still becomes exact. - return frames, float(out_fps), False - - idx = torch.linspace(0, n_in - 1, steps=n_out, device=frames.device) - idx = idx.round().to(dtype=torch.long) - out = frames.index_select(0, idx) - return out, float(out_fps), True - - -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) - - Modes: Auto / 480p / 720p (legacy "Safe Auto" still accepted) - - Fit (no pad): proportional resize ≤ target; returns resized dims. - - Crop (no pad): resize-to-cover then center-crop to exact target. - - Square handling: - * Auto & 480p: ~square → 624×624 - * 720p: ~square -> 1024x1024 - - “Safe Auto”: - * If input is already a valid WAN 2.2 bucket, passthrough. - * If input is far outside 480p–720p range, error early. - * Otherwise, same logic as Auto. - * Perfect for video-extend workflows. - """ - - TITLE = "Image Bucket Scaler MXD (No Pad)" - CATEGORY = "image/processing" - RETURN_TYPES = ("IMAGE",) - FUNCTION = "scale" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "tier": (["Auto", "480p", "720p"], {"default": "Auto"}), - "crop_to_fit": ("BOOLEAN", { - "default": True, - "label_on": "Perfect Fit (Crops Edges)", - "label_off": "Closest Fit (No Crop)" - }), - "aspect_mode": (["Auto", "Tall", "Wide", "Square"], { - "default": "Auto", - "tooltip": "Auto picks wide/tall/square from the source. Use Square/Tall/Wide to force the target bucket shape." - }), - } - } - - # ----------------------------- - # Internal helpers - # ----------------------------- - def _pick_bucket(self, 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 (1024, 1024) - else: - 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) - - # ----------------------------- - # Main function - # ----------------------------- - def scale(self, image, tier="Auto", crop_to_fit=False, aspect_mode="Auto"): - # Keep legacy "Safe Auto" values from old workflows working, but expose only one Auto in UI. - internal_tier = "Safe Auto" if tier == "Auto" else tier - out, _, _, _ = _wan22_scale_image_core( - image, - tier=internal_tier, - crop_to_fit=crop_to_fit, - aspect_mode=aspect_mode, - ) - return (out,) - - _, 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,) - - 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 ≈ 832×480 (or 480×832)\n" - " • 720p tier ≈ 1280×720 (or 720×1280)\n" - " • Squares: 624×624 or 1024×1024\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 = self._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,) - -class WAN22_I2V_Match_Resolution_MXD: - """ - Match a second image (or image batch) to a reference image resolution for WAN 2.2 - first/last-frame workflows. - """ - TITLE = "WAN 2.2 I2V Match Resolution" - CATEGORY = "image/processing" - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("matched_image",) - FUNCTION = "match_resolution" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "reference_image": ("IMAGE", { - "tooltip": "Reference size source (usually the first image after WAN bucket scaling)." - }), - "image_to_match": ("IMAGE", { - "tooltip": "Image or batch to resize using the reference image resolution." - }), - "match_mode": (["crop_to_match", "fit_inside_only", "stretch_exact"], { - "default": "crop_to_match", - "tooltip": "crop_to_match = exact size via cover+center crop; fit_inside_only = no crop, may be smaller; stretch_exact = exact size with distortion." - }), - "enforce_wan_bucket": ("BOOLEAN", { - "default": False, - "label_on": "Validate WAN Bucket", - "label_off": "No WAN Validation", - "tooltip": "If enabled, reference_image must already be a WAN 2.2 bucket size." - }), - } - } - - def match_resolution(self, reference_image, image_to_match, match_mode="crop_to_match", enforce_wan_bucket=False): - node_name = "WAN22_I2V_Match_Resolution_MXD" - reference_image = _validate_image_batch_4d(reference_image, node_name, "reference_image") - image_to_match = _validate_image_batch_4d(image_to_match, node_name, "image_to_match") - - _, ref_h, ref_w, _ = reference_image.shape - - if enforce_wan_bucket and not _wan22_is_valid_dim(ref_w, ref_h): - raise ValueError( - f"[{node_name}] Reference image resolution {ref_w}x{ref_h} is not a valid WAN 2.2 bucket.\n" - "Valid WAN 2.2 buckets are:\n" - " - 832x480 / 480x832\n" - " - 1280x720 / 720x1280\n" - " - 624x624 / 1024x1024\n\n" - "Recommended workflow:\n" - " 1. Scale the first image with 'Image Scaler Wan 2.2 I2V MXD'\n" - " 2. Use this node to match the second image to the scaled first image" - ) - - matched = _resize_to_explicit_resolution( - image_to_match, - out_w=ref_w, - out_h=ref_h, - match_mode=match_mode, - ) - return (matched,) - -# ---------- MXD Frames Select Start/End (from start or end of sequence) ---------- -class Frames_Select_StartEnd_MXD: - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "frames": ("IMAGE",), - "count": ("INT", { - "default": 1, - "min": 1, - "max": 10000, - "tooltip": "Number of frames to select" - }), - "offset": ("INT", { - "default": 1, - "min": 1, - "max": 10000, - "tooltip": "How far into the video to start selection (from start or end)" - }), - "mode": (["start", "end"], { - "default": "end", - "tooltip": "Select frames from the start or end of the sequence" - }), - }, - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("image",) - FUNCTION = "main" - 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 - 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() - else: # mode == "end" - start_idx = max(0, total - offset - count + 1) - end_idx = start_idx + count - selected = frames[start_idx:end_idx].clone() - - return (selected,) - -# ---------- MXD Frames Select Start/End (from start or end of sequence) ---------- -class Frames_Remove_From_Start_MXD: - def __init__(self): - pass - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "frames": ("IMAGE",), - "count": ("INT", { - "default": 10, - "min": 1, - "max": 10000, - "tooltip": "Number of frames to remove from the start" - }), - }, - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("image",) - FUNCTION = "main" - CATEGORY = "MXD/images" - - def main(self, frames=None, count=10): - # ✅ Skip the first `count` frames instead of keeping them - frames_after = frames[count:].clone() - return (frames_after,) - - -if HAVE_COMFY_API: - class CombineVideos_MXD: - """ - Combine two VIDEO inputs end-to-end (sequentially). - """ - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "front_video": ("VIDEO", {"tooltip": "The first video (plays first)"}), - "back_video": ("VIDEO", {"tooltip": "The second video (plays after the first)"}), - }, - } - - RETURN_TYPES = ("VIDEO",) - RETURN_NAMES = ("video",) - FUNCTION = "combine" - CATEGORY = "MXD/video" - - def combine(self, front_video, back_video): - comp_a = front_video.get_components() - comp_b = back_video.get_components() - - # Check frame rate consistency - if comp_a.frame_rate != comp_b.frame_rate: - raise ValueError(f"FPS mismatch: {comp_a.frame_rate} vs {comp_b.frame_rate}") - - # ✅ 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: - 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 - - - - combined_video = VideoFromComponents( - VideoComponents( - images=combined_images, - audio=combined_audio, - frame_rate=comp_a.frame_rate, - ) - ) - - 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 the scaled frame batch directly - - keep default workflow simple for common use - """ - CATEGORY = "MXD/video" - FUNCTION = "prepare" - RETURN_TYPES = ("VIDEO", "IMAGE", "FLOAT") - RETURN_NAMES = ("scaled_video", "images", "fps") - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "video": ("VIDEO",), - "tier": (["Auto", "480p", "720p"], {"default": "Auto"}), - "crop_to_fit": ("BOOLEAN", { - "default": True, - "label_on": "Perfect Fit (Crops Edges)", - "label_off": "Closest Fit (No Crop)" - }), - "force_fps": ("BOOLEAN", { - "default": False, - "label_on": "Force FPS", - "label_off": "Keep Source FPS", - "tooltip": "When enabled, resample frames (drop/duplicate) and set exact target fps." - }), - "target_fps": ("INT", { - "default": 16, - "min": 1, - "max": 1000, - "step": 1, - "tooltip": "Used when Force FPS is enabled. Output video fps will be set exactly to this value." - }), - "aspect_mode": (["Auto", "Tall", "Wide", "Square"], { - "default": "Auto", - "tooltip": "Auto picks wide/tall/square from the source. Use Square/Tall/Wide to force the target bucket shape." - }), - }, - } - - def prepare(self, video, tier="Auto", crop_to_fit=True, force_fps=False, target_fps=16, aspect_mode="Auto"): - 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.") - - out_frame_rate = float(comp.frame_rate) if comp.frame_rate is not None else None - if force_fps: - frames, out_frame_rate, _ = _resample_video_frames_to_fps( - frames, comp.frame_rate, target_fps - ) - - # "Auto" in video prep uses the safer extend-friendly behavior. - # Keep accepting legacy "Safe Auto" values from older saved workflows. - internal_tier = "Safe Auto" if tier == "Auto" else tier - scaled_frames, _, _, _ = _wan22_scale_image_core( - frames, - tier=internal_tier, - crop_to_fit=crop_to_fit, - aspect_mode=aspect_mode, - ) - - scaled_video = VideoFromComponents( - VideoComponents( - images=scaled_frames, - audio=comp.audio, - frame_rate=out_frame_rate, - ) - ) - - fps = float(out_frame_rate) if out_frame_rate is not None else 0.0 - return (scaled_video, scaled_frames, fps) - - # ---------- Load Video MXD (video-only picker with refresh) ---------- - class LoadVideoMXD: - """Load a video from /input with a refresh button (videos only).""" - - CATEGORY = "image/video" - FUNCTION = "load" - RETURN_TYPES = ("VIDEO", "STRING") - RETURN_NAMES = ("video", "video_path") - TITLE = "Load Video MXD" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "file": ("COMBO", { - # Only allow video uploads in the picker - "video_upload": True, - # Custom route that returns ONLY videos in /input - "remote": { - "route": "/mxd/videos/input", - "refresh_button": True, - "control_after_refresh": "first", - }, - }), - } - } - - # --- helpers -------------------------------------------------------------- - - @staticmethod - def _resolve_video_path(file: str) -> str: - """ - Try to resolve `file` in a backwards-compatible way: - 1. If it's an annotated path, let folder_paths handle it. - 2. Otherwise treat it as relative to the input directory. - """ - # 1) Try annotated style (old workflows / uploads) - try: - return folder_paths.get_annotated_filepath(file) - except Exception: - pass - - # 2) Fall back to /input relative - base = folder_paths.get_input_directory() - candidate = os.path.join(base, file) - if os.path.isfile(candidate): - return candidate - - # If all else fails, just return what we got (will error later) - return candidate - - @staticmethod - def _is_video_file(path: str) -> bool: - _, ext = os.path.splitext(path) - return ext.lower() in VIDEO_EXTS - - # --- main function -------------------------------------------------------- - - def load(self, file: str): - video_path = self._resolve_video_path(file) - - if not os.path.isfile(video_path): - raise FileNotFoundError(f"[LoadVideoMXD] File not found: {video_path}") - - if not self._is_video_file(video_path): - raise ValueError(f"[LoadVideoMXD] Not a video file: {video_path}") - - print(f"[LoadVideoMXD] Loaded exactly: {video_path}") - return (VideoFromFile(video_path), video_path) - - # --- nice-to-haves -------------------------------------------------------- - - @classmethod - def IS_CHANGED(cls, file: str): - try: - p = cls._resolve_video_path(file) - return os.path.getmtime(p) - except Exception: - return 0 - - @classmethod - def VALIDATE_INPUTS(cls, file: str): - # First, try the annotated path (for backwards compat) - if folder_paths.exists_annotated_filepath(file): - resolved = folder_paths.get_annotated_filepath(file) - if not cls._is_video_file(resolved): - return f"This node only accepts video files ({', '.join(sorted(VIDEO_EXTS))})." - return True - - # Then, try treating it as /input-relative - base = folder_paths.get_input_directory() - candidate = os.path.join(base, file) - if os.path.isfile(candidate): - if not cls._is_video_file(candidate): - return f"This node only accepts video files ({', '.join(sorted(VIDEO_EXTS))})." - return True - - return f"Invalid video file: {file}" - - # ---------- Save Video MXD ---------- - class SaveVideoMXD(io.ComfyNode): - @classmethod - def define_schema(cls): - return io.Schema( - node_id="SaveVideoMXD", - display_name="Save Video MXD", - category="image/video", - description="Saves the input video to your ComfyUI output directory.", - inputs=[ - io.Video.Input("video", tooltip="The video to save."), - io.String.Input("filename_prefix", default="video/ComfyUI", tooltip="The prefix for the file to save. This may include formatting information such as %date:yyyy-MM-dd% or %Empty Latent Image.width% to include values from nodes."), - io.Combo.Input("format", options=VideoContainer.as_input(), default="auto", tooltip="The format to save the video as."), - io.Combo.Input("codec", options=VideoCodec.as_input(), default="auto", tooltip="The codec to use for the video."), - io.Boolean.Input( - "embed_workflow", - default=True, - label_on="embed", - label_off="skip", - tooltip="When high_workflow is connected, merge it into this video's embedded workflow " - "so dragging the final video into ComfyUI shows both the high-noise stage and " - "this stage together.", - ), - io.String.Input( - "high_workflow", - optional=True, - force_input=True, - tooltip="Connect a Load Latent node's 'high_workflow' output here to carry the " - "high-noise stage's workflow into this video's metadata.", - ), - ], - hidden=[io.Hidden.prompt, io.Hidden.extra_pnginfo], - is_output_node=True, - ) - - @classmethod - def execute(cls, video: VideoInput, filename_prefix: str, format: str, codec: str, - embed_workflow: bool = True, high_workflow: str = "") -> io.NodeOutput: - width, height = video.get_dimensions() - full_output_folder, filename, counter, subfolder, filename_prefix = folder_paths.get_save_image_path( - filename_prefix, - folder_paths.get_output_directory(), - width, - height - ) - - saved_metadata = None - if not args.disable_metadata: - metadata = {} - if cls.hidden.extra_pnginfo is not None: - metadata.update(cls.hidden.extra_pnginfo) - if cls.hidden.prompt is not None: - metadata["prompt"] = cls.hidden.prompt - if embed_workflow and high_workflow: - current_workflow = metadata.get("workflow") - merged_workflow = _merge_prior_workflow_into_current(high_workflow, current_workflow) - if merged_workflow is not current_workflow: - metadata["workflow"] = merged_workflow - if len(metadata) > 0: - saved_metadata = metadata - - file = f"{filename}_{counter:05}_.{VideoContainer.get_extension(format)}" - video.save_to( - os.path.join(full_output_folder, file), - format=VideoContainer(format), - codec=codec, - metadata=saved_metadata - ) - - return io.NodeOutput(ui=ui.PreviewVideo([ui.SavedResult(file, subfolder, io.FolderType.output)])) - - class PreviewVideoMXD(io.ComfyNode): - @classmethod - def define_schema(cls): - return io.Schema( - node_id="PreviewVideoMXD", - display_name="Preview Video MXD", - category="image/video", - 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 - def execute(cls, input_video: VideoInput): - # Save a temporary H264 file so ComfyUI has something to preview - out_dir = os.path.join(folder_paths.get_output_directory(), "previews") - os.makedirs(out_dir, exist_ok=True) - - preview_path = os.path.join(out_dir, "preview_temp.mp4") - input_video.save_to(preview_path, format="mp4", codec="h264") - - # ✅ Return the raw video object (not a tuple) - return io.NodeOutput( - input_video, - ui=ui.PreviewVideo([ - ui.SavedResult("preview_temp.mp4", "previews", io.FolderType.output) - ]) - ) - - -class GroupVideoFramesMXD: - CATEGORY = "MXD/Video" - TITLE = "Group Video Frames (MXD)" - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = ("IMAGE_GROUPS",) - OUTPUT_IS_LIST = (True,) - FUNCTION = "group_frames" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "frames": ("IMAGE",), - "group_size": ("INT", {"default": 81, "min": 1, "max": 5000, "step": 1}), - } - } - - def group_frames(self, frames, group_size): - import math, torch - - all_frames = list(frames) - total = len(all_frames) - num_groups = math.ceil(total / group_size) - grouped_tensors = [] - - for i in range(num_groups): - start = i * group_size - end = min(start + group_size, total) - group = all_frames[start:end] - - clean = [] - for f in group: - # ✅ drop redundant singleton batch dim if present - if f.ndim == 4 and f.shape[0] == 1: - f = f.squeeze(0) # (H,W,C) - # ✅ ensure shape (H,W,C) - if f.ndim != 3: - print(f"[GroupVideoFramesMXD] weird frame shape {f.shape}") - continue - clean.append(f) - - # ✅ stack back to (N,H,W,C) - if len(clean) == 0: - continue - stacked = torch.stack(clean, dim=0) - grouped_tensors.append(stacked) - - print(f"[GroupVideoFramesMXD] Split {total} frames into {len(grouped_tensors)} groups of up to {group_size}.") - return (grouped_tensors,) - -if HAVE_COMFY_API: - class Wan22FirstLastImageToVideoMXD(io.ComfyNode): - @classmethod - def define_schema(cls): - return io.Schema( - node_id="Wan22FirstLastImageToVideoMXD", - display_name="WAN 2.2 First & Last I2V MXD", - category="conditioning/video_models", - inputs=[ - io.Conditioning.Input("positive"), - io.Conditioning.Input("negative"), - io.Vae.Input("vae"), - io.Int.Input("length", default=81, min=1, max=nodes.MAX_RESOLUTION, 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), - ], - 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) -> io.NodeOutput: - spacial_scale = vae.spacial_compression_encode() - - # Assume incoming images are already pre-sized by upstream nodes. - height, width = start_image.shape[1], start_image.shape[2] if start_image is not None else (vae.latent_channels * spacial_scale, vae.latent_channels * spacial_scale) - - latent = torch.zeros( - [batch_size, vae.latent_channels, ((length - 1) // 4) + 1, height // spacial_scale, width // spacial_scale], - device=comfy.model_management.intermediate_device() - ) - - image = torch.ones((length, height, width, 3)) * 0.5 - mask = torch.ones((1, 1, latent.shape[2] * 4, latent.shape[-2], latent.shape[-1])) - - if start_image is not None: - image[:start_image.shape[0]] = start_image - mask[:, :, :start_image.shape[0] + 3] = 0.0 - - if end_image is not None: - image[-end_image.shape[0]:] = end_image - mask[:, :, -end_image.shape[0]:] = 0.0 - - concat_latent_image = vae.encode(image[:, :, :, :3]) - mask = mask.view(1, mask.shape[2] // 4, 4, mask.shape[3], mask.shape[4]).transpose(1, 2) - - positive = node_helpers.conditioning_set_values(positive, {"concat_latent_image": concat_latent_image, "concat_mask": mask}) - negative = node_helpers.conditioning_set_values(negative, {"concat_latent_image": concat_latent_image, "concat_mask": mask}) - - out_latent = {"samples": latent} - return io.NodeOutput(positive, negative, out_latent) - - -# ============================================================ -# LTX Video Image Scaler MXD -# ============================================================ -# Official LTX-2.3 rules (Lightricks model card + example workflows): -# - Width & height must be divisible by 32; frame count must be 8n+1. -# - The distilled two-stage workflow generates Stage 1 low-res, then the -# ltx-2.3-spatial-upscaler-x2 doubles it (exactly 2x) for Stage 2. -# - The one published two-stage resolution is Stage 1 960x544 -> 1920x1088. -# -# Tiers below are FINAL (Stage 2) sizes; Stage 1 is exactly half. Finals are -# kept /64 so Stage 1 stays /32 (the latent constraint). Only the 1080p 16:9 -# row is officially published by Lightricks; the portrait/square rows and the -# 720p/576p tiers are /32-aligned siblings at the same pixel budget. -# -# Buckets (FINAL size, all /64) -> Stage 1 (half, all /32): -# 1080p: 1920x1088 / 1088x1920 / 1408x1408 (Stage 1: 960x544 / 544x960 / 704x704) -# 720p: 1280x704 / 704x1280 / 960x960 (Stage 1: 640x352 / 352x640 / 480x480) -# 576p: 1024x576 / 576x1024 / 768x768 (Stage 1: 512x288 / 288x512 / 384x384) -# -# Fit (no pad): proportional resize <= target, /64 aligned. -# Crop (no pad): resize-to-cover then center-crop to exact bucket. -# Square images map to each tier's square bucket. -# ============================================================ - -_LTX_BUCKETS = { - "1080p": {"landscape": (1920, 1088), "portrait": (1088, 1920), "square": (1408, 1408)}, - "720p": {"landscape": (1280, 704), "portrait": (704, 1280), "square": (960, 960)}, - "576p": {"landscape": (1024, 576), "portrait": (576, 1024), "square": (768, 768)}, -} - - -def _ceil32(x): - x = (int(x) + 31) // 32 * 32 - return max(32, x) - - -def _floor32(x): - x = int(x) // 32 * 32 - return max(32, x) - - -def _floor64(x): - x = int(x) // 64 * 64 - return max(64, x) - - -def _ltx_stage1_dims(final_w, final_h): - """Return Stage 1 dimensions that upscale exactly to the final size.""" - return max(32, int(final_w) // 2), max(32, int(final_h) // 2) - - -def _ltx_resize_fit_inside(img, out_w, out_h): - """Resize to fit inside (out_w, out_h), output /64 aligned on both sides.""" - _, ih, iw, _ = img.shape - s = min(out_w / iw, out_h / ih) - tw = _floor64(iw * s) - th = _floor64(ih * s) - tw = max(32, min(tw, nodes.MAX_RESOLUTION)) - th = max(32, min(th, nodes.MAX_RESOLUTION)) - resized = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) - return resized, tw, th - - -def _ltx_resize_then_center_crop(img, out_w, out_h): - """Resize to cover (out_w, out_h) then center-crop to exact /32 target.""" - _, ih, iw, _ = img.shape - s = max(out_w / iw, out_h / ih) - tw = _ceil32(iw * s) - th = _ceil32(ih * s) - tmp = comfy.utils.common_upscale(img.movedim(-1, 1), tw, th, "bilinear", "center").movedim(1, -1) - y0 = max(0, (th - out_h) // 2) - x0 = max(0, (tw - out_w) // 2) - return tmp[:, y0:y0+out_h, x0:x0+out_w, :] - - -def _ltx_pick_bucket(iw, ih, tier): - """Pick the landscape / portrait / square bucket for the given tier.""" - tier_map = _LTX_BUCKETS[tier] - if _is_squareish(iw, ih): - return tier_map["square"] - return tier_map["landscape"] if iw >= ih else tier_map["portrait"] - - -def _ltx_scale_image_core(image, tier="1080p", crop_to_fit=True): - """ - Core LTX scaler. Returns (scaled_image, final_w, final_h, stage1_w, stage1_h). - 'tier' is the FINAL (Stage 2) size budget; Stage 1 is exactly half. - """ - _, ih, iw, _ = image.shape - - bw, bh = _ltx_pick_bucket(iw, ih, tier) - - if _is_squareish(iw, ih): - crop_to_fit = False - - if crop_to_fit: - out = _ltx_resize_then_center_crop(image, bw, bh) - else: - out, bw, bh = _ltx_resize_fit_inside(image, bw, bh) - - final_w = int(out.shape[2]) - final_h = int(out.shape[1]) - stage1_w, stage1_h = _ltx_stage1_dims(final_w, final_h) - return out, final_w, final_h, stage1_w, stage1_h - - -class LTX_Image_Scaler_MXD: - """ - MXD Image Scaler for LTX Video (distilled two-stage workflow). - - 'tier' is the FINAL (Stage 2) size; Stage 1 is exactly half. Finals are /64 - so Stage 1 stays /32 (the LTX latent constraint). Wire stage1_width / - stage1_height into the empty latent for the low-res pass; the spatial - upscaler-x2 then doubles it back to the final size. - - Tiers (final / Stage 1): - 1080p 1920x1088 (official 16:9) / 1088x1920 / 1408x1408 -> half - 720p 1280x704 / 704x1280 / 960x960 -> half - 576p 1024x576 / 576x1024 / 768x768 -> half - - Modes: - Perfect Fit (Crops Edges) resize-to-cover + center-crop to exact bucket. - Closest Fit (No Crop) proportional resize, /64-aligned; may be smaller. - - Square images (within +-3% of 1:1) map to the tier's square bucket. - Outputs the scaled image at final size plus the Stage 1 dimensions. - """ - - TITLE = "LTX Video Image Scaler MXD" - CATEGORY = "image/processing" - RETURN_TYPES = ("IMAGE", "INT", "INT") - RETURN_NAMES = ("image", "width", "height") - FUNCTION = "scale" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "tier": (["1080p", "720p", "576p"], {"default": "1080p"}), - "crop_to_fit": ("BOOLEAN", { - "default": True, - "label_on": "Crop Edges", - "label_off": "Closest Fit (No Crop)", - }), - } - } - - def scale(self, image, tier="1080p", crop_to_fit=True): - image = _validate_image_batch_4d(image, "LTX_Image_Scaler_MXD", "image") - out, _final_w, _final_h, stage1_w, stage1_h = _ltx_scale_image_core( - image, tier=tier, crop_to_fit=crop_to_fit - ) - return (out, stage1_w, stage1_h) - - -class PadImageForOutpaintingMXD: - SEARCH_ALIASES = ["extend canvas", "expand image", "outpaint pad"] - - RETURN_TYPES = ("IMAGE", "MASK") - FUNCTION = "expand_image" - CATEGORY = "image/transform" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "image": ("IMAGE",), - "left": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), - "top": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), - "right": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), - "bottom": ("INT", {"default": 0, "min": 0, "max": nodes.MAX_RESOLUTION, "step": 2}), - "round_to": (["None", "2", "8", "16", "32", "64"], {"default": "16"}), - } - } - - @staticmethod - def _nearest_multiple(value: int, multiple: int, padded: bool) -> int: - if multiple <= 1 or value % multiple == 0: - return value - lower = (value // multiple) * multiple - upper = lower + multiple - if lower <= 0: - return upper - if not padded: - return lower - return lower if value - lower <= upper - value else upper - - @staticmethod - def _axis_plan(size: int, before: int, after: int, multiple: int) -> Tuple[int, int, int, int, int]: - target = size + before + after - if multiple > 1: - target = PadImageForOutpaintingMXD._nearest_multiple(target, multiple, before + after > 0) - - delta = target - (size + before + after) - if delta < 0: - remove = -delta - from_after = min(after, remove) - after -= from_after - remove -= from_after - from_before = min(before, remove) - before -= from_before - remove -= from_before - crop_before = remove // 2 - crop_after = remove - crop_before - else: - crop_before = 0 - crop_after = 0 - if before > 0 and after > 0: - add_before = delta // 2 - before += add_before - after += delta - add_before - elif before > 0: - before += delta - else: - after += delta - - final_size = size - crop_before - crop_after + before + after - if final_size <= 0: - raise ValueError("[PadImageForOutpaintingMXD] Rounding removed the full image on one axis.") - return before, after, crop_before, crop_after, final_size - - def expand_image(self, image, left, top, right, bottom, round_to="16"): - image = _validate_image_batch_4d(image, "PadImageForOutpaintingMXD", "image") - batch, height, width, channels = image.size() - multiple = 1 if round_to == "None" else int(round_to) - - left, right, crop_left, crop_right, final_width = self._axis_plan(width, left, right, multiple) - top, bottom, crop_top, crop_bottom, final_height = self._axis_plan(height, top, bottom, multiple) - - cropped = image[:, crop_top:height - crop_bottom, crop_left:width - crop_right, :] - crop_height = cropped.shape[1] - crop_width = cropped.shape[2] - - new_image = torch.full( - (batch, final_height, final_width, channels), - 0.5, - dtype=image.dtype, - device=image.device, - ) - new_image[:, top:top + crop_height, left:left + crop_width, :] = cropped - - mask = torch.ones( - (final_height, final_width), - dtype=torch.float32, - device=image.device, - ) - mask[top:top + crop_height, left:left + crop_width] = 0.0 - - return (new_image, mask.unsqueeze(0)) - - -# ---------- Node registration ---------- -NODE_CLASS_MAPPINGS = { - "Wan2_2EmptyLatentImageMXD": Wan2_2EmptyLatentImageMXD, - "wan22EmptyHunyuanLatentVideoMXD": wan22EmptyHunyuanLatentVideoMXD, - "SaveLatent_I2V_MXD": SaveLatent_I2V_MXD, - "LoadLatent_I2V_MXD": LoadLatent_I2V_MXD, - "LoadLatents_FromFolder_I2V_MXD": LoadLatents_FromFolder_I2V_MXD, - "LoadLatent_I2V_Pipe_MXD": LoadLatent_I2V_Pipe_MXD, - "LoadLatents_FromFolder_I2V_Pipe_MXD": LoadLatents_FromFolder_I2V_Pipe_MXD, - "LatentPipeUnpack_MXD": LatentPipeUnpack_MXD, - "WAN22_I2V_Image_Scaler_MXD": WAN22_I2V_Image_Scaler_MXD, - "LTX_Image_Scaler_MXD": LTX_Image_Scaler_MXD, - "WAN22_I2V_Match_Resolution_MXD": WAN22_I2V_Match_Resolution_MXD, - "Frames_Remove_From_Start_MXD": Frames_Remove_From_Start_MXD, - "GroupVideoFramesMXD": GroupVideoFramesMXD, - "Frames_Select_StartEnd_MXD": Frames_Select_StartEnd_MXD, - "PadImageForOutpaintingMXD": PadImageForOutpaintingMXD, -} - -if HAVE_COMFY_API: - NODE_CLASS_MAPPINGS.update({ - "Wan22ImageToVideoMXD": Wan22ImageToVideoMXD, - "WAN22_I2V_Video_Prep_MXD": WAN22_I2V_Video_Prep_MXD, - "CombineVideos_MXD": CombineVideos_MXD, - "LoadVideoMXD": LoadVideoMXD, - "SaveVideoMXD": SaveVideoMXD, - "PreviewVideoMXD": PreviewVideoMXD, - "Wan22FirstLastImageToVideoMXD": Wan22FirstLastImageToVideoMXD, - }) - -NODE_DISPLAY_NAME_MAPPINGS = { - "Wan2_2EmptyLatentImageMXD": "Wan 2.2 Empty Latent Image MXD", - "wan22EmptyHunyuanLatentVideoMXD": "WAN2.2 Empty Latent Video MXD", - "SaveLatent_I2V_MXD": "Save Latent MXD", - "LoadLatent_I2V_MXD": "Load Latent MXD", - "LoadLatents_FromFolder_I2V_MXD": "Load Latent Batch MXD", - "LoadLatent_I2V_Pipe_MXD": "Load Latent Pipe MXD", - "LoadLatents_FromFolder_I2V_Pipe_MXD": "Load Latent Batch Pipe MXD", - "LatentPipeUnpack_MXD": "Unpack Latent Pipe MXD", - "WAN22_I2V_Image_Scaler_MXD": "Image Scaler Wan 2.2 I2V MXD", - "LTX_Image_Scaler_MXD": "LTX Video Image Scaler MXD", - "WAN22_I2V_Match_Resolution_MXD": "Match Resolution Wan 2.2 I2V MXD", - "Frames_Remove_From_Start_MXD": "Remove Frames From Start MXD", - "GroupVideoFramesMXD": "Group Video Frames MXD", - "Frames_Select_StartEnd_MXD": "Select Frames MXD", - "PadImageForOutpaintingMXD": "Pad Image for Outpainting MXD", -} - -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", - "CombineVideos_MXD": "Combine Videos MXD", - "LoadVideoMXD": "Load Video MXD", - "SaveVideoMXD": "Save Video MXD", - "PreviewVideoMXD": "Preview Video MXD", - "Wan22FirstLastImageToVideoMXD": "Wan 2.2 I2V First & Last Frame MXD", - }) - -def _add_mxd_aliases(class_map, display_map): - alias_sources = {} - for key in list(class_map.keys()): - if "MXD" in key.upper(): - continue - alias = f"{key} MXD" - if alias in class_map: - continue - class_map[alias] = class_map[key] - alias_sources[alias] = key - for alias, source in alias_sources.items(): - if alias not in display_map: - display_map[alias] = display_map.get(source, alias) - return alias_sources - -_add_mxd_aliases(NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS)