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: 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
This commit is contained in:
+10
@@ -84,6 +84,13 @@ except Exception as e:
|
|||||||
STEADYDANCER_NODE_CLASS_MAPPINGS = {}
|
STEADYDANCER_NODE_CLASS_MAPPINGS = {}
|
||||||
STEADYDANCER_NODE_DISPLAY_NAME_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(RECAM_MASTER_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(SKYREELS_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(FLASHVSR_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
|
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
|
||||||
NODE_CLASS_MAPPINGS.update(STEADYDANCER_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(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_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(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
|
||||||
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_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(STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS.update(ONETOALL_NODE_DISPLAY_NAME_MAPPINGS)
|
||||||
|
|
||||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
@@ -10,18 +10,13 @@ from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscal
|
|||||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||||
from diffusers.models.modeling_utils import ModelMixin
|
from diffusers.models.modeling_utils import ModelMixin
|
||||||
from diffusers.models.transformers.transformer_wan import (
|
from diffusers.models.transformers.transformer_wan import (
|
||||||
WanTimeTextImageEmbedding,
|
WanTimeTextImageEmbedding,
|
||||||
WanRotaryPosEmbed,
|
WanRotaryPosEmbed,
|
||||||
WanTransformerBlock
|
WanTransformerBlock
|
||||||
)
|
)
|
||||||
|
|
||||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
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):
|
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
|
||||||
r"""
|
r"""
|
||||||
@@ -69,7 +64,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
|||||||
_no_split_modules = ["WanTransformerBlock"]
|
_no_split_modules = ["WanTransformerBlock"]
|
||||||
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
|
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
|
||||||
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
|
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
|
||||||
|
|
||||||
@register_to_config
|
@register_to_config
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -100,10 +95,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
|||||||
## Spatial compression with time awareness
|
## Spatial compression with time awareness
|
||||||
nn.Sequential(
|
nn.Sequential(
|
||||||
nn.Conv3d(
|
nn.Conv3d(
|
||||||
in_channels,
|
in_channels,
|
||||||
input_channels[0],
|
input_channels[0],
|
||||||
kernel_size=(3, downscale_coef + 1, downscale_coef + 1),
|
kernel_size=(3, downscale_coef + 1, downscale_coef + 1),
|
||||||
stride=(1, downscale_coef, downscale_coef),
|
stride=(1, downscale_coef, downscale_coef),
|
||||||
padding=(1, downscale_coef // 2, downscale_coef // 2)
|
padding=(1, downscale_coef // 2, downscale_coef // 2)
|
||||||
),
|
),
|
||||||
nn.GELU(approximate="tanh"),
|
nn.GELU(approximate="tanh"),
|
||||||
@@ -122,9 +117,9 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
|||||||
nn.GroupNorm(2, input_channels[2]),
|
nn.GroupNorm(2, input_channels[2]),
|
||||||
)
|
)
|
||||||
])
|
])
|
||||||
|
|
||||||
inner_dim = num_attention_heads * attention_head_dim
|
inner_dim = num_attention_heads * attention_head_dim
|
||||||
|
|
||||||
# 1. Patch & position embedding
|
# 1. Patch & position embedding
|
||||||
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
|
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
|
||||||
self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size)
|
self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size)
|
||||||
@@ -153,11 +148,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
|||||||
|
|
||||||
for _ in range(len(self.blocks)):
|
for _ in range(len(self.blocks)):
|
||||||
controlnet_block = nn.Linear(inner_dim, out_proj_dim)
|
controlnet_block = nn.Linear(inner_dim, out_proj_dim)
|
||||||
controlnet_block = zero_module(controlnet_block)
|
|
||||||
self.controlnet_blocks.append(controlnet_block)
|
self.controlnet_blocks.append(controlnet_block)
|
||||||
|
|
||||||
self.gradient_checkpointing = False
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
@@ -187,7 +181,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
|||||||
# 0. Controlnet encoder
|
# 0. Controlnet encoder
|
||||||
for control_encoder_block in self.control_encoder:
|
for control_encoder_block in self.control_encoder:
|
||||||
controlnet_states = control_encoder_block(controlnet_states)
|
controlnet_states = control_encoder_block(controlnet_states)
|
||||||
|
|
||||||
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
|
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
|
||||||
|
|
||||||
## 1. Patch embedding and stack
|
## 1. Patch embedding and stack
|
||||||
@@ -216,7 +210,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
|||||||
|
|
||||||
if encoder_hidden_states_image is not None:
|
if encoder_hidden_states_image is not None:
|
||||||
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||||
|
|
||||||
# 4. Transformer blocks
|
# 4. Transformer blocks
|
||||||
controlnet_hidden_states = ()
|
controlnet_hidden_states = ()
|
||||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||||
@@ -239,43 +233,4 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
|
|||||||
return (controlnet_hidden_states,)
|
return (controlnet_hidden_states,)
|
||||||
|
|
||||||
return Transformer2DModelOutput(sample=controlnet_hidden_states)
|
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
@@ -6,7 +6,7 @@ import numpy as np
|
|||||||
from tqdm import tqdm
|
from tqdm import tqdm
|
||||||
import re
|
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.t5 import T5EncoderModel
|
||||||
from .wanvideo.modules.clip import CLIPModel
|
from .wanvideo.modules.clip import CLIPModel
|
||||||
from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
|
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,
|
total=param_count,
|
||||||
leave=True):
|
leave=True):
|
||||||
block_idx = vace_block_idx = None
|
block_idx = vace_block_idx = None
|
||||||
if "vace_blocks." in name:
|
if name.startswith("vace_blocks."):
|
||||||
try:
|
try:
|
||||||
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
||||||
except Exception:
|
except Exception:
|
||||||
vace_block_idx = None
|
vace_block_idx = None
|
||||||
elif "blocks." in name and "face" not in name:
|
elif name.startswith("blocks.") and "face" not in name:
|
||||||
try:
|
try:
|
||||||
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
||||||
except Exception:
|
except Exception:
|
||||||
block_idx = None
|
block_idx = None
|
||||||
|
|
||||||
if "loras" in name or "controlnet" in name:
|
if "loras" in name:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
# GGUF: skip GGUFParameter params
|
# GGUF: skip GGUFParameter params
|
||||||
@@ -1310,7 +1310,7 @@ class WanVideoModelLoader:
|
|||||||
if dim == 1536:
|
if dim == 1536:
|
||||||
model_variant = "1_3B"
|
model_variant = "1_3B"
|
||||||
if dim == 3072:
|
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 "high" in model.lower() or "low" in model.lower():
|
||||||
if "i2v" in model.lower():
|
if "i2v" in model.lower():
|
||||||
@@ -1375,7 +1375,7 @@ class WanVideoModelLoader:
|
|||||||
with init_empty_weights():
|
with init_empty_weights():
|
||||||
transformer.audio_model = WanModel(**TRANSFORMER_CONFIG).eval()
|
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:
|
for block in transformer.blocks:
|
||||||
block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim)
|
block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim)
|
||||||
@@ -1477,16 +1477,12 @@ class WanVideoModelLoader:
|
|||||||
|
|
||||||
# Additional cond latents
|
# Additional cond latents
|
||||||
if "add_conv_in.weight" in sd:
|
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]
|
inner_dim = sd["add_conv_in.weight"].shape[0]
|
||||||
add_cond_in_dim = sd["add_conv_in.weight"].shape[1]
|
add_cond_in_dim = sd["add_conv_in.weight"].shape[1]
|
||||||
attn_cond_in_dim = sd["attn_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_conv_in = 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.add_proj = 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.attn_conv_in = nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
|
||||||
|
|
||||||
# Bindweave text_projection
|
# Bindweave text_projection
|
||||||
if "text_projection.0.weight" in sd:
|
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
|
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.diffusion_model = transformer
|
||||||
comfy_model.load_device = transformer_load_device
|
comfy_model.load_device = transformer_load_device
|
||||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||||
@@ -1580,7 +1615,7 @@ class WanVideoModelLoader:
|
|||||||
if gguf:
|
if gguf:
|
||||||
raise ValueError("GGUF models don't support vram management")
|
raise ValueError("GGUF models don't support vram management")
|
||||||
from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
|
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())
|
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}")
|
log.info(f"Total number of parameters in the loaded model: {total_params_in_model}")
|
||||||
|
|||||||
+38
-5
@@ -1182,6 +1182,31 @@ class WanVideoSampler:
|
|||||||
sdancer_data = sdancer_embeds.copy()
|
sdancer_data = sdancer_embeds.copy()
|
||||||
sdancer_data = dict_to_device(sdancer_data, device, dtype)
|
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
|
#region model pred
|
||||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
|
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,
|
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
|
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
|
||||||
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
|
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
|
||||||
"sdancer_input": sdancer_input, # SteadyDancer input
|
"sdancer_input": sdancer_input, # SteadyDancer input
|
||||||
|
"one_to_all_input": one_to_all_data, # One-to-All input
|
||||||
}
|
}
|
||||||
|
|
||||||
batch_size = 1
|
batch_size = 1
|
||||||
@@ -3070,11 +3096,16 @@ class WanVideoSampler:
|
|||||||
new_latent.append(latent_slice[:, j:j+1])
|
new_latent.append(latent_slice[:, j:j+1])
|
||||||
latent = torch.cat(new_latent, dim=1)
|
latent = torch.cat(new_latent, dim=1)
|
||||||
else:
|
else:
|
||||||
latent = sample_scheduler.step(
|
if latents_to_not_step > 0:
|
||||||
noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else noise_pred.unsqueeze(0),
|
raw_latent = latent[:, :latents_to_not_step]
|
||||||
timestep,
|
noise_pred_in = noise_pred[:, latents_to_not_step:]
|
||||||
latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else latent.unsqueeze(0),
|
latent = latent[:, latents_to_not_step:]
|
||||||
**scheduler_step_args)[0].squeeze(0)
|
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:
|
if noise_pred_flipped is not None:
|
||||||
latent_backwards = sample_scheduler_flipped.step(
|
latent_backwards = sample_scheduler_flipped.step(
|
||||||
noise_pred_flipped.unsqueeze(0),
|
noise_pred_flipped.unsqueeze(0),
|
||||||
@@ -3083,6 +3114,8 @@ class WanVideoSampler:
|
|||||||
**scheduler_step_args)[0].squeeze(0)
|
**scheduler_step_args)[0].squeeze(0)
|
||||||
latent_backwards = torch.flip(latent_backwards, dims=[1])
|
latent_backwards = torch.flip(latent_backwards, dims=[1])
|
||||||
latent = latent * 0.5 + latent_backwards * 0.5
|
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:
|
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)
|
latent_ovi = sample_scheduler_ovi.step(noise_pred_ovi.unsqueeze(0), t, latent_ovi.to(device).unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -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",
|
||||||
|
}
|
||||||
@@ -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
|
||||||
@@ -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
|
||||||
+3
-10
@@ -115,16 +115,9 @@ class WanRotaryPosEmbed(nn.Module):
|
|||||||
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
|
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
|
||||||
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
|
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
|
||||||
return freqs
|
return freqs
|
||||||
|
|
||||||
from ..wanvideo.modules.attention import sageattn_func
|
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:
|
class SimpleAttnProcessor2_0:
|
||||||
def __init__(self, attention_mode):
|
def __init__(self, attention_mode):
|
||||||
self.attention_mode = attention_mode
|
self.attention_mode = attention_mode
|
||||||
@@ -278,7 +271,7 @@ class MaskCamEmbed(nn.Module):
|
|||||||
mid_channels = controlnet_cfg.get("mid_channels", 64)
|
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)),
|
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())
|
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):
|
def forward(self, add_inputs: torch.Tensor):
|
||||||
# render_mask.shape [b,c,f,h,w]
|
# render_mask.shape [b,c,f,h,w]
|
||||||
@@ -321,7 +314,7 @@ class WanControlNet(ModelMixin):
|
|||||||
)
|
)
|
||||||
self.proj_out = nn.ModuleList(
|
self.proj_out = nn.ModuleList(
|
||||||
[
|
[
|
||||||
zero_module(nn.Linear(self.dim, 5120))
|
nn.Linear(self.dim, 5120)
|
||||||
for _ in range(controlnet_cfg["num_layers"])
|
for _ in range(controlnet_cfg["num_layers"])
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
|
|||||||
+224
-80
@@ -29,6 +29,18 @@ from comfy import model_management as mm
|
|||||||
|
|
||||||
__all__ = ['WanModel']
|
__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):
|
class AdaLayerNorm(nn.Module):
|
||||||
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5):
|
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5):
|
||||||
super().__init__()
|
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)
|
rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1)
|
||||||
return motion_lat, rope
|
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'):
|
def torch_dfs(model: nn.Module, parent_name='root'):
|
||||||
module_names, modules = [], []
|
module_names, modules = [], []
|
||||||
current_name = parent_name if parent_name else 'root'
|
current_name = parent_name if parent_name else 'root'
|
||||||
@@ -450,7 +457,7 @@ class WanSelfAttention(nn.Module):
|
|||||||
def qkv_fn_v(self, x):
|
def qkv_fn_v(self, x):
|
||||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||||
return self.v(x).view(b, s, n, d)
|
return self.v(x).view(b, s, n, d)
|
||||||
|
|
||||||
def qkv_fn_ip(self, x):
|
def qkv_fn_ip(self, x):
|
||||||
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
|
||||||
q = self.norm_q(self.q(x) + self.q_loras(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
q = self.norm_q(self.q(x) + self.q_loras(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d)
|
||||||
@@ -458,7 +465,7 @@ class WanSelfAttention(nn.Module):
|
|||||||
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
|
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
|
||||||
return q, k, v
|
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"""
|
r"""
|
||||||
Args:
|
Args:
|
||||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||||
@@ -478,9 +485,12 @@ class WanSelfAttention(nn.Module):
|
|||||||
if self.ref_adapter is not None and lynx_ref_feature is not None:
|
if self.ref_adapter is not None and lynx_ref_feature is not None:
|
||||||
x = x.add(ref_x, alpha=lynx_ref_scale)
|
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
|
# output
|
||||||
return self.o(x.flatten(2))
|
return self.o(x.flatten(2))
|
||||||
|
|
||||||
def forward_ip(self, q, k, v, q_ip, k_ip, v_ip, seq_lens, attention_mode_override=None):
|
def forward_ip(self, q, k, v, q_ip, k_ip, v_ip, seq_lens, attention_mode_override=None):
|
||||||
attention_mode = self.attention_mode
|
attention_mode = self.attention_mode
|
||||||
if attention_mode_override is not None:
|
if attention_mode_override is not None:
|
||||||
@@ -883,6 +893,8 @@ class WanAttentionBlock(nn.Module):
|
|||||||
self.kv_cache = None
|
self.kv_cache = None
|
||||||
self.use_motion_attn = use_motion_attn
|
self.use_motion_attn = use_motion_attn
|
||||||
self.has_face_fuser_block = face_fuser_block
|
self.has_face_fuser_block = face_fuser_block
|
||||||
|
self.ref_attn_k_img = None
|
||||||
|
self.ref_attn_v_img = None
|
||||||
|
|
||||||
# layers
|
# layers
|
||||||
self.norm1 = WanLayerNorm(self.dim, eps)
|
self.norm1 = WanLayerNorm(self.dim, eps)
|
||||||
@@ -965,7 +977,7 @@ class WanAttentionBlock(nn.Module):
|
|||||||
return norm_x
|
return norm_x
|
||||||
else:
|
else:
|
||||||
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
|
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
|
||||||
|
|
||||||
def ffn_chunked(self, mod_x, num_chunks=4):
|
def ffn_chunked(self, mod_x, num_chunks=4):
|
||||||
seq_len = mod_x.shape[1]
|
seq_len = mod_x.shape[1]
|
||||||
if seq_len <= 8192 or num_chunks <= 1:
|
if seq_len <= 8192 or num_chunks <= 1:
|
||||||
@@ -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
|
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,
|
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
|
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"""
|
r"""
|
||||||
Args:
|
Args:
|
||||||
@@ -1012,6 +1026,13 @@ class WanAttentionBlock(nn.Module):
|
|||||||
self.seg_idx = [0, self.seg_idx, x.size(1)]
|
self.seg_idx = [0, self.seg_idx, x.size(1)]
|
||||||
e = e[0]
|
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)
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
|
||||||
del e
|
del e
|
||||||
input_dtype = x.dtype
|
input_dtype = x.dtype
|
||||||
@@ -1020,6 +1041,13 @@ class WanAttentionBlock(nn.Module):
|
|||||||
is_longcat = C == 4096
|
is_longcat = C == 4096
|
||||||
if is_longcat:
|
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)
|
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:
|
else:
|
||||||
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype)
|
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:
|
if lynx_ref_feature is None and self.self_attn.ref_adapter is not None:
|
||||||
lynx_ref_feature = input_x
|
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
|
#RoPE and QKV computation
|
||||||
if inner_t is not None:
|
if inner_t is not None:
|
||||||
#query, key, value
|
#query, key, value
|
||||||
@@ -1104,10 +1149,10 @@ class WanAttentionBlock(nn.Module):
|
|||||||
# FETA
|
# FETA
|
||||||
if enhance_enabled:
|
if enhance_enabled:
|
||||||
feta_scores = get_feta_scores(q, k)
|
feta_scores = get_feta_scores(q, k)
|
||||||
|
|
||||||
#self-attention
|
#self-attention
|
||||||
split_attn = (context is not None
|
split_attn = (context is not None
|
||||||
and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1))
|
and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1))
|
||||||
and x.shape[0] == 1
|
and x.shape[0] == 1
|
||||||
and inner_t is None
|
and inner_t is None
|
||||||
and x_ip is None # Don't split when using IP-Adapter
|
and x_ip is None # Don't split when using IP-Adapter
|
||||||
@@ -1144,18 +1189,18 @@ class WanAttentionBlock(nn.Module):
|
|||||||
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
|
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
|
||||||
# process the condition tokens
|
# process the condition tokens
|
||||||
x_cond = self.self_attn.forward(
|
x_cond = self.self_attn.forward(
|
||||||
q[:, :num_cond_latents_thw].contiguous(),
|
q[:, :num_cond_latents_thw].contiguous(),
|
||||||
k[:, :num_cond_latents_thw].contiguous(),
|
k[:, :num_cond_latents_thw].contiguous(),
|
||||||
v[:, :num_cond_latents_thw].contiguous(),
|
v[:, :num_cond_latents_thw].contiguous(),
|
||||||
seq_lens)
|
seq_lens)
|
||||||
# process the noise tokens
|
# process the noise tokens
|
||||||
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
|
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
|
||||||
# merge x_cond and x_noise
|
# merge x_cond and x_noise
|
||||||
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
|
y = torch.cat([x_cond, x_noise], dim=1).contiguous()
|
||||||
else:
|
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
|
# FETA
|
||||||
if enhance_enabled:
|
if enhance_enabled:
|
||||||
@@ -1173,17 +1218,23 @@ class WanAttentionBlock(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# S2V
|
# S2V
|
||||||
if zero_timestep:
|
if zero_timestep:
|
||||||
z = []
|
z = []
|
||||||
for i in range(2):
|
for i in range(2):
|
||||||
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1])
|
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1])
|
||||||
y = torch.cat(z, dim=1)
|
y = torch.cat(z, dim=1)
|
||||||
x = x.add(y)
|
x = x.add(y)
|
||||||
else:
|
else:
|
||||||
if not is_longcat:
|
if is_longcat:
|
||||||
x = x.addcmul(y, gate_msa)
|
|
||||||
else:
|
|
||||||
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C)
|
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
|
del y, gate_msa
|
||||||
|
|
||||||
# cross-attention & ffn function
|
# cross-attention & ffn function
|
||||||
@@ -1225,7 +1276,7 @@ class WanAttentionBlock(nn.Module):
|
|||||||
x = x.add(x_audio, alpha=audio_scale)
|
x = x.add(x_audio, alpha=audio_scale)
|
||||||
|
|
||||||
# MTV-Crafter Motion Attention
|
# MTV-Crafter Motion Attention
|
||||||
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
|
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
|
||||||
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
|
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
|
||||||
x = x.add(x_motion, alpha=mtv_strength)
|
x = x.add(x_motion, alpha=mtv_strength)
|
||||||
|
|
||||||
@@ -1248,14 +1299,22 @@ class WanAttentionBlock(nn.Module):
|
|||||||
norm2_x = torch.cat(parts, dim=1)
|
norm2_x = torch.cat(parts, dim=1)
|
||||||
x_ffn = self.ffn(norm2_x)
|
x_ffn = self.ffn(norm2_x)
|
||||||
else:
|
else:
|
||||||
if not is_longcat:
|
if is_longcat:
|
||||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
|
|
||||||
else:
|
|
||||||
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
|
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
|
del shift_mlp, scale_mlp
|
||||||
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
|
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
|
||||||
del mod_x
|
del mod_x
|
||||||
|
|
||||||
# gate_mlp
|
# gate_mlp
|
||||||
if zero_timestep:
|
if zero_timestep:
|
||||||
z = []
|
z = []
|
||||||
@@ -1264,10 +1323,16 @@ class WanAttentionBlock(nn.Module):
|
|||||||
x_ffn = torch.cat(z, dim=1)
|
x_ffn = torch.cat(z, dim=1)
|
||||||
x = x.add(x_ffn)
|
x = x.add(x_ffn)
|
||||||
else:
|
else:
|
||||||
if not is_longcat:
|
if is_longcat:
|
||||||
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
|
|
||||||
else:
|
|
||||||
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
|
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
|
del gate_mlp
|
||||||
|
|
||||||
if x_ip is not None: #stand-in
|
if x_ip is not None: #stand-in
|
||||||
@@ -1389,7 +1454,7 @@ class BaseWanAttentionBlock(WanAttentionBlock):
|
|||||||
x, x_ip, lynx_ref_feature, x_ovi = super().forward(x, **kwargs)
|
x, x_ip, lynx_ref_feature, x_ovi = super().forward(x, **kwargs)
|
||||||
if vace_hints is None:
|
if vace_hints is None:
|
||||||
return x, x_ip, lynx_ref_feature, x_ovi
|
return x, x_ip, lynx_ref_feature, x_ovi
|
||||||
|
|
||||||
if self.block_id is not None:
|
if self.block_id is not None:
|
||||||
for i in range(len(vace_hints)):
|
for i in range(len(vace_hints)):
|
||||||
x.add_(vace_hints[i][self.block_id].to(x.device), alpha=vace_context_scale[i])
|
x.add_(vace_hints[i][self.block_id].to(x.device), alpha=vace_context_scale[i])
|
||||||
@@ -1419,17 +1484,26 @@ class Head(nn.Module):
|
|||||||
e = (self.modulation.unsqueeze(2) + e.unsqueeze(1)).chunk(2, dim=1)
|
e = (self.modulation.unsqueeze(2) + e.unsqueeze(1)).chunk(2, dim=1)
|
||||||
return [ei.squeeze(1) for ei in e]
|
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"""
|
r"""
|
||||||
Args:
|
Args:
|
||||||
x(Tensor): Shape [B, L1, C]
|
x(Tensor): Shape [B, L1, C]
|
||||||
e(Tensor): Shape [B, C]
|
e(Tensor): Shape [B, C]
|
||||||
"""
|
"""
|
||||||
|
|
||||||
e = self.get_mod(e.to(x.device))
|
e = self.get_mod(e.to(x.device))
|
||||||
x = self.head(self.norm(x.float()).to(x.dtype).mul_(1 + e[1]).add_(e[0]))
|
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
|
return x
|
||||||
|
|
||||||
class Head_adaLN(nn.Module):
|
class Head_adaLN(nn.Module):
|
||||||
|
|
||||||
def __init__(self, dim, out_dim, patch_size, eps=1e-6, adaln_tembed_dim=512):
|
def __init__(self, dim, out_dim, patch_size, eps=1e-6, adaln_tembed_dim=512):
|
||||||
@@ -1459,7 +1533,7 @@ class Head_adaLN(nn.Module):
|
|||||||
self.modulation.to(torch.float32)
|
self.modulation.to(torch.float32)
|
||||||
shift, scale = self.modulation(e).unsqueeze(2).chunk(2, dim=-1) # [B, T, 1, C]
|
shift, scale = self.modulation(e).unsqueeze(2).chunk(2, dim=-1) # [B, T, 1, C]
|
||||||
return self.head(self.norm(x.view(B, T, -1, C).float()).mul_(1 + scale).add_(shift).view(B, N, C).to(x.dtype))
|
return self.head(self.norm(x.view(B, T, -1, C).float()).mul_(1 + scale).add_(shift).view(B, N, C).to(x.dtype))
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
class MLPProj(torch.nn.Module):
|
class MLPProj(torch.nn.Module):
|
||||||
@@ -1732,7 +1806,7 @@ class WanModel(torch.nn.Module):
|
|||||||
nn.SiLU(),
|
nn.SiLU(),
|
||||||
ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
|
ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
|
||||||
)
|
)
|
||||||
|
|
||||||
self.original_patch_embedding = self.patch_embedding
|
self.original_patch_embedding = self.patch_embedding
|
||||||
self.expanded_patch_embedding = self.patch_embedding
|
self.expanded_patch_embedding = self.patch_embedding
|
||||||
|
|
||||||
@@ -1803,18 +1877,12 @@ class WanModel(torch.nn.Module):
|
|||||||
self.head = Head_adaLN(dim, out_dim, patch_size, eps, adaln_tembed_dim=512)
|
self.head = Head_adaLN(dim, out_dim, patch_size, eps, adaln_tembed_dim=512)
|
||||||
|
|
||||||
d = self.dim // self.num_heads
|
d = self.dim // self.num_heads
|
||||||
self.rope_embedder = EmbedND_RifleX(
|
self.rope_embedder = EmbedND_RifleX(d, 10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=None, k=None)
|
||||||
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
|
self.cached_freqs = self.cached_shape = self.cached_cond = None
|
||||||
|
|
||||||
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
# buffers (don't use register_buffer otherwise dtype will be changed in to())
|
||||||
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
|
||||||
|
|
||||||
if model_type == 'i2v' or model_type == 'fl2v':
|
if model_type == 'i2v' or model_type == 'fl2v':
|
||||||
self.img_emb = MLPProj(1280, dim, fl_pos_emb=model_type == 'fl2v')
|
self.img_emb = MLPProj(1280, dim, fl_pos_emb=model_type == 'fl2v')
|
||||||
|
|
||||||
@@ -2149,7 +2217,8 @@ class WanModel(torch.nn.Module):
|
|||||||
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
|
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
|
||||||
num_cond_latents=None,
|
num_cond_latents=None,
|
||||||
add_text_emb=None,
|
add_text_emb=None,
|
||||||
sdancer_input=None # SteadyDancer
|
sdancer_input=None, # SteadyDancer
|
||||||
|
one_to_all_input=None, # One-to-All
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
Forward pass through the diffusion model
|
Forward pass through the diffusion model
|
||||||
@@ -2173,9 +2242,9 @@ class WanModel(torch.nn.Module):
|
|||||||
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
|
||||||
"""
|
"""
|
||||||
# Stand-In only used on first positive pass, then cached in kv_cache
|
# Stand-In only used on first positive pass, then cached in kv_cache
|
||||||
if is_uncond or current_step > 0:
|
if is_uncond or current_step > 0:
|
||||||
standin_input = None
|
standin_input = None
|
||||||
|
|
||||||
# MTV Crafter motion projection
|
# MTV Crafter motion projection
|
||||||
if mtv_motion_tokens is not None:
|
if mtv_motion_tokens is not None:
|
||||||
bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1]
|
bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1]
|
||||||
@@ -2213,14 +2282,14 @@ class WanModel(torch.nn.Module):
|
|||||||
|
|
||||||
lynx_ip_scale = lynx_embeds.get("ip_scale", 1.0)
|
lynx_ip_scale = lynx_embeds.get("ip_scale", 1.0)
|
||||||
lynx_ref_scale = lynx_embeds.get("ref_scale", 1.0)
|
lynx_ref_scale = lynx_embeds.get("ref_scale", 1.0)
|
||||||
|
|
||||||
|
|
||||||
#s2v
|
#s2v
|
||||||
if self.model_type == 's2v' and s2v_audio_input is not None:
|
if self.model_type == 's2v' and s2v_audio_input is not None:
|
||||||
if is_uncond:
|
if is_uncond:
|
||||||
s2v_audio_input = s2v_audio_input * 0 # to match original code
|
s2v_audio_input = s2v_audio_input * 0 # to match original code
|
||||||
s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, s2v_motion_frames[0]), s2v_audio_input], dim=-1)
|
s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, s2v_motion_frames[0]), s2v_audio_input], dim=-1)
|
||||||
|
|
||||||
audio_emb_res = self.casual_audio_encoder(s2v_audio_input)
|
audio_emb_res = self.casual_audio_encoder(s2v_audio_input)
|
||||||
if self.enable_adain:
|
if self.enable_adain:
|
||||||
audio_emb_global, audio_emb = audio_emb_res
|
audio_emb_global, audio_emb = audio_emb_res
|
||||||
@@ -2241,7 +2310,7 @@ class WanModel(torch.nn.Module):
|
|||||||
if sdancer_input is not None and sdancer_input['start_percent'] <= current_step_percentage <= sdancer_input['end_percent']:
|
if sdancer_input is not None and sdancer_input['start_percent'] <= current_step_percentage <= sdancer_input['end_percent']:
|
||||||
sdancer_enabled = True
|
sdancer_enabled = True
|
||||||
x_noise_clone = torch.stack(x)
|
x_noise_clone = torch.stack(x)
|
||||||
|
|
||||||
# I2V
|
# I2V
|
||||||
if y is not None:
|
if y is not None:
|
||||||
if hasattr(self, "randomref_embedding_pose") and unianim_data is not None:
|
if hasattr(self, "randomref_embedding_pose") and unianim_data is not None:
|
||||||
@@ -2251,6 +2320,40 @@ class WanModel(torch.nn.Module):
|
|||||||
y[0].add_(random_ref_emb, alpha=unianim_data["strength"])
|
y[0].add_(random_ref_emb, alpha=unianim_data["strength"])
|
||||||
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
|
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
|
#uni3c controlnet
|
||||||
if uni3c_data is not None:
|
if uni3c_data is not None:
|
||||||
render_latent = uni3c_data["render_latent"].to(self.base_dtype)
|
render_latent = uni3c_data["render_latent"].to(self.base_dtype)
|
||||||
@@ -2278,15 +2381,13 @@ class WanModel(torch.nn.Module):
|
|||||||
else:
|
else:
|
||||||
self.original_patch_embedding.to(self.main_device)
|
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]
|
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
|
# ovi audio model
|
||||||
if self.audio_model is not None:
|
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]
|
x_ovi = [self.audio_model.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x_ovi[0].dtype) for u in x_ovi]
|
||||||
grid_sizes_ovi = torch.stack([torch.tensor(u.shape[1:2], dtype=torch.long) for u in x_ovi])
|
grid_sizes_ovi = torch.stack([torch.tensor(u.shape[1:2], dtype=torch.long) for u in x_ovi])
|
||||||
seq_lens_ovi = torch.tensor([u.size(1) for u in x_ovi], dtype=torch.int32)
|
seq_lens_ovi = torch.tensor([u.size(1) for u in x_ovi], dtype=torch.int32)
|
||||||
x_ovi = torch.cat([torch.cat([u, u.new_zeros(1, seq_len_ovi - u.size(1), u.size(2))], dim=1) for u in x_ovi])
|
x_ovi = torch.cat([torch.cat([u, u.new_zeros(1, seq_len_ovi - u.size(1), u.size(2))], dim=1) for u in x_ovi])
|
||||||
d = self.dim // self.num_heads
|
d = self.dim // self.num_heads
|
||||||
freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device)
|
freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device)
|
||||||
x_ovi = x_ovi.to(self.main_device, self.base_dtype)
|
x_ovi = x_ovi.to(self.main_device, self.base_dtype)
|
||||||
@@ -2375,7 +2476,7 @@ class WanModel(torch.nn.Module):
|
|||||||
seq_len += end_ref_latent_seq_len
|
seq_len += end_ref_latent_seq_len
|
||||||
x = [torch.cat([u, end_ref_latent.unsqueeze(0)], dim=1) for end_ref_latent, u in zip(end_ref_latent, x)]
|
x = [torch.cat([u, end_ref_latent.unsqueeze(0)], dim=1) for end_ref_latent, u in zip(end_ref_latent, x)]
|
||||||
|
|
||||||
|
|
||||||
x = torch.cat([torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1) for u in x])
|
x = torch.cat([torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1) for u in x])
|
||||||
|
|
||||||
if self.trainable_cond_mask is not None:
|
if self.trainable_cond_mask is not None:
|
||||||
@@ -2397,11 +2498,11 @@ class WanModel(torch.nn.Module):
|
|||||||
|
|
||||||
if freqs is None and "comfy" in self.rope_func: #comfy rope
|
if freqs is None and "comfy" in self.rope_func: #comfy rope
|
||||||
current_shape = (F, H, W)
|
current_shape = (F, H, W)
|
||||||
|
|
||||||
has_cond = attn_cond is not None
|
has_cond = attn_cond is not None
|
||||||
|
|
||||||
if (self.cached_freqs is not None and
|
if (self.cached_freqs is not None and
|
||||||
self.cached_shape == current_shape and
|
self.cached_shape == current_shape and
|
||||||
self.cached_cond == has_cond and
|
self.cached_cond == has_cond and
|
||||||
self.cached_rope_k == self.rope_embedder.k and
|
self.cached_rope_k == self.rope_embedder.k and
|
||||||
self.cached_ntk_alphas == ntk_alphas
|
self.cached_ntk_alphas == ntk_alphas
|
||||||
@@ -2411,9 +2512,9 @@ class WanModel(torch.nn.Module):
|
|||||||
freqs = self.rope_encode_comfy(F, H, W, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
|
freqs = self.rope_encode_comfy(F, H, W, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
|
||||||
if s2v_ref_latent is not None:
|
if s2v_ref_latent is not None:
|
||||||
freqs_ref = self.rope_encode_comfy(
|
freqs_ref = self.rope_encode_comfy(
|
||||||
s2v_ref_latent.shape[2],
|
s2v_ref_latent.shape[2],
|
||||||
s2v_ref_latent.shape[3],
|
s2v_ref_latent.shape[3],
|
||||||
s2v_ref_latent.shape[4],
|
s2v_ref_latent.shape[4],
|
||||||
t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
|
t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
|
||||||
freqs = torch.cat([freqs, freqs_ref], dim=1)
|
freqs = torch.cat([freqs, freqs_ref], dim=1)
|
||||||
|
|
||||||
@@ -2463,6 +2564,9 @@ class WanModel(torch.nn.Module):
|
|||||||
time_embed_dtype = self.base_dtype
|
time_embed_dtype = self.base_dtype
|
||||||
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim
|
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
|
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:
|
else:
|
||||||
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
|
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
|
||||||
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
||||||
@@ -2480,7 +2584,7 @@ class WanModel(torch.nn.Module):
|
|||||||
last_timestep = t[:, -1:]
|
last_timestep = t[:, -1:]
|
||||||
padding = last_timestep.expand(t.size(0), seq_len_ovi - t.size(1))
|
padding = last_timestep.expand(t.size(0), seq_len_ovi - t.size(1))
|
||||||
t_ovi = torch.cat([t, padding], dim=1)
|
t_ovi = torch.cat([t, padding], dim=1)
|
||||||
|
|
||||||
e_ovi = self.audio_model.time_embedding(sinusoidal_embedding_1d(self.audio_model.freq_dim, t_ovi.flatten()).to(time_embed_dtype)).unsqueeze(0) # b, dim
|
e_ovi = self.audio_model.time_embedding(sinusoidal_embedding_1d(self.audio_model.freq_dim, t_ovi.flatten()).to(time_embed_dtype)).unsqueeze(0) # b, dim
|
||||||
e0_ovi = self.audio_model.time_projection(e_ovi).unflatten(2, (6, self.dim)).movedim(1, 2) # B, seq_len, 6, dim
|
e0_ovi = self.audio_model.time_projection(e_ovi).unflatten(2, (6, self.dim)).movedim(1, 2) # B, seq_len, 6, dim
|
||||||
else:
|
else:
|
||||||
@@ -2516,14 +2620,14 @@ class WanModel(torch.nn.Module):
|
|||||||
if expanded_timesteps:
|
if expanded_timesteps:
|
||||||
e = e.view(b, f, 1, 1, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], self.dim)
|
e = e.view(b, f, 1, 1, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], self.dim)
|
||||||
e0 = e0.view(b, f, 1, 1, 6, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], 6, self.dim)
|
e0 = e0.view(b, f, 1, 1, 6, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], 6, self.dim)
|
||||||
|
|
||||||
e = e.flatten(1, 3)
|
e = e.flatten(1, 3)
|
||||||
e0 = e0.flatten(1, 3)
|
e0 = e0.flatten(1, 3)
|
||||||
|
|
||||||
e0 = e0.transpose(1, 2)
|
e0 = e0.transpose(1, 2)
|
||||||
if not e0.is_contiguous():
|
if not e0.is_contiguous():
|
||||||
e0 = e0.contiguous()
|
e0 = e0.contiguous()
|
||||||
|
|
||||||
e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
|
e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||||
|
|
||||||
# clip vision embedding
|
# clip vision embedding
|
||||||
@@ -2580,7 +2684,7 @@ class WanModel(torch.nn.Module):
|
|||||||
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
|
||||||
for u in nag_context
|
for u in nag_context
|
||||||
]).to(text_embed_dtype))
|
]).to(text_embed_dtype))
|
||||||
|
|
||||||
if self.offload_txt_emb:
|
if self.offload_txt_emb:
|
||||||
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
|
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||||
|
|
||||||
@@ -2637,7 +2741,7 @@ class WanModel(torch.nn.Module):
|
|||||||
accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device)
|
accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device)
|
||||||
if pred_id is None:
|
if pred_id is None:
|
||||||
pred_id = self.teacache_state.new_prediction(cache_device=self.cache_device)
|
pred_id = self.teacache_state.new_prediction(cache_device=self.cache_device)
|
||||||
should_calc = True
|
should_calc = True
|
||||||
else:
|
else:
|
||||||
previous_modulated_input = self.teacache_state.get(pred_id)['previous_modulated_input']
|
previous_modulated_input = self.teacache_state.get(pred_id)['previous_modulated_input']
|
||||||
previous_modulated_input = previous_modulated_input.to(device)
|
previous_modulated_input = previous_modulated_input.to(device)
|
||||||
@@ -2665,7 +2769,7 @@ class WanModel(torch.nn.Module):
|
|||||||
accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(self.cache_device)
|
accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(self.cache_device)
|
||||||
|
|
||||||
previous_modulated_input = e.to(self.cache_device).clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.to(self.cache_device).clone()
|
previous_modulated_input = e.to(self.cache_device).clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.to(self.cache_device).clone()
|
||||||
|
|
||||||
if not should_calc:
|
if not should_calc:
|
||||||
x = x.to(previous_residual.dtype) + previous_residual.to(x.device)
|
x = x.to(previous_residual.dtype) + previous_residual.to(x.device)
|
||||||
self.teacache_state.update(
|
self.teacache_state.update(
|
||||||
@@ -2803,7 +2907,11 @@ class WanModel(torch.nn.Module):
|
|||||||
lynx_x_ip=lynx_x_ip,
|
lynx_x_ip=lynx_x_ip,
|
||||||
lynx_ip_scale=lynx_ip_scale,
|
lynx_ip_scale=lynx_ip_scale,
|
||||||
lynx_ref_scale=lynx_ref_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:
|
if self.audio_model is not None:
|
||||||
kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)
|
kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)
|
||||||
@@ -2812,7 +2920,7 @@ class WanModel(torch.nn.Module):
|
|||||||
kwargs['seq_lens_ovi'] = seq_lens_ovi
|
kwargs['seq_lens_ovi'] = seq_lens_ovi
|
||||||
kwargs['freqs_ovi'] = freqs_ovi
|
kwargs['freqs_ovi'] = freqs_ovi
|
||||||
|
|
||||||
|
|
||||||
if vace_data is not None:
|
if vace_data is not None:
|
||||||
vace_hint_list = []
|
vace_hint_list = []
|
||||||
vace_scale_list = []
|
vace_scale_list = []
|
||||||
@@ -2828,7 +2936,7 @@ class WanModel(torch.nn.Module):
|
|||||||
vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs)
|
vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs)
|
||||||
vace_hint_list.append(vace_hints)
|
vace_hint_list.append(vace_hints)
|
||||||
vace_scale_list.append(1.0)
|
vace_scale_list.append(1.0)
|
||||||
|
|
||||||
kwargs['vace_hints'] = vace_hint_list
|
kwargs['vace_hints'] = vace_hint_list
|
||||||
kwargs['vace_context_scale'] = vace_scale_list
|
kwargs['vace_context_scale'] = vace_scale_list
|
||||||
|
|
||||||
@@ -2898,7 +3006,17 @@ class WanModel(torch.nn.Module):
|
|||||||
if b in self.slg_blocks and is_uncond:
|
if b in self.slg_blocks and is_uncond:
|
||||||
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
|
||||||
continue
|
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:
|
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
|
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:
|
if block.has_face_fuser_block and motion_vec is not None:
|
||||||
@@ -2924,6 +3042,31 @@ class WanModel(torch.nn.Module):
|
|||||||
#controlnet
|
#controlnet
|
||||||
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
|
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"]
|
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:
|
if lynx_ref_feature_extractor:
|
||||||
return lynx_ref_buffer
|
return lynx_ref_buffer
|
||||||
@@ -2953,20 +3096,20 @@ class WanModel(torch.nn.Module):
|
|||||||
accumulated_error = 0.0,
|
accumulated_error = 0.0,
|
||||||
cache_ovi = x_ovi.clone().to(original_x.device) - original_x_ovi if x_ovi is not None else None
|
cache_ovi = x_ovi.clone().to(original_x.device) - original_x_ovi if x_ovi is not None else None
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
if self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None:
|
if self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None:
|
||||||
self.easycache_state.update(
|
self.easycache_state.update(
|
||||||
pred_id,
|
pred_id,
|
||||||
previous_raw_output=x.clone(),
|
previous_raw_output=x.clone(),
|
||||||
)
|
)
|
||||||
|
|
||||||
if self.ref_conv is not None and fun_ref is not None:
|
if self.ref_conv is not None and fun_ref is not None:
|
||||||
fun_ref_length = fun_ref.size(1)
|
fun_ref_length = fun_ref.size(1)
|
||||||
x = x[:, fun_ref_length:]
|
x = x[:, fun_ref_length:]
|
||||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||||
|
|
||||||
if end_ref_latent is not None:
|
if end_ref_latent is not None:
|
||||||
end_ref_latent_length = end_ref_latent.size(1)
|
end_ref_latent_length = end_ref_latent.size(1)
|
||||||
x = x[:, :-end_ref_latent_length]
|
x = x[:, :-end_ref_latent_length]
|
||||||
@@ -2976,10 +3119,11 @@ class WanModel(torch.nn.Module):
|
|||||||
x = x[:, :self.original_seq_len]
|
x = x[:, :self.original_seq_len]
|
||||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||||
|
|
||||||
|
|
||||||
x = x[:, :self.original_seq_len]
|
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:
|
if x_ovi is not None:
|
||||||
x_ovi = self.audio_model.head(x_ovi, e_ovi.to(x_ovi.device))
|
x_ovi = self.audio_model.head(x_ovi, e_ovi.to(x_ovi.device))
|
||||||
@@ -2987,9 +3131,9 @@ class WanModel(torch.nn.Module):
|
|||||||
assert len(x) == len(grid_sizes_ovi)
|
assert len(x) == len(grid_sizes_ovi)
|
||||||
x_ovi = [u[:gs] for u, gs in zip(x_ovi, grid_sizes_ovi)]
|
x_ovi = [u[:gs] for u, gs in zip(x_ovi, grid_sizes_ovi)]
|
||||||
x_ovi = [u.float() for u in x_ovi]
|
x_ovi = [u.float() for u in x_ovi]
|
||||||
|
|
||||||
x = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type]
|
x = self.unpatchify(x, original_grid_sizes)
|
||||||
x = [u[:, :orig_frames, ...].float() for u in x]
|
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)
|
return (x, x_ovi, pred_id) if pred_id is not None else (x, x_ovi, None)
|
||||||
|
|
||||||
def unpatchify(self, x, grid_sizes):
|
def unpatchify(self, x, grid_sizes):
|
||||||
|
|||||||
Reference in New Issue
Block a user