diff --git a/.gitignore b/.gitignore index 387f1e9..3a9489c 100644 --- a/.gitignore +++ b/.gitignore @@ -32,6 +32,7 @@ logs/ # Local testing workspace testing/ userdata/ +local_notes/ # ComfyUI local cache & configs ComfyUI/output/ diff --git a/__init__.py b/__init__.py index 1fab914..db1a73e 100644 --- a/__init__.py +++ b/__init__.py @@ -24,8 +24,8 @@ for _name in ( "mediacomparers", "wan22nodes", "loraloader_mxd", - "wan_svi_first_last_mxd", "CharacterPrompts", + "ltxnodes", ): _mod = _safe_import(_name) _class_map, _display_map = _get_mappings(_mod) diff --git a/loraloader_mxd/server/routes_model_info.py b/loraloader_mxd/server/routes_model_info.py index 572b38c..5d14c14 100644 --- a/loraloader_mxd/server/routes_model_info.py +++ b/loraloader_mxd/server/routes_model_info.py @@ -20,6 +20,14 @@ def _check_valid_model_type(request): return None +def _file_details_sort_key(file_info): + modified = file_info.get('modified') + if not isinstance(modified, (int, float)): + modified = 0 + file = str(file_info.get('file') or '').replace('\\', '/').lower() + return (-modified, file) + + @routes.get('/loraloader-mxd/api/{type}') async def api_get_models_list(request): """Returns a list of model types from user configuration. @@ -59,6 +67,8 @@ async def api_get_models_list(request): id=f'no_file_details_{model_type}', at_most_secs=30 ) + if model_type == 'loras': + response.sort(key=_file_details_sort_key) return web.json_response(response) return web.json_response(list(files)) diff --git a/ltxnodes.py b/ltxnodes.py new file mode 100644 index 0000000..0a32e38 --- /dev/null +++ b/ltxnodes.py @@ -0,0 +1,610 @@ +from __future__ import annotations +import os +import re +import struct +import time +import urllib.error +import urllib.request +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 + +_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( + 'VHS_latentpreview', + {'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) + for preview in previews: + img = Image.fromarray(preview.numpy()) + buf = BytesIO() + buf.write((1).to_bytes(length=4, byteorder='big') * 2) + buf.write(ind.to_bytes(length=4, byteorder='big')) + buf.write(struct.pack('16p', _serv.last_node_id.encode('ascii'))) + img.save(buf, format="JPEG", quality=95, compress_level=1) + _serv.send_sync(server.BinaryEventTypes.PREVIEW_IMAGE, buf.getvalue(), _serv.client_id) + # 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) + + +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) + + # 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 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,) + + +######################################################################################################################## +NODE_CLASS_MAPPINGS = { + "LTXVideoEmptyLatent_MXD": LTXVideoEmptyLatentMXD, + "LTXKSampler_MXD": LTXKSamplerMXD, + "LTXKSampler2_MXD": LTXKSampler2MXD, +} + +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", +} diff --git a/maxedoutnodes.py b/maxedoutnodes.py index b8b612d..8c525df 100644 --- a/maxedoutnodes.py +++ b/maxedoutnodes.py @@ -1,9 +1,10 @@ from __future__ import annotations -import torch, math, comfy, os, folder_paths, node_helpers, comfy.model_management, comfy.utils, json, hashlib, re +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 +from PIL import Image, ImageOps, ImageSequence, ImageFilter, ImageColor +from nodes import SaveImage try: from comfy_api.latest import io HAVE_COMFY_API = True @@ -1514,6 +1515,189 @@ class SmartCropByMaskMXD: ######################################################################################################################## +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, @@ -1535,6 +1719,8 @@ NODE_CLASS_MAPPINGS = { "Save Image MXD": SaveImage_MXD, "Extract Workflow From Image MXD": ExtractWorkflowFromImageMXD, "SmartCropByMaskMXD": SmartCropByMaskMXD, + "BboxDetectorCombinedBatchMXD": BboxDetectorCombinedBatchMXD, + "ImageAndMaskPreviewMXD": ImageAndMaskPreviewMXD, } if HAVE_COMFY_API: @@ -1563,6 +1749,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "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: diff --git a/wan22nodes.py b/wan22nodes.py index a9bef46..35455f2 100644 --- a/wan22nodes.py +++ b/wan22nodes.py @@ -196,41 +196,51 @@ class SaveLatent_I2V_MXD: def save_only(self, samples, positive, negative, filename_prefix="I2V", prompt=None, extra_pnginfo=None, unique_id=None): - - # ---- save latent (.latent) ---- - 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 + _save_i2v_latent_bundle( + samples=samples, + positive=positive, + negative=negative, + filename_prefix=filename_prefix, + prompt=prompt, + extra_pnginfo=extra_pnginfo, + unique_id=unique_id, ) + return {} - # Metadata - meta = None - if not args.disable_metadata: - meta = {} - if prompt is not None: - try: meta["prompt"] = json.dumps(prompt) - except: pass - if extra_pnginfo is not None: - for k, v in extra_pnginfo.items(): - try: meta[k] = json.dumps(v) - except: pass - _attach_source_ksampler_metadata(meta, prompt, unique_id) +class SaveLatent_VACE22_MXD(SaveLatent_I2V_MXD): + """ + VACE 2.2 saver: I2V latent + conditioning sidecar + trim_latent value. + Kept as a separate node so existing I2V workflows stay unchanged. + """ + TITLE = "Save Latent Vace 2.2" + CATEGORY = "MXD/Latents (VACE 2.2)" - 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([]), + @classmethod + def INPUT_TYPES(cls): + inputs = SaveLatent_I2V_MXD.INPUT_TYPES() + inputs["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." + }), } - comfy.utils.save_torch_file(payload, latent_path, metadata=meta) + return inputs - # ---- save conditioning sidecar (.cond.pt) ---- - cond_path = latent_path.replace(".latent", ".cond.pt") - torch.save({"positive": positive, "negative": negative}, cond_path) - - # No preview logic at all + def save_only(self, samples, positive, negative, filename_prefix="I2V", + 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 ---------- @@ -453,6 +463,97 @@ def _attach_source_ksampler_metadata(meta: Dict[str, Any], prompt: Any, unique_i 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 = "" @@ -1267,6 +1368,129 @@ class LoadLatents_FromFolder_I2V_MXD(LoadLatents_FromFolder_WithParams): filename_prefixes, ) +class LoadLatent_VACE22_MXD(LoadLatent_I2V_MXD): + """ + I2V loader plus the VACE 2.2 trim_latent value saved by Save Latent Vace 2.2. + """ + TITLE = "Load Latent Vace 2.2" + CATEGORY = "MXD/Latents (VACE 2.2)" + + RETURN_TYPES = ( + "FLOAT", + "CONDITIONING", + "CONDITIONING", + "LATENT", + "INT", + "FLOAT", + "STRING", + "STRING", + "INT", + "STRING", + "INT", + ) + RETURN_NAMES = ( + "shift", + "positive", + "negative", + "samples", + "steps", + "cfg", + "sampler_name", + "scheduler", + "end_at_step", + "filename_prefix", + "trim_latent", + ) + + @classmethod + def INPUT_TYPES(s): + inputs = LoadLatent_I2V_MXD.INPUT_TYPES.__func__(s) + sampler_type = s.RETURN_TYPES[6] + scheduler_type = s.RETURN_TYPES[7] + s.RETURN_TYPES = ( + "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", + "INT", "FLOAT", sampler_type, scheduler_type, + "INT", "STRING", "INT", + ) + return inputs + + def load(self, latent): + base_tuple = super().load(latent) + latent_ref = latent if str(latent).startswith("latents/") else f"latents/{latent}" + latent_path = folder_paths.get_annotated_filepath(latent_ref) + _pos, _neg, sidecar = _load_i2v_conditioning_sidecar(latent_path) + _sample_dict, meta, _keys = _load_latent_file(latent_path) + trim_latent = _coerce_trim_latent(sidecar.get("trim_latent", meta.get("trim_latent", 0))) + return (*base_tuple, trim_latent) + + +class LoadLatents_FromFolder_VACE22_MXD(LoadLatents_FromFolder_I2V_MXD): + """ + Batch I2V loader plus a trim_latent list aligned with each returned latent slice. + """ + TITLE = "Load Latents (Folder, Vace 2.2)" + CATEGORY = "MXD/Latents (VACE 2.2)" + FUNCTION = "load_batch_vace22" + + RETURN_TYPES = ( + "FLOAT", + "CONDITIONING", + "CONDITIONING", + "LATENT", + "INT", + "FLOAT", + "STRING", + "STRING", + "INT", + "STRING", + "INT", + ) + 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): + inputs = LoadLatents_FromFolder_I2V_MXD.INPUT_TYPES.__func__(s) + sampler_type = s.RETURN_TYPES[6] + scheduler_type = s.RETURN_TYPES[7] + s.RETURN_TYPES = ( + "FLOAT", "CONDITIONING", "CONDITIONING", "LATENT", + "INT", "FLOAT", sampler_type, scheduler_type, + "INT", "STRING", "INT", + ) + return inputs + + def load_batch_vace22(self, subfolder): + base_tuple = super().load_batch_i2v(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) + + trims = [] + for path in files: + sample_dict, meta, _keys = _load_latent_file(path) + _pos, _neg, sidecar = _load_i2v_conditioning_sidecar(path) + trim_latent = _coerce_trim_latent(sidecar.get("trim_latent", meta.get("trim_latent", 0))) + t = sample_dict["samples"] + slice_count = int(t.size(0)) if isinstance(t, torch.Tensor) and t.dim() >= 4 and t.size(0) > 1 else 1 + trims.extend([trim_latent] * slice_count) + + return (*base_tuple, trims) + # ---------- Empty latent image generator (for video nodes) ---------- class Wan2_2EmptyLatentImageMXD: """ @@ -1348,6 +1572,7 @@ class wan22EmptyHunyuanLatentVideoMXD: 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), @@ -1462,9 +1687,10 @@ if HAVE_COMFY_API: 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)] # 16:9, 9:16 -SQUARE_TOL = 0.03 # ±3% aspect-ratio tolerance counts as "square-ish" +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) @@ -1486,6 +1712,32 @@ 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. @@ -1578,22 +1830,26 @@ def _resize_to_explicit_resolution(img, out_w, out_h, match_mode="crop_to_match" _WAN22_VALID_RES = { (832, 480), (480, 832), (1280, 720), (720, 1280), - (624, 624), (720, 720), + (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): +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: - if tier == "720p": - return (720, 720) - return (624, 624) + 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": @@ -1615,7 +1871,7 @@ def _wan22_pick_bucket(iw, ih, tier, 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): +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). @@ -1639,7 +1895,7 @@ def _wan22_scale_image_core(image, tier="Auto", crop_to_fit=False): "WAN 2.2 works best around:\n" " - 480p tier ~= 832x480 (or 480x832)\n" " - 720p tier ~= 1280x720 (or 720x1280)\n" - " - Squares: 624x624 or 720x720\n\n" + " - 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." ) @@ -1647,12 +1903,7 @@ def _wan22_scale_image_core(image, tier="Auto", crop_to_fit=False): tier = "Auto" # --- Normal path (Auto / 480p / 720p) --- - bw, bh = _wan22_pick_bucket(iw, ih, tier, crop_to_fit) - is_squareish = _is_squareish(iw, ih) - - if is_squareish: - crop_to_fit = False - + 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) @@ -1733,7 +1984,7 @@ class WAN22_I2V_Image_Scaler_MXD: - Crop (no pad): resize-to-cover then center-crop to exact target. - Square handling: * Auto & 480p: ~square → 624×624 - * 720p: ~square → 720×720 + * 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. @@ -1757,6 +2008,10 @@ class WAN22_I2V_Image_Scaler_MXD: "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." + }), } } @@ -1770,7 +2025,7 @@ class WAN22_I2V_Image_Scaler_MXD: # --- Square handling --- if is_squareish: if tier == "720p": - return (720, 720) + return (1024, 1024) else: return (624, 624) @@ -1796,10 +2051,15 @@ class WAN22_I2V_Image_Scaler_MXD: # ----------------------------- # Main function # ----------------------------- - def scale(self, image, tier="Auto", crop_to_fit=False): + 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) + 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 @@ -1821,7 +2081,7 @@ class WAN22_I2V_Image_Scaler_MXD: "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 720×720\n\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." ) @@ -1891,7 +2151,7 @@ class WAN22_I2V_Match_Resolution_MXD: "Valid WAN 2.2 buckets are:\n" " - 832x480 / 480x832\n" " - 1280x720 / 720x1280\n" - " - 624x624 / 720x720\n\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" @@ -2127,13 +2387,13 @@ if HAVE_COMFY_API: """ Prepare a source video for iterative WAN 2.2 extension: - scale entire video using WAN bucket logic - - output start/end frames from the full scaled video + - output the scaled frame batch directly - keep default workflow simple for common use """ CATEGORY = "MXD/video" FUNCTION = "prepare" - RETURN_TYPES = ("VIDEO", "IMAGE", "IMAGE", "INT", "INT", "FLOAT") - RETURN_NAMES = ("scaled_video", "start_image", "end_image", "width", "height", "fps") + RETURN_TYPES = ("VIDEO", "IMAGE", "FLOAT") + RETURN_NAMES = ("scaled_video", "images", "fps") @classmethod def INPUT_TYPES(cls): @@ -2146,21 +2406,27 @@ if HAVE_COMFY_API: "label_on": "Perfect Fit (Crops Edges)", "label_off": "Closest Fit (No Crop)" }), - "fps_mode": (["none", "force"], { - "default": "none", - "tooltip": "none = keep source fps. force = resample frames (drop/duplicate) and set exact target fps." + "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": ("FLOAT", { - "default": 16.0, - "min": 0.001, - "max": 1000.0, - "step": 0.01, - "tooltip": "Used when fps_mode=force. Output video fps will be set exactly to this value." + "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, fps_mode="none", target_fps=16.0): + 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: @@ -2179,7 +2445,7 @@ if HAVE_COMFY_API: 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 fps_mode == "force": + if force_fps: frames, out_frame_rate, _ = _resample_video_frames_to_fps( frames, comp.frame_rate, target_fps ) @@ -2187,13 +2453,13 @@ if HAVE_COMFY_API: # "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, out_w, out_h, _ = _wan22_scale_image_core( - frames, tier=internal_tier, crop_to_fit=crop_to_fit + scaled_frames, _, _, _ = _wan22_scale_image_core( + frames, + tier=internal_tier, + crop_to_fit=crop_to_fit, + aspect_mode=aspect_mode, ) - start_image = scaled_frames[0:1].clone() - end_image = scaled_frames[-1:].clone() - scaled_video = VideoFromComponents( VideoComponents( images=scaled_frames, @@ -2203,92 +2469,7 @@ if HAVE_COMFY_API: ) fps = float(out_frame_rate) if out_frame_rate is not None else 0.0 - return (scaled_video, start_image, end_image, out_w, out_h, fps) - - class WAN22_I2V_Video_Prep_Advanced_MXD: - """ - Advanced variant of WAN22_I2V_Video_Prep_MXD with frame-selection controls. - """ - CATEGORY = "MXD/video" - FUNCTION = "prepare" - RETURN_TYPES = ("VIDEO", "IMAGE", "IMAGE", "IMAGE", "INT", "INT", "FLOAT") - RETURN_NAMES = ("scaled_video", "selected_frames", "start_image", "end_image", "width", "height", "fps") - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "video": ("VIDEO",), - "tier": (["Auto", "480p", "720p"], {"default": "Auto"}), - "crop_to_fit": ("BOOLEAN", { - "default": True, - "label_on": "Perfect Fit (Crops Edges)", - "label_off": "Closest Fit (No Crop)" - }), - "fps_mode": (["none", "force"], { - "default": "none", - "tooltip": "none = keep source fps. force = resample frames (drop/duplicate) and set exact target fps." - }), - "target_fps": ("FLOAT", { - "default": 16.0, - "min": 0.001, - "max": 1000.0, - "step": 0.01, - "tooltip": "Used when fps_mode=force. Output video fps will be set exactly to this value." - }), - "mode": (["start", "end"], {"default": "end"}), - "count": ("INT", {"default": 1, "min": 1, "max": 10000}), - "offset": ("INT", {"default": 1, "min": 1, "max": 10000}), - }, - } - - def prepare(self, video, tier="Auto", crop_to_fit=True, fps_mode="none", target_fps=16.0, mode="end", count=1, offset=1): - comp = video.get_components() - if isinstance(comp.images, list): - if len(comp.images) == 0: - raise ValueError("[WAN22_I2V_Video_Prep_Advanced_MXD] Input video has zero frames.") - frames = torch.stack(comp.images) - else: - frames = comp.images - - if frames is None: - raise ValueError("[WAN22_I2V_Video_Prep_Advanced_MXD] Input video has no frames.") - if frames.ndim == 3: - frames = frames.unsqueeze(0) - if frames.ndim != 4: - raise ValueError(f"[WAN22_I2V_Video_Prep_Advanced_MXD] Unexpected frame tensor shape: {tuple(frames.shape)}") - if frames.shape[0] <= 0: - raise ValueError("[WAN22_I2V_Video_Prep_Advanced_MXD] Input video has zero frames.") - - out_frame_rate = float(comp.frame_rate) if comp.frame_rate is not None else None - if fps_mode == "force": - 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, out_w, out_h, _ = _wan22_scale_image_core( - frames, tier=internal_tier, crop_to_fit=crop_to_fit - ) - - selected_frames = _select_frames_start_end( - scaled_frames, count=count, offset=offset, mode=mode - ) - start_image = selected_frames[0:1].clone() - end_image = selected_frames[-1:].clone() - - scaled_video = VideoFromComponents( - VideoComponents( - images=scaled_frames, - audio=comp.audio, - frame_rate=out_frame_rate, - ) - ) - - fps = float(out_frame_rate) if out_frame_rate is not None else 0.0 - return (scaled_video, selected_frames, start_image, end_image, out_w, out_h, fps) + return (scaled_video, scaled_frames, fps) # ---------- Load Video MXD (video-only picker with refresh) ---------- class LoadVideoMXD: @@ -2607,31 +2788,33 @@ if HAVE_COMFY_API: # ============================================================ # LTX Video Image Scaler MXD # ============================================================ -# LTX Video requires all dimensions to be multiples of 32. -# Tiers: 480p / 768 / 1024 (or Auto to pick nearest by area) -# Fit (no pad): proportional resize <= target, /32 aligned. +# 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. -# Buckets (all /32): -# 480p: 832x480 / 480x832 / 512x512 -# 768: 1280x768 / 768x1280 / 768x768 -# 1024: 1792x1024 / 1024x1792 / 1024x1024 # ============================================================ _LTX_BUCKETS = { - "480p": {"landscape": (832, 480), "portrait": (480, 832), "square": (512, 512)}, - "768": {"landscape": (1280, 768), "portrait": (768, 1280), "square": (768, 768)}, - "1024": {"landscape": (1792, 1024), "portrait": (1024, 1792), "square": (1024, 1024)}, + "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)}, } -_LTX_TIER_AREAS = { - "480p": 832 * 480, # 399,360 - "768": 1280 * 768, # 983,040 - "1024": 1792 * 1024, # 1,835,008 -} - -_LTX_VALID_RES = {b for t in _LTX_BUCKETS.values() for b in t.values()} - def _ceil32(x): x = (int(x) + 31) // 32 * 32 @@ -2643,16 +2826,22 @@ def _floor32(x): return max(32, x) -def _ltx_is_valid_res(w, h): - return (w, h) in _LTX_VALID_RES +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 /32 aligned on both sides.""" + """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 = _floor32(iw * s) - th = _floor32(ih * s) + 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) @@ -2671,12 +2860,6 @@ def _ltx_resize_then_center_crop(img, out_w, out_h): return tmp[:, y0:y0+out_h, x0:x0+out_w, :] -def _ltx_pick_tier_auto(iw, ih): - """Pick the LTX tier whose reference area is closest to the input area.""" - area = iw * ih - return min(_LTX_TIER_AREAS, key=lambda t: abs(area - _LTX_TIER_AREAS[t])) - - def _ltx_pick_bucket(iw, ih, tier): """Pick the landscape / portrait / square bucket for the given tier.""" tier_map = _LTX_BUCKETS[tier] @@ -2685,35 +2868,13 @@ def _ltx_pick_bucket(iw, ih, tier): return tier_map["landscape"] if iw >= ih else tier_map["portrait"] -def _ltx_scale_image_core(image, tier="Auto", crop_to_fit=True): +def _ltx_scale_image_core(image, tier="1080p", crop_to_fit=True): """ - Core LTX scaler. Returns (scaled_image, out_w, out_h, passthrough). - passthrough=True only when Safe Auto detects an already-valid resolution. + 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 - if tier == "Safe Auto": - if _ltx_is_valid_res(iw, ih): - return image, iw, ih, True - area = iw * ih - min_area = int(_LTX_TIER_AREAS["480p"] * 0.5) - max_area = int(_LTX_TIER_AREAS["1024"] * 1.8) - if area < min_area or area > max_area: - size_label = "small" if area < min_area else "large" - raise ValueError( - f"[LTX_Image_Scaler_MXD] Input {iw}x{ih} is too {size_label} for LTX Video buckets.\n" - "LTX Video works best around:\n" - " - 480p tier: 832x480 / 480x832 / 512x512\n" - " - 768 tier: 1280x768 / 768x1280 / 768x768\n" - " - 1024 tier: 1792x1024 / 1024x1792 / 1024x1024\n\n" - "Use a source image closer to one of these tiers, or process it " - "through your LTX workflow first." - ) - tier = "Auto" - - if tier == "Auto": - tier = _ltx_pick_tier_auto(iw, ih) - bw, bh = _ltx_pick_bucket(iw, ih, tier) if _is_squareish(iw, ih): @@ -2724,25 +2885,32 @@ def _ltx_scale_image_core(image, tier="Auto", crop_to_fit=True): else: out, bw, bh = _ltx_resize_fit_inside(image, bw, bh) - return out, int(out.shape[2]), int(out.shape[1]), False + 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 — all outputs are multiples of 32. + MXD Image Scaler for LTX Video (distilled two-stage workflow). - Tiers: - Auto — picks the tier whose area is closest to the input. - 480p — targets 832x480 / 480x832 / 512x512. - 768 — targets 1280x768 / 768x1280 / 768x768. - 1024 — targets 1792x1024 / 1024x1792 / 1024x1024. + '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 size. - Closest Fit (No Crop) proportional resize, /32-aligned; may be smaller than bucket. + 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 (aspect ratio within +-3% of 1:1) map to the tier's square bucket. - Returns image + width + height so downstream nodes can read the final dims directly. + 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" @@ -2756,19 +2924,117 @@ class LTX_Image_Scaler_MXD: return { "required": { "image": ("IMAGE",), - "tier": (["Auto", "480p", "768", "1024"], {"default": "Auto"}), + "tier": (["1080p", "720p", "576p"], {"default": "1080p"}), "crop_to_fit": ("BOOLEAN", { "default": True, - "label_on": "Perfect Fit (Crops Edges)", + "label_on": "Crop Edges", "label_off": "Closest Fit (No Crop)", }), } } - def scale(self, image, tier="Auto", crop_to_fit=True): + def scale(self, image, tier="1080p", crop_to_fit=True): image = _validate_image_batch_4d(image, "LTX_Image_Scaler_MXD", "image") - out, ow, oh, _ = _ltx_scale_image_core(image, tier=tier, crop_to_fit=crop_to_fit) - return (out, ow, oh) + 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 ---------- @@ -2781,19 +3047,22 @@ NODE_CLASS_MAPPINGS = { "SaveLatent_I2V_MXD": SaveLatent_I2V_MXD, "LoadLatent_I2V_MXD": LoadLatent_I2V_MXD, "LoadLatents_FromFolder_I2V_MXD": LoadLatents_FromFolder_I2V_MXD, + "SaveLatent_VACE22_MXD": SaveLatent_VACE22_MXD, + "LoadLatent_VACE22_MXD": LoadLatent_VACE22_MXD, + "LoadLatents_FromFolder_VACE22_MXD": LoadLatents_FromFolder_VACE22_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, - "WAN22_I2V_Video_Prep_Advanced_MXD": WAN22_I2V_Video_Prep_Advanced_MXD, "CombineVideos_MXD": CombineVideos_MXD, "LoadVideoMXD": LoadVideoMXD, "SaveVideoMXD": SaveVideoMXD, @@ -2810,19 +3079,22 @@ NODE_DISPLAY_NAME_MAPPINGS = { "SaveLatent_I2V_MXD": "Save Latent I2V MXD", "LoadLatent_I2V_MXD": "Load Latent I2V MXD", "LoadLatents_FromFolder_I2V_MXD": "Load Latent Batch I2V MXD", + "SaveLatent_VACE22_MXD": "Save Latent Vace 2.2 MXD", + "LoadLatent_VACE22_MXD": "Load Latent Vace 2.2 MXD", + "LoadLatents_FromFolder_VACE22_MXD": "Load Latent Batch Vace 2.2 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", - "WAN22_I2V_Video_Prep_Advanced_MXD": "WAN 2.2 Video Prep I2V MXD Advanced", "CombineVideos_MXD": "Combine Videos MXD", "LoadVideoMXD": "Load Video MXD", "SaveVideoMXD": "Save Video MXD", diff --git a/web/index.js b/web/index.js index 82f14a2..e6a7f34 100644 --- a/web/index.js +++ b/web/index.js @@ -3,4 +3,6 @@ import './js/image_comparer.js'; import './addons/zip_loader/js/zip_loader.js'; import './loraloader_mxd_entry.js'; import './mxd_character_prompts.js'; +import './ltx_sampler_mxd.js'; +import './wan22_video_prep_mxd.js'; diff --git a/web/ltx_sampler_mxd.js b/web/ltx_sampler_mxd.js new file mode 100644 index 0000000..9ef05cd --- /dev/null +++ b/web/ltx_sampler_mxd.js @@ -0,0 +1,266 @@ +import { app } from "../../scripts/app.js"; +import { api } from "../../scripts/api.js"; + +const LTX_SAMPLER_NODE_TYPES = new Set(["LTXKSampler_MXD", "LTXKSampler2_MXD"]); +const CUSTOM_SIGMAS_MODE = "Custom Sigmas"; +const ltxPreviewImages = {}; +const ltxPreviewTimers = {}; +const ltxPreviewPaused = {}; +const ltxPreviewAutoPaused = {}; +const textDecoder = new TextDecoder(); + +function getWidget(node, name) { + return node.widgets?.find((widget) => widget.name === name); +} + +function hideWidget(widget) { + if (!widget._mxdOriginalComputeSize) { + widget._mxdOriginalComputeSize = widget.computeSize; + } + + widget.hidden = true; + widget.disabled = true; + widget.computeSize = () => [0, -4]; +} + +function showWidget(widget) { + widget.hidden = false; + widget.disabled = false; + + if (widget._mxdOriginalComputeSize) { + widget.computeSize = widget._mxdOriginalComputeSize; + } +} + +function resizeNodeToWidgets(node) { + if (!node.computeSize || !node.setSize) { + return; + } + + const computed = node.computeSize(); + const currentWidth = node.size?.[0] ?? computed[0]; + node.setSize([Math.max(currentWidth, computed[0]), computed[1]]); +} + +function updateCustomSigmasVisibility(node) { + const modeWidget = getWidget(node, "mode"); + const sigmasWidget = getWidget(node, "custom_sigmas"); + + if (!modeWidget || !sigmasWidget) { + return; + } + + if (modeWidget.value === CUSTOM_SIGMAS_MODE) { + showWidget(sigmasWidget); + } else { + hideWidget(sigmasWidget); + } + + resizeNodeToWidgets(node); + app.canvas?.setDirty(true, true); +} + +function getNodeById(id) { + return app.graph?._nodes_by_id?.[id] ?? app.graph?.getNodeById?.(id); +} + +function updatePauseButton(id, buttonEl) { + if (!buttonEl) { + return; + } + + if (ltxPreviewPaused[id]) { + buttonEl.textContent = "Play"; + buttonEl.title = "Resume latent preview playback"; + } else { + buttonEl.textContent = "Pause"; + buttonEl.title = "Pause latent preview playback"; + } +} + +function setLatentPreviewPaused(id, paused) { + ltxPreviewPaused[id] = paused; + const node = getNodeById(id); + const widget = node ? getWidget(node, "ltxlatentpreview") : null; + updatePauseButton(id, widget?.pauseEl); +} + +function getPreviewContext(id, width, height) { + const node = getNodeById(id); + if (!node) { + return null; + } + + let widget = getWidget(node, "ltxlatentpreview"); + if (!widget) { + const previewEl = document.createElement("div"); + previewEl.style.width = "100%"; + previewEl.style.position = "relative"; + + const canvasEl = document.createElement("canvas"); + canvasEl.style.width = "100%"; + canvasEl.style.display = "block"; + previewEl.appendChild(canvasEl); + + const pauseEl = document.createElement("button"); + pauseEl.textContent = "Pause"; + pauseEl.style.position = "absolute"; + pauseEl.style.right = "6px"; + pauseEl.style.bottom = "6px"; + pauseEl.style.padding = "1px 6px"; + pauseEl.style.fontSize = "11px"; + pauseEl.style.lineHeight = "1.2"; + pauseEl.style.opacity = "0.85"; + pauseEl.style.cursor = "pointer"; + previewEl.appendChild(pauseEl); + + widget = node.addDOMWidget("ltxlatentpreview", "ltxcanvas", previewEl, { + serialize: false, + hideOnZoom: false, + }); + widget.serialize = false; + widget.canvasEl = canvasEl; + widget.pauseEl = pauseEl; + widget.computeSize = function (availableWidth) { + if (!this.aspectRatio) { + return [availableWidth, -4]; + } + return [availableWidth, (node.size[0] - 20) / this.aspectRatio + 10]; + }; + + pauseEl.addEventListener("pointerdown", (event) => { + event.preventDefault(); + event.stopImmediatePropagation(); + event.stopPropagation(); + }, true); + pauseEl.addEventListener("click", (event) => { + event.preventDefault(); + event.stopImmediatePropagation(); + event.stopPropagation(); + setLatentPreviewPaused(id, !ltxPreviewPaused[id]); + }, true); + } + updatePauseButton(id, widget.pauseEl); + + const canvasEl = widget.canvasEl || widget.element; + if (canvasEl.width !== width || canvasEl.height !== height) { + widget.aspectRatio = width / height; + canvasEl.width = width; + canvasEl.height = height; + resizeNodeToWidgets(node); + } + return canvasEl.getContext("2d"); +} + +function beginLatentPreview(id, rate) { + clearInterval(ltxPreviewTimers[id]); + let displayIndex = 0; + ltxPreviewAutoPaused[id] = false; + setLatentPreviewPaused(id, false); + const startNode = getNodeById(id); + if (startNode) { + startNode.progress = 0; + } + + ltxPreviewTimers[id] = setInterval(() => { + const node = getNodeById(id); + if (!node) { + clearInterval(ltxPreviewTimers[id]); + delete ltxPreviewTimers[id]; + delete ltxPreviewAutoPaused[id]; + return; + } + if (node.progress == null) { + if (!ltxPreviewAutoPaused[id]) { + ltxPreviewAutoPaused[id] = true; + setLatentPreviewPaused(id, true); + } + } else { + ltxPreviewAutoPaused[id] = false; + } + if (ltxPreviewPaused[id]) { + return; + } + + const images = ltxPreviewImages[id]; + const image = images?.[displayIndex]; + if (!image) { + return; + } + getPreviewContext(id, image.width, image.height)?.drawImage(image, 0, 0); + displayIndex = (displayIndex + 1) % images.length; + app.canvas?.setDirty(true, true); + }, 1000 / Math.max(1, rate || 8)); +} + +api.addEventListener("VHS_latentpreview", ({ detail }) => { + if (detail.id == null) { + return; + } + + ltxPreviewImages[detail.id] = []; + ltxPreviewImages[detail.id].length = detail.length; + const idParts = String(detail.id).split(":"); + for (let i = 1; i <= idParts.length; i++) { + const id = idParts.slice(0, i).join(":"); + ltxPreviewImages[id] = ltxPreviewImages[detail.id]; + beginLatentPreview(id, detail.rate); + } +}); + +api.addEventListener("b_preview", async (event) => { + if (Object.keys(ltxPreviewTimers).length === 0) { + return; + } + + const header = new DataView(await event.detail.slice(0, 24).arrayBuffer()); + const index = header.getUint32(4); + const idLength = header.getUint8(8); + const id = textDecoder.decode(header.buffer.slice(9, 9 + idLength)); + const images = ltxPreviewImages[id]; + if (!images) { + return; + } + + event.preventDefault(); + event.stopImmediatePropagation(); + event.stopPropagation(); + images[index] = await window.createImageBitmap(event.detail.slice(24)); +}, true); + +app.registerExtension({ + name: "ComfyUI-MaxedOut.LTXSamplerMXD", + + beforeRegisterNodeDef(nodeType, nodeData) { + if (!LTX_SAMPLER_NODE_TYPES.has(nodeData.name)) { + return; + } + + const onNodeCreated = nodeType.prototype.onNodeCreated; + nodeType.prototype.onNodeCreated = function () { + const result = onNodeCreated?.apply(this, arguments); + const node = this; + const modeWidget = getWidget(node, "mode"); + + if (modeWidget && !modeWidget._mxdLtxCallbackWrapped) { + const originalCallback = modeWidget.callback; + modeWidget.callback = function () { + const callbackResult = originalCallback?.apply(this, arguments); + updateCustomSigmasVisibility(node); + return callbackResult; + }; + modeWidget._mxdLtxCallbackWrapped = true; + } + + updateCustomSigmasVisibility(node); + return result; + }; + + const onConfigure = nodeType.prototype.onConfigure; + nodeType.prototype.onConfigure = function () { + const result = onConfigure?.apply(this, arguments); + requestAnimationFrame(() => updateCustomSigmasVisibility(this)); + return result; + }; + }, +}); diff --git a/web/mxd_buttons.css b/web/mxd_buttons.css index 2c9f13c..d21e395 100644 --- a/web/mxd_buttons.css +++ b/web/mxd_buttons.css @@ -1,4 +1,4 @@ -:not(#fakeid) .rgthree-button-reset { +:not(#fakeid) .mxd-button-reset { position: relative; appearance: none; cursor: pointer; @@ -9,7 +9,7 @@ margin: 0; } -:not(#fakeid) .rgthree-button { +:not(#fakeid) .mxd-button { --padding-top: 7px; --padding-bottom: 9px; --padding-x: 16px; @@ -34,7 +34,7 @@ align-items: center; justify-content: center; } -:not(#fakeid) .rgthree-button::before, :not(#fakeid) .rgthree-button::after { +:not(#fakeid) .mxd-button::before, :not(#fakeid) .mxd-button::after { content: ""; display: block; position: absolute; @@ -47,55 +47,55 @@ background: linear-gradient(to bottom, rgba(255, 255, 255, 0.06), rgba(0, 0, 0, 0.15)); mix-blend-mode: screen; } -:not(#fakeid) .rgthree-button::after { +:not(#fakeid) .mxd-button::after { mix-blend-mode: multiply; } -:not(#fakeid) .rgthree-button:hover { +:not(#fakeid) .mxd-button:hover { background: #303030; } -:not(#fakeid) .rgthree-button:active { +:not(#fakeid) .mxd-button:active { box-shadow: 0px 0px 0px rgba(0, 0, 0, 0); background: #121212; padding: calc(var(--padding-top) + 1px) calc(var(--padding-x) - 1px) calc(var(--padding-bottom) - 1px) calc(var(--padding-x) + 1px); } -:not(#fakeid) .rgthree-button:active::before, :not(#fakeid) .rgthree-button:active::after { +:not(#fakeid) .mxd-button:active::before, :not(#fakeid) .mxd-button:active::after { box-shadow: 1px 1px 0px rgba(255, 255, 255, 0.15), inset 1px 1px 0px rgba(0, 0, 0, 0.5), inset 1px 3px 5px rgba(0, 0, 0, 0.33); } -:not(#fakeid) .rgthree-button.-blue { +:not(#fakeid) .mxd-button.-blue { background: #346599 !important; } -:not(#fakeid) .rgthree-button.-blue:hover { +:not(#fakeid) .mxd-button.-blue:hover { background: #3b77b8 !important; } -:not(#fakeid) .rgthree-button.-blue:active { +:not(#fakeid) .mxd-button.-blue:active { background: #1d5086 !important; } -:not(#fakeid) .rgthree-button.-green { +:not(#fakeid) .mxd-button.-green { background: linear-gradient(to bottom, rgba(255, 255, 255, 0.06), rgba(0, 0, 0, 0.15)), #14580b; } -:not(#fakeid) .rgthree-button.-green:hover { +:not(#fakeid) .mxd-button.-green:hover { background: linear-gradient(to bottom, rgba(255, 255, 255, 0.06), rgba(0, 0, 0, 0.15)), #1a6d0f; } -:not(#fakeid) .rgthree-button.-green:active { +:not(#fakeid) .mxd-button.-green:active { background: linear-gradient(to bottom, rgba(0, 0, 0, 0.15), rgba(255, 255, 255, 0.06)), #0f3f09; } -:not(#fakeid) .rgthree-button[disabled] { +:not(#fakeid) .mxd-button[disabled] { box-shadow: none; background: #666 !important; color: #aaa; pointer-events: none; } -:not(#fakeid) .rgthree-button[disabled]::before, :not(#fakeid) .rgthree-button[disabled]::after { +:not(#fakeid) .mxd-button[disabled]::before, :not(#fakeid) .mxd-button[disabled]::after { display: none; } -:not(#fakeid) .rgthree-comfybar-top-button-group { +:not(#fakeid) .mxd-comfybar-top-button-group { font-size: 0; flex: 1 1 auto; display: flex; align-items: stretch; } -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button { +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button { margin: 0; flex: 1 1; height: 36px; @@ -104,27 +104,27 @@ background: var(--p-button-secondary-background); color: var(--p-button-secondary-color); } -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button.-primary { +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button.-primary { background: var(--p-button-primary-background); color: var(--p-button-primary-color); } -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button::before, :not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button::after { +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button::before, :not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button::after { border-radius: 0; } -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button svg { +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button svg { fill: currentColor; width: 28px; height: 28px; } -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button:first-of-type, -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button:first-of-type::before, -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button:first-of-type::after { +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button:first-of-type, +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button:first-of-type::before, +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button:first-of-type::after { border-top-left-radius: 0.33rem; border-bottom-left-radius: 0.33rem; } -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button:last-of-type, -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button:last-of-type::before, -:not(#fakeid) .rgthree-comfybar-top-button-group .rgthree-comfybar-top-button:last-of-type::after { +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button:last-of-type, +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button:last-of-type::before, +:not(#fakeid) .mxd-comfybar-top-button-group .mxd-comfybar-top-button:last-of-type::after { border-top-right-radius: 0.33rem; border-bottom-right-radius: 0.33rem; } diff --git a/web/mxd_dialog.css b/web/mxd_dialog.css index c9efdf6..7361ee3 100644 --- a/web/mxd_dialog.css +++ b/web/mxd_dialog.css @@ -1,5 +1,5 @@ @charset "UTF-8"; -.rgthree-dialog { +.mxd-dialog { outline: 0; border: 0; border-radius: 6px; @@ -13,21 +13,21 @@ padding: 0; max-height: calc(100% - 32px); } -.rgthree-dialog *, .rgthree-dialog *::before, .rgthree-dialog *::after { +.mxd-dialog *, .mxd-dialog *::before, .mxd-dialog *::after { box-sizing: inherit; } -.rgthree-dialog-container > * { +.mxd-dialog-container > * { padding: 8px 16px; } -.rgthree-dialog-container > *:first-child { +.mxd-dialog-container > *:first-child { padding-top: 16px; } -.rgthree-dialog-container > *:last-child { +.mxd-dialog-container > *:last-child { padding-bottom: 16px; } -.rgthree-dialog.-iconed::after { +.mxd-dialog.-iconed::after { content: ""; font-size: 276px; position: absolute; @@ -43,68 +43,68 @@ z-index: -1; } -.rgthree-dialog.-iconed.-help::after { +.mxd-dialog.-iconed.-help::after { content: "🛟"; } -.rgthree-dialog.-iconed.-settings::after { +.mxd-dialog.-iconed.-settings::after { content: "⚙️"; } @media (max-width: 832px) { - .rgthree-dialog { + .mxd-dialog { max-width: calc(100% - 32px); } } -.rgthree-dialog-container-title { +.mxd-dialog-container-title { display: flex; flex-direction: row; align-items: center; justify-content: start; } -.rgthree-dialog-container-title > svg:first-child { +.mxd-dialog-container-title > svg:first-child { width: 36px; height: 36px; margin-right: 16px; } -.rgthree-dialog-container-title h2 { +.mxd-dialog-container-title h2 { font-size: 1.375rem; margin: 0; font-weight: bold; } -.rgthree-dialog-container-title h2 small { +.mxd-dialog-container-title h2 small { font-size: 0.8125rem; font-weight: normal; opacity: 0.75; } -.rgthree-dialog-container-content { +.mxd-dialog-container-content { overflow: auto; max-height: calc(100vh - 200px); /* Arbitrary height to copensate for margin, title, and footer.*/ } -.rgthree-dialog-container-content p { +.mxd-dialog-container-content p { font-size: 0.8125rem; margin-top: 0; } -.rgthree-dialog-container-content ul li p { +.mxd-dialog-container-content ul li p { margin-bottom: 4px; } -.rgthree-dialog-container-content ul li p + p { +.mxd-dialog-container-content ul li p + p { margin-top: 0.5em; } -.rgthree-dialog-container-content ul li ul { +.mxd-dialog-container-content ul li ul { margin-top: 0.5em; margin-bottom: 1em; } -.rgthree-dialog-container-content p code { +.mxd-dialog-container-content p code { display: inline-block; padding: 2px 4px; margin: 0px 2px; @@ -113,12 +113,12 @@ background: rgba(255, 255, 255, 0.1); } -.rgthree-dialog-container-footer { +.mxd-dialog-container-footer { display: flex; align-items: center; justify-content: center; } -body.rgthree-dialog-open > *:not(.rgthree-dialog):not(.rgthree-top-messages-container) { +body.mxd-dialog-open > *:not(.mxd-dialog):not(.mxd-top-messages-container) { filter: blur(5px); } diff --git a/web/mxd_dialog.js b/web/mxd_dialog.js index fdc03e7..7d9c2ef 100644 --- a/web/mxd_dialog.js +++ b/web/mxd_dialog.js @@ -3,16 +3,16 @@ export class MxdDialog extends EventTarget { constructor(options) { super(); this.options = options; - let container = $el("div.rgthree-dialog-container"); + let container = $el("div.mxd-dialog-container"); this.element = $el("dialog", { - classes: ["rgthree-dialog", options.class || ""], + classes: ["mxd-dialog", options.class || ""], child: container, parent: document.body, events: { click: (event) => { if (!this.element.open || event.target === container || - getClosestOrSelf(event.target, `.rgthree-dialog-container`) === container) { + getClosestOrSelf(event.target, `.mxd-dialog-container`) === container) { return; } return this.close(); @@ -22,7 +22,7 @@ export class MxdDialog extends EventTarget { this.element.addEventListener("close", (event) => { this.onDialogElementClose(); }); - this.titleElement = $el("div.rgthree-dialog-container-title", { + this.titleElement = $el("div.mxd-dialog-container-title", { parent: container, children: !options.title ? null @@ -34,11 +34,11 @@ export class MxdDialog extends EventTarget { : options.title : options.title, }); - this.contentElement = $el("div.rgthree-dialog-container-content", { + this.contentElement = $el("div.mxd-dialog-container-content", { parent: container, child: options.content, }); - const footerEl = $el("footer.rgthree-dialog-container-footer", { parent: container }); + const footerEl = $el("footer.mxd-dialog-container-footer", { parent: container }); for (const button of options.buttons || []) { $el("button", { text: button.label, @@ -56,7 +56,7 @@ export class MxdDialog extends EventTarget { if (options.closeButtonLabel !== false) { $el("button", { text: options.closeButtonLabel || "Close", - className: "rgthree-button", + className: "mxd-button", parent: footerEl, events: { click: (e) => { @@ -76,7 +76,7 @@ export class MxdDialog extends EventTarget { setAttributes(this.contentElement, { children: content }); } show() { - document.body.classList.add("rgthree-dialog-open"); + document.body.classList.add("mxd-dialog-open"); this.element.showModal(); this.dispatchEvent(new CustomEvent("show")); return this; @@ -88,7 +88,7 @@ export class MxdDialog extends EventTarget { this.element.close(); } onDialogElementClose() { - document.body.classList.remove("rgthree-dialog-open"); + document.body.classList.remove("mxd-dialog-open"); this.element.remove(); this.dispatchEvent(new CustomEvent("close", this.getCloseEventDetail())); } diff --git a/web/mxd_dialog_info.js b/web/mxd_dialog_info.js index 2d1a165..5fc3812 100644 --- a/web/mxd_dialog_info.js +++ b/web/mxd_dialog_info.js @@ -18,7 +18,7 @@ const EXTENSION_BASE = new URL(".", import.meta.url).pathname.replace(/\/$/, "") class MxdInfoDialog extends MxdDialog { constructor(file) { const dialogOptions = { - class: "rgthree-info-dialog", + class: "mxd-info-dialog", title: `

