Files
kijai-ComfyUI-WanVideoWrapper/onetoall/nodes.py
T
kijai 8b037bce2e Squashed commit of the following:
commit c3eb0f49faf68ab953f1b08b7e00225e041e5d0b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 12:55:49 2025 +0200

    move workflow

commit e129e25c26f9b55b527dd3e9f15c6e3f215af11f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 11:17:17 2025 +0200

    Fix padding

commit f252f34eff5cc15ec6fc475f929cafa3e5b7f46c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 01:38:17 2025 +0200

    Add long video example

commit 09ceab808b67a3b2fb7d1ee5fc0a1ad667739e2a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 01:31:48 2025 +0200

    Support extension

commit 7ca221874e8a2cabfc766c51bb63774fde3c851b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 12:28:29 2025 +0200

    Might as well not even do control pass on uncond...

commit b55caf299e4d89148f5885e8e56bf8e411472dc3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 12:15:59 2025 +0200

    Cfg fixes

commit fd54ba23e6746acb33a8bf124e5bc7de9d947ff1
Merge: 2f97b1b e867e64
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 10:39:55 2025 +0200

    Merge branch 'main' into onetoall

commit 2f97b1bd887367962542b9a6058f9f6e3c4ad4d7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 09:32:09 2025 +0200

    Add ref_mask input

commit 74cad232fd35347c50f2ed7465ff13e179ef8402
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 03:44:42 2025 +0200

    Update nodes_model_loading.py

commit 01a038eb4a30f29d868fbaef190e6e90da1a058d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 03:11:08 2025 +0200

    Fix indentation

commit a95f4d6eaa4468e818910fec7ba11e1f92423d9b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:54:47 2025 +0200

    Update model.py

commit ad006985a1bafdf5941c0fa85a47852eb20a818a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:54:19 2025 +0200

    Fix token replace

commit b5f0f44f1720586950756ad142a538e04814270f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:50:52 2025 +0200

    Don't use token replace by default

commit 874174ec2921c528a4373097fd0bebbbb5257606
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:24:47 2025 +0200

    Create WanToAllAnimation_test.json

commit 9e6175855618c94c1bcb89c4b89879219410ce53
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:23:15 2025 +0200

    Add token replacement

commit 41fd76dfcbf0e70a3a7308a6fa0652fb492ed1f6
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 00:45:33 2025 +0200

    Use correct norm for reference attn

commit 705f5dcc8b6cd5fa6fe453f9bd01ffdf43a23078
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 00:11:17 2025 +0200

    cleanup

commit 4f095d97f80da807417d49d9aa7e9ee47145c85f
Merge: 3e4e4db 2369cdb
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 7 18:44:01 2025 +0200

    Merge branch 'main' into onetoall

commit 3e4e4db35d3e266c39d48cd683f60384a737eca5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 7 00:27:23 2025 +0200

    handle controlnet better

commit c5742552a9af4a3ae208f9c2ead6e1105cc2c348
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 17:24:45 2025 +0200

    cleanup

commit c06ff9c06651c32953236802bd7fb385b9cf93ab
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 03:41:02 2025 +0200

    3D rope for controlnet

commit 948ea6b783f54892515cbc9cfe66484913904ee7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 03:08:04 2025 +0200

    pose input scaling

commit 90c2eff3b2d30d3a92ff5c27e4327a0ac80b642c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 02:37:48 2025 +0200

    Cleanup

commit 9f7683422c1aa8ebe4d3380a86be98d6c589b270
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Dec 5 23:29:05 2025 +0200

    pose control

commit 0f217be4d8742741b0f89db50138214302a58dc3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Dec 5 20:55:10 2025 +0200

    Support reference input
2025-12-09 12:56:11 +02:00

150 lines
7.5 KiB
Python

