From eb085ad39a5bcf6e3f53970fae05cdcf058f5919 Mon Sep 17 00:00:00 2001 From: Maxed-Out-99 Date: Fri, 26 Sep 2025 10:53:34 -0700 Subject: [PATCH] Add I2V-specific latent save/load nodes Introduces SaveLatent_I2V_MXD and LoadLatent_I2V_MXD classes for image-to-video workflows. These nodes handle saving/loading latents with conditioning sidecars and preview images, and are registered for use in the node system. --- wan22nodes.py | 176 ++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 176 insertions(+) diff --git a/wan22nodes.py b/wan22nodes.py index 0fd2751..3143bba 100644 --- a/wan22nodes.py +++ b/wan22nodes.py @@ -822,6 +822,178 @@ class wan22EmptyHunyuanLatentVideoMXD: device=comfy.model_management.intermediate_device() ) return ({"samples": latent},) + +# ---------- I2V-specific latent save/load (sidecar conditioning; subclassed loader) ---------- + +class SaveLatent_I2V_MXD: + """ + I2V-only saver that persists: + • latent tensor -> .latent (safetensors via comfy.utils.save_torch_file) + • pos/neg CONDITIONING -> .cond.pt (torch.save; robust for nested tensors) + • preview images to TEMP for UI + """ + TITLE = "Save Latent I2V (with Conditioning)" + CATEGORY = "MXD/Latents (I2V)" + OUTPUT_NODE = True + RETURN_TYPES = () + FUNCTION = "save_and_preview" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "samples": ("LATENT", {"tooltip": "High-noise latent to save for later low-noise finishing."}), + "positive": ("CONDITIONING", {"tooltip": "Positive CONDITIONING after WAN image→video."}), + "negative": ("CONDITIONING", {"tooltip": "Negative CONDITIONING after WAN image→video."}), + "vae": ("VAE", {"tooltip": "Used to decode preview images for UI convenience."}), + "filename_prefix": ("STRING", {"default": "I2V", "tooltip": "Prefix for saved files"}), + }, + "hidden": {"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"}, + } + + def save_and_preview(self, samples, positive, negative, vae, filename_prefix="I2V", + prompt=None, extra_pnginfo=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 + ) + + meta = None + if not args.disable_metadata: + 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 + + 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) + + # ---- save conditioning sidecar (.cond.pt) ---- + cond_path = latent_path.replace(".latent", ".cond.pt") + torch.save({"positive": positive, "negative": negative}, cond_path) + + # ---- previews to TEMP for UI ---- + images = vae.decode(samples["samples"]) + if len(images.shape) == 5: + images = images.reshape(-1, images.shape[-3], images.shape[-2], images.shape[-1]) + + temp_dir = folder_paths.get_temp_directory() + w, h = images[0].shape[1], images[0].shape[0] + preview_prefix = filename_prefix + "_preview" + full_temp_folder, preview_name, temp_counter, temp_subfolder, _ = folder_paths.get_save_image_path( + preview_prefix, temp_dir, w, h + ) + + results = [] + for b, image in enumerate(images): + np_img = (255.0 * image.cpu().numpy()) + img = Image.fromarray(np.clip(np_img, 0, 255).astype(np.uint8)) + fn_with_batch = preview_name.replace("%batch_num%", str(b)) + preview_file = f"{fn_with_batch}_{temp_counter:05}_.png" + img.save(os.path.join(full_temp_folder, preview_file), compress_level=1) + results.append({"filename": preview_file, "subfolder": temp_subfolder, "type": "temp"}) + temp_counter += 1 + + return {"ui": {"images": results}} + + +class LoadLatent_I2V_MXD(LoadLatent_WithParams): + """ + Same outputs as LoadLatent_WithParams plus two CONDITIONING outputs at the end. + Fixes sampler/scheduler enum wiring by setting enums on THIS subclass. + """ + TITLE = "Load Latent I2V (With Params + Conditioning)" + CATEGORY = "MXD/Latents (I2V)" + FUNCTION = "load" + + # Base tuple + two extra + RETURN_TYPES = LoadLatent_WithParams.RETURN_TYPES + ("CONDITIONING", "CONDITIONING") + RETURN_NAMES = LoadLatent_WithParams.RETURN_NAMES + ("positive_conditioning", "negative_conditioning") + + @classmethod + def INPUT_TYPES(s): + # mirror base: build file list + 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) + files.sort() + options = [os.path.relpath(f, folder_paths.get_input_directory()).replace(os.sep, "/") for f in files] + + # pull live enums from KSamplerAdvanced and attach them to THIS CLASS + ks_inputs = KSamplerAdvanced.INPUT_TYPES().get("required", {}) + samplers_enum = ks_inputs.get("sampler_name", ("STRING",))[0] + schedulers_enum = ks_inputs.get("scheduler", ("STRING",))[0] + + # rebuild RETURN_TYPES on THIS CLASS so ports wire correctly + base_rts = ( + "LATENT", + "STRING", + "STRING", + "INT", + "FLOAT", + samplers_enum, + schedulers_enum, + "INT", + "FLOAT", # shift + "STRING", # filename_prefix + ) + s.RETURN_TYPES = base_rts + ("CONDITIONING", "CONDITIONING") + s._SAMPLERS_ENUM = samplers_enum + s._SCHEDULERS_ENUM = schedulers_enum + + return {"required": {"latent": (options, )}} + + @classmethod + def IS_CHANGED(s, latent): + p = folder_paths.get_annotated_filepath(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): + return LoadLatent_WithParams.VALIDATE_INPUTS(latent) + + def load(self, latent): + # Use base loader to get (samples, pos, neg, steps, cfg, sampler_name, scheduler, end_at_step, shift, prefix) + base_tuple = super().load(latent) + + # sidecar conditioning + latent_path = folder_paths.get_annotated_filepath(latent) + cond_path = latent_path.replace(".latent", ".cond.pt") + positive_conditioning, negative_conditioning = [], [] + if os.path.exists(cond_path): + try: + d = torch.load(cond_path, map_location="cpu") + positive_conditioning = d.get("positive", []) + negative_conditioning = d.get("negative", []) + except Exception: + positive_conditioning, negative_conditioning = [], [] + + return base_tuple + (positive_conditioning, negative_conditioning) + # ---------- Node registration ---------- @@ -831,6 +1003,8 @@ NODE_CLASS_MAPPINGS = { "LoadLatents_FromFolder_WithParams": LoadLatents_FromFolder_WithParams, "Wan2_2EmptyLatentImageMXD": Wan2_2EmptyLatentImageMXD, "wan22EmptyHunyuanLatentVideoMXD": wan22EmptyHunyuanLatentVideoMXD, + "SaveLatent_I2V_MXD": SaveLatent_I2V_MXD, + "LoadLatent_I2V_MXD": LoadLatent_I2V_MXD, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -839,4 +1013,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "LoadLatents_FromFolder_WithParams": "Load Latent Batch MXD", "Wan2_2EmptyLatentImageMXD": "Wan 2.2 Empty Latent Image MXD", "wan22EmptyHunyuanLatentVideoMXD": "WAN2.2 Empty Latent Video MXD", + "SaveLatent_I2V_MXD": "Save Latent I2V MXD", + "LoadLatent_I2V_MXD": "Load Latent I2V MXD", }