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
This commit is contained in:
kijai
2025-12-09 12:56:11 +02:00
parent e867e642d4
commit 8b037bce2e
11 changed files with 10605 additions and 166 deletions
+10
View File
@@ -84,6 +84,13 @@ except Exception as e:
STEADYDANCER_NODE_CLASS_MAPPINGS = {}
STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS = {}
try:
from .onetoall.nodes import NODE_CLASS_MAPPINGS as ONETOALL_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ONETOALL_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: OneToAll nodes not available due to error in importing them: {e}")
ONETOALL_NODE_CLASS_MAPPINGS = {}
ONETOALL_NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
@@ -108,6 +115,8 @@ NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(STEADYDANCER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(ONETOALL_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
@@ -133,5 +142,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(ONETOALL_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
-45
View File
@@ -17,11 +17,6 @@ from diffusers.models.transformers.transformer_wan import (
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def zero_module(module):
for p in module.parameters():
nn.init.zeros_(p)
return module
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r"""
@@ -153,7 +148,6 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
for _ in range(len(self.blocks)):
controlnet_block = nn.Linear(inner_dim, out_proj_dim)
controlnet_block = zero_module(controlnet_block)
self.controlnet_blocks.append(controlnet_block)
self.gradient_checkpointing = False
@@ -240,42 +234,3 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
return Transformer2DModelOutput(sample=controlnet_hidden_states)
if __name__ == "__main__":
parameters = {
"added_kv_proj_dim": None,
"attention_head_dim": 128,
"cross_attn_norm": True,
"eps": 1e-06,
"ffn_dim": 8960,
"freq_dim": 256,
"image_dim": None,
"in_channels": 3,
"num_attention_heads": 12,
"num_layers": 2,
"patch_size": [1, 2, 2],
"qk_norm": "rms_norm_across_heads",
"rope_max_seq_len": 1024,
"text_dim": 4096,
"downscale_coef": 8,
"out_proj_dim": 12 * 128,
"vae_channels": 16
}
controlnet = WanControlnet(**parameters)
hidden_states = torch.rand(1, 16, 13, 60, 90)
timestep = torch.tensor([1000]).repeat(17550).unsqueeze(0) #torch.randint(low=0, high=1000, size=(1,), dtype=torch.long)
encoder_hidden_states = torch.rand(1, 512, 4096)
controlnet_states = torch.rand(1, 3, 49, 480, 720)
controlnet_hidden_states = controlnet(
hidden_states=hidden_states,
timestep=timestep,
encoder_hidden_states=encoder_hidden_states,
controlnet_states=controlnet_states,
return_dict=False
)
print("Output states count", len(controlnet_hidden_states[0]))
for out_hidden_states in controlnet_hidden_states[0]:
print(out_hidden_states.shape)
File diff suppressed because it is too large Load Diff
+49 -14
View File
@@ -6,7 +6,7 @@ import numpy as np
from tqdm import tqdm
import re
from .wanvideo.modules.model import WanModel, LoRALinearLayer
from .wanvideo.modules.model import WanModel, LoRALinearLayer, WanRMSNorm
from .wanvideo.modules.t5 import T5EncoderModel
from .wanvideo.modules.clip import CLIPModel
from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
@@ -852,18 +852,18 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
total=param_count,
leave=True):
block_idx = vace_block_idx = None
if "vace_blocks." in name:
if name.startswith("vace_blocks."):
try:
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
except Exception:
vace_block_idx = None
elif "blocks." in name and "face" not in name:
elif name.startswith("blocks.") and "face" not in name:
try:
block_idx = int(name.split("blocks.")[1].split(".")[0])
except Exception:
block_idx = None
if "loras" in name or "controlnet" in name:
if "loras" in name:
continue
# GGUF: skip GGUFParameter params
@@ -1310,7 +1310,7 @@ class WanVideoModelLoader:
if dim == 1536:
model_variant = "1_3B"
if dim == 3072:
log.info(f"5B model detected, no Teacache or MagCache coefficients available, consider using EasyCache for this model")
log.info("5B model detected, no Teacache or MagCache coefficients available, consider using EasyCache for this model")
if "high" in model.lower() or "low" in model.lower():
if "i2v" in model.lower():
@@ -1375,7 +1375,7 @@ class WanVideoModelLoader:
with init_empty_weights():
transformer.audio_model = WanModel(**TRANSFORMER_CONFIG).eval()
from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm
from .wanvideo.modules.model import WanLayerNorm
for block in transformer.blocks:
block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim)
@@ -1477,16 +1477,12 @@ class WanVideoModelLoader:
# Additional cond latents
if "add_conv_in.weight" in sd:
def zero_module(module):
for p in module.parameters():
torch.nn.init.zeros_(p)
return module
inner_dim = sd["add_conv_in.weight"].shape[0]
add_cond_in_dim = sd["add_conv_in.weight"].shape[1]
attn_cond_in_dim = sd["attn_conv_in.weight"].shape[1]
transformer.add_conv_in = torch.nn.Conv3d(add_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
transformer.add_proj = zero_module(torch.nn.Linear(inner_dim, inner_dim))
transformer.attn_conv_in = torch.nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
transformer.add_conv_in = nn.Conv3d(add_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
transformer.add_proj = nn.Linear(inner_dim, inner_dim)
transformer.attn_conv_in = nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
# Bindweave text_projection
if "text_projection.0.weight" in sd:
@@ -1516,6 +1512,45 @@ class WanVideoModelLoader:
transformer.condition_embedding_align = PoseRefNetNoBNV3(in_channels_x=16, in_channels_c=16, hidden_dim=128, num_heads=8) # Frame-wise Attention Alignment Unit
if "image_to_cond.conv_in.bias" in sd:
# One-to-all
from .onetoall.controlnet import MiniHunyuanEncoder, MiniEncoder2D
from .onetoall.refextractor_2d import WanRefextractor, WanAttentionBlock
log.info("One-to-all model detected, patching model...")
with init_empty_weights():
transformer.image_to_cond = MiniEncoder2D(
in_channels = sd["image_to_cond.conv_in.bias"].shape[0],
out_channels = in_channels,
down_block_types= ("DownEncoderBlockInflated","DownEncoderBlockInflated","DownEncoderBlockInflated"),
block_out_channels=(16, 16, 16),
norm_num_groups = 4,
layers_per_block = 1,
spatial_compression_ratio=1
)
transformer.input_hint_block = MiniHunyuanEncoder(
in_channels=3,
out_channels=in_channels,
block_out_channels=(16, 16, 16, 16),
norm_num_groups=4,
layers_per_block=1,
spatial_compression_ratio=16
)
controlnet_layers = 1
transformer.controlnet = nn.Module()
transformer.controlnet.blocks = nn.ModuleList([WanAttentionBlock(in_features, out_features, ffn_dim, ffn2_dim, num_heads) for _ in range(controlnet_layers)])
transformer.controlnet_zero = nn.ModuleList([nn.Linear(in_features, out_features) for _ in range(controlnet_layers)])
transformer.refextractor = WanRefextractor(
patch_size=(1, 2, 2), in_dim=sd["refextractor.patch_embedding.weight"].shape[1],
dim=dim, in_features=in_features, out_features=out_features, ffn_dim=ffn_dim, ffn2_dim=ffn2_dim,
num_heads=num_heads, num_layers=7)
for block in transformer.blocks:
block.ref_attn_k_img = nn.Linear(in_features, out_features)
block.ref_attn_v_img = nn.Linear(in_features, out_features)
block.ref_attn_norm_k_img = WanRMSNorm(out_features, eps=1e-6)
comfy_model.diffusion_model = transformer
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
@@ -1580,7 +1615,7 @@ class WanVideoModelLoader:
if gguf:
raise ValueError("GGUF models don't support vram management")
from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm
from .wanvideo.modules.model import WanLayerNorm
total_params_in_model = sum(p.numel() for p in patcher.model.diffusion_model.parameters())
log.info(f"Total number of parameters in the loaded model: {total_params_in_model}")
+38 -5
View File
@@ -1182,6 +1182,31 @@ class WanVideoSampler:
sdancer_data = sdancer_embeds.copy()
sdancer_data = dict_to_device(sdancer_data, device, dtype)
# One-to-all-Animation
one_to_all_embeds = image_embeds.get("one_to_all_embeds", None)
one_to_all_data = prev_latents = None
latents_to_not_step = 0
if one_to_all_embeds is not None:
log.info("Using One-to-All embeddings:")
for k, v in one_to_all_embeds.items():
log.info(f" {k}: {v.shape if isinstance(v, torch.Tensor) else v}")
one_to_all_data = one_to_all_embeds.copy()
one_to_all_data = dict_to_device(one_to_all_data, device, dtype)
if one_to_all_embeds.get("pose_images") is not None:
pose_images_in = one_to_all_data.pop("pose_images")
pose_images = transformer.input_hint_block(pose_images_in)
if one_to_all_embeds.get("ref_latent_pos") is not None:
pose_prefix_image = transformer.input_hint_block(one_to_all_data.pop("pose_prefix_image"))
pose_images = torch.cat([pose_prefix_image, pose_images],dim=2)
one_to_all_data["controlnet_tokens"] = pose_images.flatten(2).transpose(1, 2)
prev_latents = one_to_all_data.get("prev_latents", None)
if prev_latents is not None:
log.info(f"Using previous latents for One-to-All Animation with shape: {prev_latents.shape}")
latent[:, :prev_latents.shape[1]] = prev_latents.to(latent)
one_to_all_data["token_replace"] = True
latents_to_not_step = prev_latents.shape[1]
one_to_all_data["num_latent_frames_to_replace"] = latents_to_not_step
#region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
@@ -1477,6 +1502,7 @@ class WanVideoSampler:
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
"sdancer_input": sdancer_input, # SteadyDancer input
"one_to_all_input": one_to_all_data, # One-to-All input
}
batch_size = 1
@@ -3070,11 +3096,16 @@ class WanVideoSampler:
new_latent.append(latent_slice[:, j:j+1])
latent = torch.cat(new_latent, dim=1)
else:
latent = sample_scheduler.step(
noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else noise_pred.unsqueeze(0),
timestep,
latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else latent.unsqueeze(0),
**scheduler_step_args)[0].squeeze(0)
if latents_to_not_step > 0:
raw_latent = latent[:, :latents_to_not_step]
noise_pred_in = noise_pred[:, latents_to_not_step:]
latent = latent[:, latents_to_not_step:]
elif recammaster is not None or mocha_embeds is not None:
noise_pred_in = noise_pred[:, :orig_noise_len]
latent = latent[:, :orig_noise_len]
else:
noise_pred_in = noise_pred
latent = sample_scheduler.step(noise_pred_in.unsqueeze(0), timestep, latent.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
if noise_pred_flipped is not None:
latent_backwards = sample_scheduler_flipped.step(
noise_pred_flipped.unsqueeze(0),
@@ -3083,6 +3114,8 @@ class WanVideoSampler:
**scheduler_step_args)[0].squeeze(0)
latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5
if latents_to_not_step > 0:
latent = torch.cat([raw_latent, latent], dim=1)
if latent_ovi is not None:
latent_ovi = sample_scheduler_ovi.step(noise_pred_ovi.unsqueeze(0), t, latent_ovi.to(device).unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
+440
View File
@@ -0,0 +1,440 @@
import torch
from torch import nn
from torch.nn import functional as F
from einops import rearrange
import numpy as np
from typing import Tuple
from .unet_causal_3d_blocks import get_down_block3d, CausalConv3d
class ControlNetCausalConditioningEmbedding(nn.Module):
def __init__(self, conditioning_embedding_channels: int, conditioning_channels: int = 3, block_out_channels: Tuple[int, ...] = (16, 32, 96, 256)):
super().__init__()
self.conv_in = CausalConv3d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class MiniHunyuanEncoder(nn.Module):
'''
a direct copy of hunyuan encoder
'''
def __init__(
self,
in_channels = 3,
out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512],
layers_per_block = 2,
norm_num_groups = 32,
act_fn: str = "silu",
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
def forward(self, sample):
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
sample = self.conv_in(sample)
# down
for down_block in self.down_blocks:
sample = down_block(sample)
sample = self.conv_out(sample)
return sample
class ControlNetConditioningEmbedding(nn.Module):
"""
Quoting from https://arxiv.org/abs/2302.05543: "Stable Diffusion uses a pre-processing method similar to VQ-GAN
[11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized
training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the
convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides
(activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full
model) to encode image-space conditions ... into feature maps ..."
"""
def __init__(
self,
conditioning_embedding_channels: int,
conditioning_channels: int = 3,
block_out_channels: Tuple[int, ...] = (16, 32, 96, 256),
):
super().__init__()
self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class InflatedGroupNorm(nn.GroupNorm):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class InflatedConv3d(nn.Conv2d):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class ResnetBlockInflated(nn.Module):
def __init__(self, *, in_channels, out_channels=None, dropout=0.0, groups=32, groups_out=None, pre_norm=True, eps=1e-6, non_linearity="swish", output_scale_factor=1.0):
super().__init__()
self.pre_norm = pre_norm
self.pre_norm = True
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.output_scale_factor = output_scale_factor
if groups_out is None:
groups_out = groups
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
if non_linearity == "swish":
self.nonlinearity = lambda x: F.silu(x)
elif non_linearity == "silu":
self.nonlinearity = nn.SiLU()
def forward(self, input_tensor, temb):
if temb is not None:
print("Warning: temb is None in ResnetBlockInflated")
hidden_states = input_tensor
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
return output_tensor
class DownEncoderBlockInflated(nn.Module):
def __init__(self, *, num_layers: int, in_channels: int, out_channels: int, add_downsample: bool, downsample_stride: tuple = (1, 2, 2),
resnet_eps: float = 1e-6, resnet_act_fn: str = "silu", resnet_groups: int = 32):
super().__init__()
self.resnets = nn.ModuleList([ResnetBlockInflated(
in_channels=in_channels if i == 0 else out_channels,
out_channels=out_channels,
eps=resnet_eps,
non_linearity=resnet_act_fn,
groups=resnet_groups,
) for i in range(num_layers)])
self.downsamplers = nn.ModuleList()
if add_downsample:
self.downsamplers.append(
InflatedConv3d(
out_channels,
out_channels,
kernel_size=3,
stride=2,
padding=1,
)
)
self.down_stride = downsample_stride
else:
self.down_stride = (1, 1, 1)
def forward(self, x, temb=None):
for resnet in self.resnets:
x = resnet(x, temb)
for down in self.downsamplers:
x = down(x)
return x
class SFT(nn.Module): # 2D SFT
def __init__(
self, in_channels, out_channels, intermediate_channels=128, groups=32, eps=1e-6):
super().__init__()
self.out_channels = out_channels
self.norm = InflatedGroupNorm(groups, out_channels, eps, affine=True)
self.mlp_shared = nn.Sequential(InflatedConv3d(in_channels, intermediate_channels, kernel_size=3, stride=1, padding=1), nn.SiLU())
self.mlp_gamma = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.mlp_beta = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, hidden_state, condition):
"""
hidden_state : (B, Cout, T, H, W)
condition : (B, Cin, 1, H, W)
"""
hidden_state = self.norm(hidden_state) #2D SFT 2D Norm
actv = self.mlp_shared(condition)
gamma = self.mlp_gamma(actv)
beta = self.mlp_beta(actv)
return torch.addcmul(beta, hidden_state, 1 + gamma)
class MiniEncoder2D(nn.Module):
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: list = (
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
),
block_out_channels: list = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
spatial_compression_ratio: int = 8,
):
super().__init__()
# -------------------------------------------------------------------
# conv in
# -------------------------------------------------------------------
self.conv_in = InflatedConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
self.down_blocks = nn.ModuleList()
output_channel = block_out_channels[0]
num_spatial_down_layers = int(np.log2(spatial_compression_ratio))
for i, block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
# is_final_block = i == len(block_out_channels) - 1
add_spatial_downsample = bool(i < num_spatial_down_layers)
downsample_stride = (1, 2, 2) if add_spatial_downsample else (1, 1, 1)
down_block = DownEncoderBlockInflated(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=add_spatial_downsample,
downsample_stride=downsample_stride,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = InflatedConv3d(output_channel, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
# (B,C,1,H,W)
x = self.conv_in(x)
for block in self.down_blocks:
x = block(x)
return self.conv_out(x)
class Driven_Ref_PoseEncoder(nn.Module):
def __init__(
self, in_channels = 3, out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512], layers_per_block = 2, norm_num_groups = 32,
act_fn: str = "silu", time_compression_ratio: int = 4, spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.ref_pose_encoder = MiniEncoder2D(
in_channels = in_channels,
out_channels = out_channels,
block_out_channels = block_out_channels,
norm_num_groups = norm_num_groups,
layers_per_block = layers_per_block,
spatial_compression_ratio = spatial_compression_ratio,
)
self.sft_layers = nn.ModuleList()
for i, ch in enumerate(block_out_channels):
if i == 0: # 0 层 (H/2,W/2) 不做 SFT
self.sft_layers.append(None)
else: # H/4、H/8、H/16 做 SFT
self.sft_layers.append(
SFT(
in_channels=ch,
out_channels=ch,
intermediate_channels=max(8, ch // 2),
groups=norm_num_groups,
)
)
def forward(self, driven_pose, ref_pose):
# driven_pose b c t h w
# ref_pose b c 1 h w
ref_pose_cond, ref_feats = self.ref_pose_encoder(ref_pose)
x = self.conv_in(driven_pose)
for i, down_block in enumerate(self.down_blocks):
x = down_block(x)
if self.sft_layers[i] is not None:
cond_feat = ref_feats[i]
x = self.sft_layers[i](x, cond_feat)
driven_pose_cond = self.conv_out(x)
return driven_pose_cond, ref_pose_cond
+150
View File
@@ -0,0 +1,150 @@
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",
}
+220
View File
@@ -0,0 +1,220 @@
from typing import Dict, Union
import torch
import torch.nn as nn
from ..wanvideo.modules.model import WanLayerNorm, WanSelfAttention, EmbedND_RifleX, sinusoidal_embedding_1d, apply_rotary_emb_split, apply_rope_comfy1
class WanAttentionBlock(nn.Module):
def __init__(self, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default"):
super().__init__()
self.dim = out_features
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = out_features // num_heads
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.attention_mode = attention_mode
self.rope_func = rope_func
# layers
self.norm1 = WanLayerNorm(self.dim, eps)
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function, head_norm=False)
self.norm2 = WanLayerNorm(self.dim, eps)
self.ffn = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
def get_mod(self, e, modulation):
if e.dim() == 3:
if e.shape[-1] == 512:
e = self.modulation(e)
return e.unsqueeze(2).chunk(6, dim=-1)
return (modulation + e).chunk(6, dim=1) # 1, 6, dim
elif e.dim() == 4:
e_mod = modulation.unsqueeze(2) + e
return [ei.squeeze(1) for ei in e_mod.unbind(dim=1)]
def modulate(self, norm_x, shift_msa, scale_msa):
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
def ffn_chunked(self, mod_x, num_chunks=4):
seq_len = mod_x.shape[1]
if seq_len <= 8192 or num_chunks <= 1:
return self.ffn(mod_x)
return torch.cat([self.ffn(chunk.contiguous()) for chunk in mod_x.chunk(num_chunks, dim=1)], dim=1)
#region attention forward
def forward(self, x, e, seq_lens, freqs, split_rope=True, e_tr=None, tr_start=0, tr_num=0):
use_token_replace = False
if e_tr is not None and tr_num > 0:
tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = self.get_mod(e_tr.to(x.device), self.modulation)
use_token_replace = True
tr_start = tr_start or 0
tr_end = tr_start + (tr_num or 0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
del e
input_dtype = x.dtype
if use_token_replace:
norm_x = self.norm1(x.to(shift_msa.dtype))
input_x = torch.cat([
torch.addcmul(shift_msa, norm_x[:, :tr_start], 1 + scale_msa), # before replace → T
torch.addcmul(tr_shift_msa, norm_x[:, tr_start:tr_end], 1 + tr_scale_msa), # replace segment → t=0
torch.addcmul(shift_msa, norm_x[:, tr_end:], 1 + scale_msa) # after replace → T
], dim=1).to(input_dtype)
else:
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa).to(input_dtype)
del shift_msa, scale_msa
b, s, n, d = *x.shape[:2], self.self_attn.num_heads, self.self_attn.head_dim
h_dim = w_dim = 2 * (self.head_dim // 6)
t_dim = self.head_dim - h_dim - w_dim
q = self.self_attn.norm_q(self.self_attn.q(input_x)).to(self.self_attn.norm_q.weight.dtype).view(b, s, n, d)
if split_rope:
q = apply_rotary_emb_split(q, freqs, t_dim) # Apply split rotary embedding (only to H/W dimensions, leaving T unchanged)
else:
q = apply_rope_comfy1(q, freqs)
k = self.self_attn.norm_k(self.self_attn.k(input_x).to(self.self_attn.norm_k.weight.dtype)).to(input_x.dtype).view(b, s, n, d)
if split_rope:
k = apply_rotary_emb_split(k, freqs, t_dim)
else:
k = apply_rope_comfy1(k, freqs)
v = self.self_attn.v(input_x).view(b, s, n, d)
del input_x
y = self.self_attn.forward(q, k, v, seq_lens)
del q, k, v
if use_token_replace:
x = x + torch.cat([
y[:, :tr_start] * gate_msa,
y[:, tr_start:tr_end] * tr_gate_msa,
y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype)
else:
x = x.addcmul(y, gate_msa)
del y, gate_msa
# ffn
if use_token_replace:
norm2_x = self.norm2(x.to(shift_mlp.dtype))
mod_x = torch.cat([
torch.addcmul(shift_mlp, norm2_x[:, :tr_start], 1 + scale_mlp),
torch.addcmul(tr_shift_mlp, norm2_x[:, tr_start:tr_end], 1 + tr_scale_mlp),
torch.addcmul(shift_mlp, norm2_x[:, tr_end:], 1 + scale_mlp)
], dim=1)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
del shift_mlp, scale_mlp
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
del mod_x
# gate_mlp
if use_token_replace:
x = x + torch.cat([
x_ffn[:, :tr_start] * gate_mlp,
x_ffn[:, tr_start:tr_end] * tr_gate_mlp,
x_ffn[:, tr_end:] * gate_mlp
], dim=1).to(input_dtype)
else:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
del gate_mlp
return x
class WanRefextractor(nn.Module):
def __init__(self, patch_size=(1, 2, 2), in_dim=16, dim=5120, in_features=5120, out_features=5120, ffn_dim=8192, ffn2_dim=8192,
freq_dim=256, num_heads=16, num_layers=32, eps=1e-6,
qk_norm=True, cross_attn_norm=True,
attention_mode='sdpa', rope_func='comfy', rms_norm_function='default',
main_device=torch.device('cuda'), offload_device=torch.device('cpu'), dtype=torch.float16):
super().__init__()
self.patch_size = patch_size
self.freq_dim = freq_dim
self.dim = dim
self.main_device = main_device
self.base_dtype = dtype
self.attention_mode = attention_mode
self.patch_embedding = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.time_embedding = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
self.blocks = nn.ModuleList([
WanAttentionBlock(in_features, out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm, cross_attn_norm, eps, attention_mode="sdpa", rope_func=rope_func, rms_norm_function=rms_norm_function)
for i in range(num_layers)
])
self.ref_blocks = nn.ModuleList([])
for _ in range(len(self.blocks)+1):
self.ref_blocks.append(nn.Linear(in_features, out_features))
d = dim // num_heads
self.rope_embedder = EmbedND_RifleX(d,10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=1, k=0)
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if steps_t is None:
steps_t = t_len
if steps_h is None:
steps_h = h_len
if steps_w is None:
steps_w = w_len
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
return freqs
def forward(
self,
x: torch.Tensor,
timestep: torch.LongTensor,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
B, C, F, H, W = x.shape
freqs = self.rope_encode_comfy(F, H, W, device=x.device, dtype=x.dtype)
self.patch_embedding.to(self.main_device)
x = self.patch_embedding(x.float()).to(x.dtype).flatten(2).transpose(1, 2).to(self.base_dtype)
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
time_embed_dtype = self.time_embedding[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)).to(self.base_dtype) # b, 6, dim
del e
# 4. Transformer blocks
block_samples = ()
for block in self.blocks:
block_samples = block_samples + (x, )
x = block(x, e0, seq_lens, freqs)
block_samples = block_samples + (x, )
ref_block_samples = ()
for block_sample, ref_block in zip(block_samples, self.ref_blocks):
block_sample = ref_block(block_sample)
ref_block_samples = ref_block_samples + (block_sample, )
return ref_block_samples, freqs
+144
View File
@@ -0,0 +1,144 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
import comfy.ops
ops = comfy.ops.disable_weight_init
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float(
"-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class CausalConv3d(nn.Module):
def __init__(self, chan_in, chan_out, kernel_size, stride = 1, dilation = 1, pad_mode='replicate', **kwargs):
super().__init__()
self.pad_mode = pad_mode
padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0) # W, H, T
self.time_causal_padding = padding
self.conv = ops.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class DownsampleCausal3D(nn.Module):
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv", kernel_size=3, bias=True, stride=2):
super().__init__()
self.channels, self.out_channels, self.use_conv, self.padding, self.name = channels, out_channels or channels, use_conv, padding, name
self.conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, bias=bias)
def forward(self, x, scale=1.0):
return self.conv(x)
class ResnetBlockCausal3D(nn.Module):
def __init__(self, *, in_channels: int, out_channels: Optional[int] = None, groups: int = 32, eps: float = 1e-6, conv_3d_out_channels: Optional[int] = None):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.norm2 = torch.nn.GroupNorm(num_groups=groups, num_channels=out_channels, eps=eps, affine=True)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
conv_3d_out_channels = conv_3d_out_channels or out_channels
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
def forward(self, input_tensor: torch.FloatTensor, temb: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
hidden_states = input_tensor
hidden_states = self.conv1(nn.SiLU()(self.norm1(hidden_states)))
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.conv2(nn.SiLU()(self.norm2(hidden_states)))
return input_tensor + hidden_states
def get_down_block3d(down_block_type: str, num_layers: int, in_channels: int, out_channels: int,
add_downsample: bool, downsample_stride: int, resnet_eps: float, resnet_act_fn: str, resnet_groups: Optional[int] = None,
downsample_padding: Optional[int] = None, **kwargs):
down_block_type = down_block_type[7:] if down_block_type.startswith(
"UNetRes") else down_block_type
if down_block_type == "DownEncoderBlockCausal3D":
return DownEncoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
add_downsample=add_downsample,
downsample_stride=downsample_stride,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
downsample_padding=downsample_padding,
)
raise ValueError(f"{down_block_type} does not exist.")
class DownEncoderBlockCausal3D(nn.Module):
def __init__(self, in_channels: int, out_channels: int, num_layers: int = 1, resnet_eps: float = 1e-6,
resnet_groups: int = 32, add_downsample: bool = True, downsample_stride: int = 2, downsample_padding: int = 1, **kwargs):
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=out_channels,
eps=resnet_eps,
groups=resnet_groups,
)
)
self.resnets = nn.ModuleList(resnets)
if add_downsample:
self.downsamplers = nn.ModuleList([DownsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
padding=downsample_padding,
name="op",
stride=downsample_stride,
)])
else:
self.downsamplers = None
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states, scale)
return hidden_states
+2 -9
View File
@@ -118,13 +118,6 @@ class WanRotaryPosEmbed(nn.Module):
from ..wanvideo.modules.attention import sageattn_func
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
class SimpleAttnProcessor2_0:
def __init__(self, attention_mode):
self.attention_mode = attention_mode
@@ -278,7 +271,7 @@ class MaskCamEmbed(nn.Module):
mid_channels = controlnet_cfg.get("mid_channels", 64)
self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)),
nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU())
self.mask_zero_proj = zero_module(nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2)))
self.mask_zero_proj = nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))
def forward(self, add_inputs: torch.Tensor):
# render_mask.shape [b,c,f,h,w]
@@ -321,7 +314,7 @@ class WanControlNet(ModelMixin):
)
self.proj_out = nn.ModuleList(
[
zero_module(nn.Linear(self.dim, 5120))
nn.Linear(self.dim, 5120)
for _ in range(controlnet_cfg["num_layers"])
]
)
+178 -34
View File
@@ -29,6 +29,18 @@ from comfy import model_management as mm
__all__ = ['WanModel']
def apply_rotary_emb_split(hidden_states, freqs_cis, t_dim):
"""Apply rotary embedding only to the spatial (H/W) dimensions, leaving temporal (T) unchanged."""
t_part, hw_part = torch.split(hidden_states, [t_dim, hidden_states.shape[-1] - t_dim], dim=-1)
hw_freqs = freqs_cis[..., t_dim//2:, :, :]
x_ = hw_part.to(dtype=hw_freqs.dtype).reshape(*hw_part.shape[:-1], -1, 1, 2)
x_out = hw_freqs[..., 0] * x_[..., 0]
x_out.addcmul_(hw_freqs[..., 1], x_[..., 1])
out_hw = x_out.reshape(*hw_part.shape).type_as(hidden_states)
return torch.cat([t_part, out_hw], dim=-1)
class AdaLayerNorm(nn.Module):
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5):
super().__init__()
@@ -100,11 +112,6 @@ class FramePackMotioner(nn.Module):#from comfy.ldm.wan.model
rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1)
return motion_lat, rope
def zero_module(module):
for p in module.parameters():
p.detach().zero_()
return module
def torch_dfs(model: nn.Module, parent_name='root'):
module_names, modules = [], []
current_name = parent_name if parent_name else 'root'
@@ -458,7 +465,7 @@ class WanSelfAttention(nn.Module):
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
return q, k, v
def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None):
def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None, onetoall_ref=None, onetoall_ref_scale=1.0):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -478,6 +485,9 @@ class WanSelfAttention(nn.Module):
if self.ref_adapter is not None and lynx_ref_feature is not None:
x = x.add(ref_x, alpha=lynx_ref_scale)
if onetoall_ref is not None:
x = x.add(onetoall_ref, alpha=onetoall_ref_scale)
# output
return self.o(x.flatten(2))
@@ -883,6 +893,8 @@ class WanAttentionBlock(nn.Module):
self.kv_cache = None
self.use_motion_attn = use_motion_attn
self.has_face_fuser_block = face_fuser_block
self.ref_attn_k_img = None
self.ref_attn_v_img = None
# layers
self.norm1 = WanLayerNorm(self.dim, eps)
@@ -996,6 +1008,8 @@ class WanAttentionBlock(nn.Module):
lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx
x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None,
num_cond_latents=None, #longcat image cond amount
x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all
e_tr=None, tr_num=0, tr_start=0, #token replacement
):
r"""
Args:
@@ -1012,6 +1026,13 @@ class WanAttentionBlock(nn.Module):
self.seg_idx = [0, self.seg_idx, x.size(1)]
e = e[0]
use_token_replace = False
if e_tr is not None and tr_num > 0:
tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = self.get_mod(e_tr.to(x.device), self.modulation)
use_token_replace = True
tr_start = tr_start or 0
tr_end = tr_start + (tr_num or 0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
del e
input_dtype = x.dtype
@@ -1020,6 +1041,13 @@ class WanAttentionBlock(nn.Module):
is_longcat = C == 4096
if is_longcat:
input_x = self.modulate(self.norm1(x.view(B, T, -1, C).to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype).view(B, N, C)
elif use_token_replace:
norm_x = self.norm1(x.to(shift_msa.dtype))
input_x = torch.cat([
torch.addcmul(shift_msa, norm_x[:, :tr_start], 1 + scale_msa), # before replace → T
torch.addcmul(tr_shift_msa, norm_x[:, tr_start:tr_end], 1 + tr_scale_msa), # replace segment → t=0
torch.addcmul(shift_msa, norm_x[:, tr_end:], 1 + scale_msa) # after replace → T
], dim=1).to(input_dtype)
else:
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype)
@@ -1053,6 +1081,23 @@ class WanAttentionBlock(nn.Module):
if lynx_ref_feature is None and self.self_attn.ref_adapter is not None:
lynx_ref_feature = input_x
onetoall_ref = None
if x_onetoall_ref is not None:
b, s, n, d = *x_onetoall_ref.shape[:2], self.self_attn.num_heads, self.self_attn.head_dim
h_dim = w_dim = 2 * (self.head_dim // 6)
t_dim = self.head_dim - h_dim - w_dim
q_ref = self.self_attn.norm_q(self.self_attn.q(input_x)).to(input_x.dtype).view(b, N, n, d)
q_ref = apply_rotary_emb_split(q_ref, freqs, t_dim) # Apply split rotary embedding (only to H/W dimensions, leaving T unchanged)
k_ref = self.ref_attn_norm_k_img(self.ref_attn_k_img(x_onetoall_ref).to(self.ref_attn_norm_k_img.weight.dtype)).to(x_onetoall_ref.dtype).view(b, s, n, d)
k_ref = apply_rotary_emb_split(k_ref, onetoall_freqs, t_dim)
v_ref = self.ref_attn_v_img(x_onetoall_ref).view(b, s, n, d)
onetoall_ref = attention(q_ref, k_ref, v_ref, k_lens=seq_lens, attention_mode=self.attention_mode)
del q_ref, k_ref, v_ref
#RoPE and QKV computation
if inner_t is not None:
#query, key, value
@@ -1153,9 +1198,9 @@ class WanAttentionBlock(nn.Module):
# merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale)
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale, onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale)
del q, k, v
del q, k, v,
# FETA
if enhance_enabled:
@@ -1180,10 +1225,16 @@ class WanAttentionBlock(nn.Module):
y = torch.cat(z, dim=1)
x = x.add(y)
else:
if not is_longcat:
x = x.addcmul(y, gate_msa)
else:
if is_longcat:
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C)
elif use_token_replace:
x = x + torch.cat([
y[:, :tr_start] * gate_msa,
y[:, tr_start:tr_end] * tr_gate_msa,
y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype)
else:
x = x.addcmul(y, gate_msa)
del y, gate_msa
# cross-attention & ffn function
@@ -1248,10 +1299,18 @@ class WanAttentionBlock(nn.Module):
norm2_x = torch.cat(parts, dim=1)
x_ffn = self.ffn(norm2_x)
else:
if not is_longcat:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
else:
if is_longcat:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
elif use_token_replace:
norm2_x = self.norm2(x.to(shift_mlp.dtype))
mod_x = torch.cat([
torch.addcmul(shift_mlp, norm2_x[:, :tr_start], 1 + scale_mlp),
torch.addcmul(tr_shift_mlp, norm2_x[:, tr_start:tr_end], 1 + tr_scale_mlp),
torch.addcmul(shift_mlp, norm2_x[:, tr_end:], 1 + scale_mlp)
], dim=1)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
del shift_mlp, scale_mlp
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
del mod_x
@@ -1264,10 +1323,16 @@ class WanAttentionBlock(nn.Module):
x_ffn = torch.cat(z, dim=1)
x = x.add(x_ffn)
else:
if not is_longcat:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
else:
if is_longcat:
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
elif use_token_replace:
x = x + torch.cat([
x_ffn[:, :tr_start] * gate_mlp,
x_ffn[:, tr_start:tr_end] * tr_gate_mlp,
x_ffn[:, tr_end:] * gate_mlp
], dim=1).to(input_dtype)
else:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
del gate_mlp
if x_ip is not None: #stand-in
@@ -1419,14 +1484,23 @@ class Head(nn.Module):
e = (self.modulation.unsqueeze(2) + e.unsqueeze(1)).chunk(2, dim=1)
return [ei.squeeze(1) for ei in e]
def forward(self, x, e, **kwargs):
def forward(self, x, e, e_tr=None, tr_start=0, tr_num=0, **kwargs):
r"""
Args:
x(Tensor): Shape [B, L1, C]
e(Tensor): Shape [B, C]
"""
e = self.get_mod(e.to(x.device))
if tr_num > 0 and e_tr is not None:
e_tr = self.get_mod(e_tr.to(x.device))
tr_end = tr_start + tr_num
norm_x = self.norm(x.float()).to(x.dtype)
x = self.head(torch.cat([
norm_x[:, :tr_start].mul(1 + e[1]).add(e[0]),
norm_x[:, tr_start:tr_end].mul(1 + e_tr[1]).add(e_tr[0]),
norm_x[:, tr_end:].mul(1 + e[1]).add(e[0])
], dim=1))
else:
x = self.head(self.norm(x.float()).to(x.dtype).mul_(1 + e[1]).add_(e[0]))
return x
@@ -1803,13 +1877,7 @@ class WanModel(torch.nn.Module):
self.head = Head_adaLN(dim, out_dim, patch_size, eps, adaln_tembed_dim=512)
d = self.dim // self.num_heads
self.rope_embedder = EmbedND_RifleX(
d,
10000.0,
[d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)],
num_frames=None,
k=None,
)
self.rope_embedder = EmbedND_RifleX(d, 10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=None, k=None)
self.cached_freqs = self.cached_shape = self.cached_cond = None
# buffers (don't use register_buffer otherwise dtype will be changed in to())
@@ -2149,7 +2217,8 @@ class WanModel(torch.nn.Module):
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
num_cond_latents=None,
add_text_emb=None,
sdancer_input=None # SteadyDancer
sdancer_input=None, # SteadyDancer
one_to_all_input=None, # One-to-All
):
r"""
Forward pass through the diffusion model
@@ -2251,6 +2320,40 @@ class WanModel(torch.nn.Module):
y[0].add_(random_ref_emb, alpha=unianim_data["strength"])
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
suffix_frames = x[0].shape[1]
prefix_frames = 0
# One-to-all-Animation
onetoall_ref_block_samples = onetoall_freqs = prev_x = prev_control = None
onetoall_ref_scale = 1.0
onetoall_control_enabled = use_token_replace = False
e0_token_replace = token_replace_start = None
replace_token_num = token_replace_start = 0
if one_to_all_input is not None:
# reference condition
ref_cond_latent = one_to_all_input.get("ref_latent_pos", None) if not is_uncond else one_to_all_input.get("ref_latent_neg", None)
if ref_cond_latent is not None and one_to_all_input['ref_start_percent'] <= current_step_percentage <= one_to_all_input['ref_end_percent']:
onetoall_ref_scale = one_to_all_input.get("ref_strength", 1.0)
image_cond = self.image_to_cond(ref_cond_latent.to(self.main_device, self.base_dtype))[0]
x = [torch.cat([v, u], dim=1) for v, u in zip([image_cond], x)]
seq_len += math.ceil((image_cond.shape[-1] * image_cond.shape[-2]) / 4 * image_cond.shape[-3])
F += 1
prefix_frames = 1
suffix_frames += 1
onetoall_ref_block_samples, onetoall_freqs = self.refextractor(ref_cond_latent, timestep=t)
# pose controlnet
controlnet_tokens = one_to_all_input.get("controlnet_tokens", None)
if not is_uncond and controlnet_tokens is not None and one_to_all_input['controlnet_start_percent'] <= current_step_percentage <= one_to_all_input['controlnet_end_percent']:
onetoall_control_enabled = True
onetoall_control_strength = one_to_all_input.get("controlnet_strength", 1.0)
# token replace
if one_to_all_input.get("token_replace", False):
use_token_replace = True
num_latent_frames_to_replace = one_to_all_input.get("num_latent_frames_to_replace", 2)
t_token_replace = torch.zeros_like(t)
token_replace_start = (H // self.patch_size[1]) * (W // self.patch_size[2]) # skip first (ref) frame
replace_token_num = num_latent_frames_to_replace * token_replace_start # zero next frames
#uni3c controlnet
if uni3c_data is not None:
render_latent = uni3c_data["render_latent"].to(self.base_dtype)
@@ -2279,8 +2382,6 @@ class WanModel(torch.nn.Module):
self.original_patch_embedding.to(self.main_device)
x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x]
orig_frames = x[0].shape[1]
# ovi audio model
if self.audio_model is not None:
x_ovi = [self.audio_model.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x_ovi[0].dtype) for u in x_ovi]
@@ -2463,6 +2564,9 @@ class WanModel(torch.nn.Module):
time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim
if use_token_replace:
e_token_replace = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t_token_replace.flatten()).to(time_embed_dtype)) # b, dim
e0_token_replace = self.time_projection(e_token_replace).unflatten(1, (6, self.dim)) # b, 6, dim
else:
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
@@ -2803,7 +2907,11 @@ class WanModel(torch.nn.Module):
lynx_x_ip=lynx_x_ip,
lynx_ip_scale=lynx_ip_scale,
lynx_ref_scale=lynx_ref_scale,
num_cond_latents=num_cond_latents
num_cond_latents=num_cond_latents,
onetoall_ref_scale=onetoall_ref_scale,
e_tr=e0_token_replace if use_token_replace else None,
tr_start=token_replace_start,
tr_num=replace_token_num,
)
if self.audio_model is not None:
kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)
@@ -2898,7 +3006,17 @@ class WanModel(torch.nn.Module):
if b in self.slg_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
continue
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, **kwargs) #run block
x_onetoall_ref = None
if onetoall_ref_block_samples is not None:
interval_ref = len(self.blocks) / len(onetoall_ref_block_samples)
interval_ref = int(np.ceil(interval_ref))
x_onetoall_ref = onetoall_ref_block_samples[b // interval_ref]
# ---run block----#
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_onetoall_ref=x_onetoall_ref, onetoall_freqs=onetoall_freqs, **kwargs)
# ---post block----#
if self.audio_injector is not None and s2v_audio_input is not None:
x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v
if block.has_face_fuser_block and motion_vec is not None:
@@ -2924,6 +3042,31 @@ class WanModel(torch.nn.Module):
#controlnet
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
x[:, :self.original_seq_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
# One-to-All-Animation controlnet
if onetoall_control_enabled:
if prev_x is not None and (b - 1) < len(self.controlnet.blocks):
tqdm.write(f"Applying One-to-All ControlNet at block {b}")
if b == 1:
ctrl_in = prev_x + controlnet_tokens
elif prev_control is not None:
ctrl_in = prev_control
self.controlnet.blocks[b - 1].to(self.main_device)
control_out = self.controlnet.blocks[b - 1](ctrl_in, e0, seq_lens, freqs, e_tr=e0_token_replace, tr_num=replace_token_num,tr_start=token_replace_start, split_rope=False)
self.controlnet.blocks[b - 1].to(self.offload_device, non_blocking=self.use_non_blocking)
prev_control = control_out
control_out_proj = self.controlnet_zero[b - 1](control_out)
x = x + control_out_proj * onetoall_control_strength
if b < len(self.controlnet.blocks): # Store prev_x only while controlnet is active
prev_x = x
elif b == len(self.controlnet.blocks): # Controlnet done, free memory
prev_x = None
prev_control = None
if controlnet_tokens is not None:
del controlnet_tokens
controlnet_tokens = None
mm.soft_empty_cache()
if lynx_ref_feature_extractor:
return lynx_ref_buffer
@@ -2979,7 +3122,8 @@ class WanModel(torch.nn.Module):
x = x[:, :self.original_seq_len]
x = self.head(x, e.to(x.device), temp_length=F)
x = self.head(x, e.to(x.device), temp_length=F,
e_tr=e_token_replace.to(x.device) if use_token_replace else None, tr_start=token_replace_start, tr_num=replace_token_num)
if x_ovi is not None:
x_ovi = self.audio_model.head(x_ovi, e_ovi.to(x_ovi.device))
@@ -2988,8 +3132,8 @@ class WanModel(torch.nn.Module):
x_ovi = [u[:gs] for u, gs in zip(x_ovi, grid_sizes_ovi)]
x_ovi = [u.float() for u in x_ovi]
x = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type]
x = [u[:, :orig_frames, ...].float() for u in x]
x = self.unpatchify(x, original_grid_sizes)
x = [u[:, prefix_frames:suffix_frames, ...].float() for u in x]
return (x, x_ovi, pred_id) if pred_id is not None else (x, x_ovi, None)
def unpatchify(self, x, grid_sizes):