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: 2f97b1be867e64Author: 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: 3e4e4db2369cdbAuthor: 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
150 lines
7.5 KiB
Python
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",
|
|
} |