Loading...

`, content: "
Loading..
", onBeforeClose: () => true, @@ -60,7 +60,7 @@ class MxdInfoDialog extends MxdDialog { this.setContent(this.getInfoContent()); this.setTitle(this.modelInfo?.name || this.modelInfo?.file || "Unknown"); } else if (action === "copy-trained-words") { - const selected = queryAll(".-rgthree-is-selected", target.closest("tr")); + const selected = queryAll(".-mxd-is-selected", target.closest("tr")); const text = selected.map((el) => el.getAttribute("data-word")).join(", "); await navigator.clipboard.writeText(text); mxdRuntime.showMessage({ @@ -70,7 +70,7 @@ class MxdInfoDialog extends MxdDialog { timeout: 3000, }); } else if (action === "toggle-trained-word") { - target?.classList.toggle("-rgthree-is-selected"); + target?.classList.toggle("-mxd-is-selected"); const tr = target.closest("tr"); if (tr) { const span = query("td:first-child > *", tr); @@ -78,7 +78,7 @@ class MxdInfoDialog extends MxdDialog { if (!small) { small = $el("small", { parent: span }); } - const num = queryAll(".-rgthree-is-selected", tr).length; + const num = queryAll(".-mxd-is-selected", tr).length; small.innerHTML = num ? `${num} selected | Copy` : ""; } } else if (action === "edit-row") { @@ -87,7 +87,7 @@ class MxdInfoDialog extends MxdDialog { const input = td.querySelector("input,textarea"); if (!input) { const fieldName = tr.dataset["fieldName"]; - tr.classList.add("-rgthree-editing"); + tr.classList.add("-mxd-editing"); const isTextarea = fieldName === "userNote"; const rowInput = $el(`${isTextarea ? "textarea" : 'input[type="text"]'}`, { value: td.textContent }); rowInput.addEventListener("keydown", (evt) => { @@ -118,13 +118,13 @@ class MxdInfoDialog extends MxdDialog { const info = this.modelInfo || {}; const civitaiLink = info.links?.find((i) => i.includes("civitai.com/models")); const html = ` -