import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddOneToAllReferenceEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"ref_image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"ref_mask": ("MASK",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, ref_mask=None):
updated = dict(embeds)
ref_latent = ref_latent_empty = None
vae.to(device)
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
ref_latent = vae.encode([ref_image_in], device, tiled=False)
ref_mask_in = None
if ref_mask is not None:
ref_mask_in = (ref_mask.unsqueeze(0).repeat(3, 1, 1, 1) * 2 - 1.).to(device, vae.dtype)
else:
ref_mask_in = torch.zeros_like(ref_image_in)-1
ref_mask_latent = vae.encode([ref_mask_in], device, tiled=False)
if ref_mask is not None and not torch.all(ref_mask == 0):
ref_latent_empty = vae.encode([torch.zeros_like(ref_image_in)-1], device, tiled=False)
else:
ref_latent_empty = ref_mask_latent
vae.to(offload_device)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["ref_latent_pos"] = torch.cat([ref_latent, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_latent_neg"] = torch.cat([ref_latent_empty, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_strength"] = strength
updated["one_to_all_embeds"]["ref_start_percent"] = start_percent
updated["one_to_all_embeds"]["ref_end_percent"] = end_percent
return (updated,)
class WanVideoAddOneToAllPoseEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
},
"optional": {
"pose_prefix_image": ("IMAGE",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, pose_images, strength, pose_prefix_image=None, start_percent=0.0, end_percent=1.0):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
pose_images_in = pose_images[..., :3].unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
updated["one_to_all_embeds"]["pose_images"] = pose_images_in
if pose_prefix_image is not None:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_prefix_image.unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
else:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_images_in[:, :, :1]
updated["one_to_all_embeds"]["controlnet_strength"] = strength
updated["one_to_all_embeds"]["controlnet_start_percent"] = start_percent
updated["one_to_all_embeds"]["controlnet_end_percent"] = end_percent
return (updated,)
class WanVideoAddOneToAllExtendEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"prev_latents": ("LATENT", {"tooltip": "Previous latents to be used to continue generation"}),
"window_size": ("INT", {"default": 81, "min": 1, "max": 256, "step": 1, "tooltip": "Number of new frames to generate" }),
"overlap": ("INT", {"default": 5, "min": 0, "max": 64, "step": 1, "tooltip": "Number of overlapping frames between previous and new frames" }),
"frames_processed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Number of frames already processed in the video" }),
"if_not_enough_frames": (["pad_with_last", "error"], {"default": "pad_with_last", "tooltip": "What to do if there are not enough frames in pose_images for the window"}),
},
"optional": {
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "IMAGE",)
RETURN_NAMES = ("image_embeds", "pose_slice",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, prev_latents, if_not_enough_frames, window_size=81, overlap=5, frames_processed=0, pose_images=None):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["prev_latents"] = prev_latents["samples"][0]
if pose_images is not None:
pose_images_in = pose_images.clone()[..., :3]
start = max(0, frames_processed - overlap)
end = start + window_size
log.info(f"Extracting pose images from {start} to {end}")
if start >= pose_images_in.shape[0]:
raise ValueError(f"start index {start} exceeds pose images length {pose_images_in.shape[0]}")
if end > pose_images_in.shape[0]:
if if_not_enough_frames == "pad_with_last":
padding_needed = end - pose_images_in.shape[0]
pose_images_in = torch.cat([pose_images_in, pose_images_in[-1:].repeat(padding_needed, 1, 1, 1)], dim=0)
log.info(f"Not enough frames, padding with {padding_needed} frames to reach {end} total frames")
else:
raise ValueError(f"end index {end} exceeds pose images length {pose_images.shape[0]}")
pose_slice = pose_images_in[start:end]
else:
pose_slice = torch.zeros((1, 64, 64, 3))
return (updated, pose_slice)
NODE_CLASS_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": WanVideoAddOneToAllReferenceEmbeds,
"WanVideoAddOneToAllPoseEmbeds": WanVideoAddOneToAllPoseEmbeds,
"WanVideoAddOneToAllExtendEmbeds": WanVideoAddOneToAllExtendEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": "WanVideo Add OneToAll Reference Embeds",
"WanVideoAddOneToAllPoseEmbeds": "WanVideo Add OneToAll Pose Embeds",
"WanVideoAddOneToAllExtendEmbeds": "WanVideo Add OneToAll Extend Embeds",
}