init
This commit is contained in:
@@ -0,0 +1,199 @@
|
|||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
from einops import rearrange
|
||||||
|
|
||||||
|
from ..wanvideo.modules.attention import attention
|
||||||
|
|
||||||
|
def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor):
|
||||||
|
return (x * (1 + scale) + shift)
|
||||||
|
|
||||||
|
|
||||||
|
def sinusoidal_embedding_1d(dim, position):
|
||||||
|
sinusoid = torch.outer(position.type(torch.float64), torch.pow(
|
||||||
|
10000, -torch.arange(dim//2, dtype=torch.float64, device=position.device).div(dim//2)))
|
||||||
|
x = torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
|
||||||
|
return x.to(position.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
def precompute_freqs_cis_3d(dim: int, end: int = 1024, theta: float = 10000.0):
|
||||||
|
# 3d rope precompute
|
||||||
|
f_freqs_cis = precompute_freqs_cis(dim - 2 * (dim // 3), end, theta)
|
||||||
|
h_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
|
||||||
|
w_freqs_cis = precompute_freqs_cis(dim // 3, end, theta)
|
||||||
|
return f_freqs_cis, h_freqs_cis, w_freqs_cis
|
||||||
|
|
||||||
|
|
||||||
|
def precompute_freqs_cis(dim: int, end: int = 1024, theta: float = 10000.0):
|
||||||
|
# 1d rope precompute
|
||||||
|
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)
|
||||||
|
[: (dim // 2)].double() / dim))
|
||||||
|
freqs = torch.outer(torch.arange(end, device=freqs.device), freqs)
|
||||||
|
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64
|
||||||
|
return freqs_cis
|
||||||
|
|
||||||
|
|
||||||
|
def rope_apply(x, freqs, num_heads):
|
||||||
|
x = rearrange(x, "b s (n d) -> b s n d", n=num_heads)
|
||||||
|
x_out = torch.view_as_complex(x.to(torch.float64).reshape(
|
||||||
|
x.shape[0], x.shape[1], x.shape[2], -1, 2))
|
||||||
|
x_out = torch.view_as_real(x_out * freqs).flatten(2)
|
||||||
|
return x_out.to(x.dtype)
|
||||||
|
|
||||||
|
|
||||||
|
class RMSNorm(nn.Module):
|
||||||
|
def __init__(self, dim, eps=1e-5):
|
||||||
|
super().__init__()
|
||||||
|
self.eps = eps
|
||||||
|
self.weight = nn.Parameter(torch.ones(dim))
|
||||||
|
|
||||||
|
def norm(self, x):
|
||||||
|
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
dtype = x.dtype
|
||||||
|
return self.norm(x.float()).to(dtype) * self.weight
|
||||||
|
|
||||||
|
|
||||||
|
class AttentionModule(nn.Module):
|
||||||
|
def __init__(self, num_heads, head_dim):
|
||||||
|
super().__init__()
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = head_dim
|
||||||
|
|
||||||
|
def forward(self, q, k, v):
|
||||||
|
b, n, d = q.size(0), self.num_heads, self.head_dim
|
||||||
|
x = attention(
|
||||||
|
q.view(b, -1, n, d),
|
||||||
|
k.view(b, -1, n, d),
|
||||||
|
v.view(b, -1, n, d)
|
||||||
|
)
|
||||||
|
return x.flatten(2)
|
||||||
|
|
||||||
|
|
||||||
|
class SelfAttention(nn.Module):
|
||||||
|
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = dim // num_heads
|
||||||
|
|
||||||
|
self.q = nn.Linear(dim, dim)
|
||||||
|
self.k = nn.Linear(dim, dim)
|
||||||
|
self.v = nn.Linear(dim, dim)
|
||||||
|
self.o = nn.Linear(dim, dim)
|
||||||
|
self.norm_q = RMSNorm(dim, eps=eps)
|
||||||
|
self.norm_k = RMSNorm(dim, eps=eps)
|
||||||
|
|
||||||
|
self.attn = AttentionModule(self.num_heads, self.head_dim)
|
||||||
|
|
||||||
|
def forward(self, x, freqs):
|
||||||
|
q = self.norm_q(self.q(x))
|
||||||
|
k = self.norm_k(self.k(x))
|
||||||
|
v = self.v(x)
|
||||||
|
q = rope_apply(q, freqs, self.num_heads)
|
||||||
|
k = rope_apply(k, freqs, self.num_heads)
|
||||||
|
x = self.attn(q, k, v)
|
||||||
|
return self.o(x)
|
||||||
|
|
||||||
|
|
||||||
|
class CrossAttention(nn.Module):
|
||||||
|
def __init__(self, dim: int, num_heads: int, eps: float = 1e-6, clip_fea: torch.Tensor = None):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.head_dim = dim // num_heads
|
||||||
|
|
||||||
|
self.q = nn.Linear(dim, dim)
|
||||||
|
self.k = nn.Linear(dim, dim)
|
||||||
|
self.v = nn.Linear(dim, dim)
|
||||||
|
self.o = nn.Linear(dim, dim)
|
||||||
|
self.norm_q = RMSNorm(dim, eps=eps)
|
||||||
|
self.norm_k = RMSNorm(dim, eps=eps)
|
||||||
|
|
||||||
|
|
||||||
|
self.k_img = nn.Linear(dim, dim)
|
||||||
|
self.v_img = nn.Linear(dim, dim)
|
||||||
|
self.norm_k_img = RMSNorm(dim, eps=eps)
|
||||||
|
|
||||||
|
self.attn = AttentionModule(self.num_heads, self.head_dim)
|
||||||
|
|
||||||
|
def forward(self, x: torch.Tensor, y: torch.Tensor, clip_fea: torch.Tensor = None):
|
||||||
|
ctx = y
|
||||||
|
q = self.norm_q(self.q(x))
|
||||||
|
k = self.norm_k(self.k(ctx))
|
||||||
|
v = self.v(ctx)
|
||||||
|
x = self.attn(q, k, v)
|
||||||
|
if clip_fea is not None:
|
||||||
|
k_img = self.norm_k_img(self.k_img(clip_fea))
|
||||||
|
v_img = self.v_img(clip_fea)
|
||||||
|
y = self.attn(q, k_img, v_img)
|
||||||
|
x = x + y
|
||||||
|
return self.o(x)
|
||||||
|
|
||||||
|
|
||||||
|
class GateModule(nn.Module):
|
||||||
|
def __init__(self,):
|
||||||
|
super().__init__()
|
||||||
|
|
||||||
|
def forward(self, x, gate, residual):
|
||||||
|
return x + gate * residual
|
||||||
|
|
||||||
|
class DiTBlock(nn.Module):
|
||||||
|
def __init__(self, dim: int, num_heads: int, ffn_dim: int, eps: float = 1e-6):
|
||||||
|
super().__init__()
|
||||||
|
self.dim = dim
|
||||||
|
self.num_heads = num_heads
|
||||||
|
self.ffn_dim = ffn_dim
|
||||||
|
|
||||||
|
self.self_attn = SelfAttention(dim, num_heads, eps)
|
||||||
|
self.cross_attn = CrossAttention(dim, num_heads, eps)
|
||||||
|
self.norm1 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
||||||
|
self.norm2 = nn.LayerNorm(dim, eps=eps, elementwise_affine=False)
|
||||||
|
self.norm3 = nn.LayerNorm(dim, eps=eps)
|
||||||
|
self.ffn = nn.Sequential(nn.Linear(dim, ffn_dim), nn.GELU(
|
||||||
|
approximate='tanh'), nn.Linear(ffn_dim, dim))
|
||||||
|
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||||
|
self.gate = GateModule()
|
||||||
|
|
||||||
|
def forward(self, x, context, t_mod, freqs, clip_fea=None):
|
||||||
|
has_seq = len(t_mod.shape) == 4
|
||||||
|
chunk_dim = 2 if has_seq else 1
|
||||||
|
# msa: multi-head self-attention mlp: multi-layer perceptron
|
||||||
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||||
|
self.modulation.to(dtype=t_mod.dtype, device=t_mod.device) + t_mod).chunk(6, dim=chunk_dim)
|
||||||
|
if has_seq:
|
||||||
|
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (
|
||||||
|
shift_msa.squeeze(2), scale_msa.squeeze(2), gate_msa.squeeze(2),
|
||||||
|
shift_mlp.squeeze(2), scale_mlp.squeeze(2), gate_mlp.squeeze(2),
|
||||||
|
)
|
||||||
|
input_x = modulate(self.norm1(x), shift_msa, scale_msa)
|
||||||
|
x = self.gate(x, gate_msa, self.self_attn(input_x, freqs))
|
||||||
|
x = x + self.cross_attn(self.norm3(x), context, clip_fea=clip_fea)
|
||||||
|
input_x = modulate(self.norm2(x), shift_mlp, scale_mlp)
|
||||||
|
x = self.gate(x, gate_mlp, self.ffn(input_x))
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class WanModelDualControl(torch.nn.Module):
|
||||||
|
def __init__(self, dim: int, ffn_dim: int, eps: float, num_heads: int, control_layers = 12):
|
||||||
|
super().__init__()
|
||||||
|
self.control_layers = control_layers
|
||||||
|
self.control_blocks_dense = nn.ModuleList([
|
||||||
|
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
|
||||||
|
for _ in range(self.control_layers)
|
||||||
|
])
|
||||||
|
|
||||||
|
self.control_blocks_sparse = nn.ModuleList([
|
||||||
|
DiTBlock(dim//2, num_heads//2, ffn_dim//2, eps)
|
||||||
|
for _ in range(self.control_layers)
|
||||||
|
])
|
||||||
|
|
||||||
|
self.control_initial_combine_linear_dense = torch.nn.Linear(dim, dim//2)
|
||||||
|
self.control_initial_combine_linear_sparse = torch.nn.Linear(dim, dim//2)
|
||||||
|
|
||||||
|
self.control_text_linear = torch.nn.Linear(dim, dim//2)
|
||||||
|
self.control_t_mod = torch.nn.Linear(dim, dim//2)
|
||||||
|
|
||||||
|
self.control_combine_linears = torch.nn.ModuleList([torch.nn.Linear(dim//2, dim) for _ in range(self.control_layers)])
|
||||||
|
head_dim = dim // num_heads
|
||||||
|
self.freqs = precompute_freqs_cis_3d(head_dim)
|
||||||
@@ -0,0 +1,88 @@
|
|||||||
|
import torch
|
||||||
|
from ..utils import log
|
||||||
|
import comfy.model_management as mm
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
|
||||||
|
class WanVideoAddDualControlEmbeds:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||||
|
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||||
|
"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"}),
|
||||||
|
"first_frame_noise_level": ("FLOAT", {"default": 0.925926, "min": 0.0, "max": 1.0, "step": 0.000001, "tooltip": "Noise level for the first frame when using previous frames"}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"dense": ("IMAGE", {"tooltip": "Dense control signal (depth) video input"}),
|
||||||
|
"sparse": ("IMAGE", {"tooltip": "Sparse control signal (tracks) video input"}),
|
||||||
|
"prev_images": ("IMAGE", {"tooltip": "Previous frames for temporal consistency, default is 8 frames"}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||||
|
RETURN_NAMES = ("image_embeds",)
|
||||||
|
FUNCTION = "add"
|
||||||
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
|
def add(self, embeds, vae, strength, start_percent, end_percent, first_frame_noise_level, dense=None, sparse=None, prev_images=None):
|
||||||
|
updated = dict(embeds)
|
||||||
|
updated.setdefault("dual_control", {})
|
||||||
|
|
||||||
|
if dense is None and sparse is None:
|
||||||
|
raise ValueError("At least one of dense or sparse inputs must be provided.")
|
||||||
|
|
||||||
|
num_frames = dense.shape[0] if dense is not None else sparse.shape[0]
|
||||||
|
height = dense.shape[1] if dense is not None else sparse.shape[1]
|
||||||
|
width = dense.shape[2] if dense is not None else sparse.shape[2]
|
||||||
|
msk = torch.ones(1, num_frames, height//8, width//8, device=device)
|
||||||
|
msk[:, 1:] = 0
|
||||||
|
msk = torch.concat([torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]], dim=1)
|
||||||
|
msk = msk.view(1, msk.shape[1] // 4, 4, height//8, width//8)
|
||||||
|
msk = msk.transpose(1, 2)
|
||||||
|
|
||||||
|
dense_input_latent = sparse_input_latent = None
|
||||||
|
|
||||||
|
vae.to(device)
|
||||||
|
if dense is not None:
|
||||||
|
dense_images = dense[..., :3].permute(3, 0, 1, 2) * 2 - 1
|
||||||
|
dense_images = 1 - dense_images # Invert colors for depth to match the usual range in comfy
|
||||||
|
dense_video_latent = vae.encode([dense_images.to(device, vae.dtype)], device, tiled=False)
|
||||||
|
dense_first = (dense_images[:, :1]).to(device, vae.dtype)
|
||||||
|
vae_input_dense = torch.cat([dense_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
|
||||||
|
dense_concat_latent = vae.encode([vae_input_dense], device, tiled=False)
|
||||||
|
dense_concat_latent = torch.cat([msk, dense_concat_latent], dim=1)
|
||||||
|
dense_input_latent = torch.cat([dense_video_latent, dense_concat_latent],dim=1)
|
||||||
|
if sparse is not None:
|
||||||
|
sparse_images = sparse[..., :3].permute(3, 0, 1, 2) * 2 - 1
|
||||||
|
sparse_video_latent = vae.encode([sparse_images.to(device, vae.dtype)], device, tiled=False)
|
||||||
|
sparse_first = (sparse_images[:, :1]).to(device, vae.dtype)
|
||||||
|
vae_input_sparse = torch.cat([sparse_first, torch.zeros(3, num_frames-1, height, width, device=device, dtype=vae.dtype)], dim=1)
|
||||||
|
sparse_concat_latent = vae.encode([vae_input_sparse], device, tiled=False)
|
||||||
|
sparse_concat_latent = torch.cat([msk, sparse_concat_latent], dim=1)
|
||||||
|
sparse_input_latent = torch.cat([sparse_video_latent, sparse_concat_latent],dim=1)
|
||||||
|
|
||||||
|
if prev_images is not None:
|
||||||
|
prev_images = prev_images[..., :3].permute(3, 0, 1, 2) * 2 - 1
|
||||||
|
prev_video_latent = vae.encode([prev_images.to(device, vae.dtype)], device, tiled=False)
|
||||||
|
updated["dual_control"]["prev_latent"] = prev_video_latent[0]
|
||||||
|
|
||||||
|
vae.to(offload_device)
|
||||||
|
updated["dual_control"]["dense_input_latent"] = dense_input_latent
|
||||||
|
updated["dual_control"]["sparse_input_latent"] = sparse_input_latent
|
||||||
|
updated["dual_control"]["strength"] = strength
|
||||||
|
updated["dual_control"]["start_percent"] = start_percent
|
||||||
|
updated["dual_control"]["end_percent"] = end_percent
|
||||||
|
updated["dual_control"]["first_frame_noise_level"] = first_frame_noise_level
|
||||||
|
return (updated,)
|
||||||
|
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"WanVideoAddDualControlEmbeds": WanVideoAddDualControlEmbeds,
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"WanVideoAddDualControlEmbeds": "WanVideo Add Dual Control Embeds",
|
||||||
|
}
|
||||||
+2
-1
@@ -49,6 +49,7 @@ OPTIONAL_MODULES = [
|
|||||||
(".WanMove.nodes", "WanMove"),
|
(".WanMove.nodes", "WanMove"),
|
||||||
(".SCAIL.nodes", "SCAIL"),
|
(".SCAIL.nodes", "SCAIL"),
|
||||||
(".LongCat.nodes", "LongCat"),
|
(".LongCat.nodes", "LongCat"),
|
||||||
|
(".LongVie2.nodes", "LongVie2"),
|
||||||
]
|
]
|
||||||
|
|
||||||
def register_nodes(module_path: str, name: str, optional: bool) -> None:
|
def register_nodes(module_path: str, name: str, optional: bool) -> None:
|
||||||
@@ -71,4 +72,4 @@ for module_path, name in REQUIRED_MODULES:
|
|||||||
for module_path, name in OPTIONAL_MODULES:
|
for module_path, name in OPTIONAL_MODULES:
|
||||||
register_nodes(module_path, name, optional=True)
|
register_nodes(module_path, name, optional=True)
|
||||||
|
|
||||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||||
|
|||||||
+1
-1
@@ -56,7 +56,7 @@ def _replace_linear(model, compute_dtype, state_dict, prefix="", patches=None, s
|
|||||||
module_prefix = module_prefix.replace("_orig_mod.", "")
|
module_prefix = module_prefix.replace("_orig_mod.", "")
|
||||||
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert)
|
_replace_linear(module, compute_dtype, state_dict, module_prefix, patches, scale_weights, compile_args, modules_to_not_convert)
|
||||||
|
|
||||||
if isinstance(module, nn.Linear) and "loras" not in module_prefix and name not in modules_to_not_convert:
|
if isinstance(module, nn.Linear) and "loras" not in module_prefix and "dual_controller" not in module_prefix and name not in modules_to_not_convert:
|
||||||
weight_key = module_prefix + "weight"
|
weight_key = module_prefix + "weight"
|
||||||
if weight_key not in state_dict:
|
if weight_key not in state_dict:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -1605,6 +1605,24 @@ class WanVideoModelLoader:
|
|||||||
block.ref_attn_v_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)
|
block.ref_attn_norm_k_img = WanRMSNorm(out_features, eps=1e-6)
|
||||||
|
|
||||||
|
if "blocks.0.control_blocks_dense.cross_attn.k.weight" in sd:
|
||||||
|
log.info("LongVie2 model detected, patching model...")
|
||||||
|
from .LongVie2.modules import WanModelDualControl
|
||||||
|
control_layers = 12
|
||||||
|
with init_empty_weights():
|
||||||
|
dual_controller = WanModelDualControl(dim=5120, ffn_dim=13824, eps=1e-06, num_heads=40, control_layers=control_layers)
|
||||||
|
for b in range(control_layers):
|
||||||
|
transformer.blocks[b].control_blocks_dense = dual_controller.control_blocks_dense[b]
|
||||||
|
transformer.blocks[b].control_blocks_sparse = dual_controller.control_blocks_sparse[b]
|
||||||
|
transformer.blocks[b].control_combine_linears = dual_controller.control_combine_linears[b]
|
||||||
|
transformer.dual_controller = nn.Module()
|
||||||
|
transformer.dual_controller.control_initial_combine_linear_dense = dual_controller.control_initial_combine_linear_dense
|
||||||
|
transformer.dual_controller.control_initial_combine_linear_sparse = dual_controller.control_initial_combine_linear_sparse
|
||||||
|
transformer.dual_controller.control_t_mod = dual_controller.control_t_mod
|
||||||
|
transformer.dual_controller.control_text_linear = dual_controller.control_text_linear
|
||||||
|
transformer.dual_controller_freqs = dual_controller.freqs
|
||||||
|
|
||||||
|
|
||||||
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)
|
||||||
|
|||||||
@@ -1259,6 +1259,18 @@ class WanVideoSampler:
|
|||||||
if context_options is None:
|
if context_options is None:
|
||||||
image_cond = replace_feature(image_cond.unsqueeze(0).clone(), track_pos.unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
|
image_cond = replace_feature(image_cond.unsqueeze(0).clone(), track_pos.unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
|
||||||
|
|
||||||
|
# LongVie2 dual control
|
||||||
|
dual_control_embeds = image_embeds.get("dual_control", None)
|
||||||
|
if dual_control_embeds is not None and context_options is None:
|
||||||
|
dual_control_input = dict_to_device(dual_control_embeds.copy(), device, dtype) if dual_control_embeds is not None else None
|
||||||
|
prev_latents = dual_control_input.get("prev_latent", None)
|
||||||
|
if prev_latents is not None:
|
||||||
|
_sigma = dual_control_embeds.get("first_frame_noise_level", 0.925926)
|
||||||
|
log.info(f"Using dual control previous latents with first frame noise level: {_sigma}")
|
||||||
|
latent[:, :1] = (1 - _sigma) * prev_latents[:, -1:].to(latent) + _sigma * noise[:, :1]
|
||||||
|
prev_ones = torch.ones(20, *prev_latents.shape[1:], device=device, dtype=dtype)
|
||||||
|
dual_control_input["prev_latent"] = torch.cat([prev_ones, prev_latents]).unsqueeze(0)
|
||||||
|
|
||||||
#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,
|
||||||
@@ -1513,6 +1525,19 @@ class WanVideoSampler:
|
|||||||
if wanmove_embeds is not None and context_window is not None:
|
if wanmove_embeds is not None and context_window is not None:
|
||||||
image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
|
image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
|
||||||
|
|
||||||
|
dual_control_in = None
|
||||||
|
if dual_control_embeds is not None:
|
||||||
|
if context_window is not None:
|
||||||
|
dual_control_in = dual_control_embeds.copy()
|
||||||
|
dense_input_latent = dual_control_embeds.get("dense_input_latent", None)
|
||||||
|
if dense_input_latent is not None:
|
||||||
|
dual_control_in["dense_input_latent"] = dual_control_embeds["dense_input_latent"][:, :, context_window]
|
||||||
|
sparse_input_latent = dual_control_embeds.get("sparse_input_latent", None)
|
||||||
|
if sparse_input_latent is not None:
|
||||||
|
dual_control_in["sparse_input_latent"] = dual_control_embeds["sparse_input_latent"][:, :, context_window]
|
||||||
|
else:
|
||||||
|
dual_control_in = dual_control_input
|
||||||
|
|
||||||
base_params = {
|
base_params = {
|
||||||
'x': [z], # latent
|
'x': [z], # latent
|
||||||
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
|
'y': [image_cond_input] if image_cond_input is not None else None, # image cond
|
||||||
@@ -1575,6 +1600,7 @@ class WanVideoSampler:
|
|||||||
"one_to_all_input": one_to_all_data, # One-to-All input
|
"one_to_all_input": one_to_all_data, # One-to-All input
|
||||||
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
|
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
|
||||||
"scail_input": scail_data_in, # SCAIL input
|
"scail_input": scail_data_in, # SCAIL input
|
||||||
|
"dual_control_input": dual_control_in, # LongVie2 dual control input
|
||||||
}
|
}
|
||||||
|
|
||||||
batch_size = 1
|
batch_size = 1
|
||||||
|
|||||||
@@ -2318,6 +2318,7 @@ class WanModel(torch.nn.Module):
|
|||||||
sdancer_input=None, # SteadyDancer
|
sdancer_input=None, # SteadyDancer
|
||||||
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
|
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
|
||||||
scail_input=None, # SCAIL pose
|
scail_input=None, # SCAIL pose
|
||||||
|
dual_control_input=None, # LongVie2 dual controlnet
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
Forward pass through the diffusion model
|
Forward pass through the diffusion model
|
||||||
@@ -2340,6 +2341,7 @@ class WanModel(torch.nn.Module):
|
|||||||
List[Tensor]:
|
List[Tensor]:
|
||||||
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]
|
||||||
"""
|
"""
|
||||||
|
print("input x shape:", x[0].shape)
|
||||||
# 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
|
||||||
@@ -2544,6 +2546,20 @@ class WanModel(torch.nn.Module):
|
|||||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||||
self.original_seq_len = x[0].shape[1]
|
self.original_seq_len = x[0].shape[1]
|
||||||
|
|
||||||
|
prev_latent = None
|
||||||
|
if dual_control_input is not None:
|
||||||
|
prev_latent = dual_control_input.get("prev_latent", None)
|
||||||
|
if prev_latent is not None:
|
||||||
|
F += prev_latent.shape[2]
|
||||||
|
print("Using Dual ControlNet prev latent shape:", prev_latent.shape)
|
||||||
|
prev_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in prev_latent]
|
||||||
|
prev_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in prev_x]
|
||||||
|
print("Dual ControlNet prev_x shape:", prev_x[0].shape)
|
||||||
|
print("x shape before concat:", x[0].shape)
|
||||||
|
seq_len += prev_x[0].shape[1]
|
||||||
|
print("F updated to:", F)
|
||||||
|
x = [torch.cat([u, v], dim=1) for u, v in zip(prev_x, x)]
|
||||||
|
|
||||||
# SCAIL pose
|
# SCAIL pose
|
||||||
if scail_input is not None:
|
if scail_input is not None:
|
||||||
scail_pose_latents = scail_input.get("pose_latent", None)
|
scail_pose_latents = scail_input.get("pose_latent", None)
|
||||||
@@ -2834,6 +2850,44 @@ class WanModel(torch.nn.Module):
|
|||||||
chunked_self_attention = False
|
chunked_self_attention = False
|
||||||
seq_chunks = 0
|
seq_chunks = 0
|
||||||
|
|
||||||
|
# dual control
|
||||||
|
if dual_control_input is not None and dual_control_input["start_percent"] <= current_step_percentage <= dual_control_input["end_percent"]:
|
||||||
|
dense_latent = dual_control_input["dense_input_latent"]
|
||||||
|
print("dense_latent shape:", dense_latent.shape)
|
||||||
|
sparse_latent = dual_control_input["sparse_input_latent"]
|
||||||
|
if dense_latent is None and sparse_latent is None:
|
||||||
|
raise ValueError("At least one of dense_input_latent or sparse_input_latent must be provided in dual_control_input")
|
||||||
|
|
||||||
|
if dense_latent is not None:
|
||||||
|
dense_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in dense_latent]
|
||||||
|
dense_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in dense_x]
|
||||||
|
dense = self.dual_controller.control_initial_combine_linear_dense(dense_x[0])
|
||||||
|
|
||||||
|
if sparse_latent is not None:
|
||||||
|
sparse_x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in sparse_latent]
|
||||||
|
sparse_x = [u.flatten(2).transpose(1, 2).to(self.base_dtype) for u in sparse_x]
|
||||||
|
sparse = self.dual_controller.control_initial_combine_linear_sparse(sparse_x[0])
|
||||||
|
|
||||||
|
if dense_latent is None:
|
||||||
|
dense = torch.zeros_like(sparse)
|
||||||
|
elif sparse_latent is None:
|
||||||
|
sparse = torch.zeros_like(dense)
|
||||||
|
|
||||||
|
control_context = clip_fea_control = None
|
||||||
|
if context != []:
|
||||||
|
control_context = self.dual_controller.control_text_linear(context)
|
||||||
|
if clip_embed is not None:
|
||||||
|
clip_fea_control = self.dual_controller.control_text_linear(clip_embed)
|
||||||
|
control_t_mod = self.dual_controller.control_t_mod(e0)
|
||||||
|
|
||||||
|
control_freqs = torch.cat([
|
||||||
|
self.dual_controller_freqs[0][:f].view(f, 1, 1, -1).expand(f, h, w, -1),
|
||||||
|
self.dual_controller_freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
|
||||||
|
self.dual_controller_freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
|
||||||
|
], dim=-1).reshape(f * h * w, 1, -1).to(x.device)
|
||||||
|
else:
|
||||||
|
dual_control_input = None
|
||||||
|
|
||||||
# MultiTalk
|
# MultiTalk
|
||||||
if multitalk_audio is not None:
|
if multitalk_audio is not None:
|
||||||
self.multitalk_audio_proj.to(self.main_device)
|
self.multitalk_audio_proj.to(self.main_device)
|
||||||
@@ -3173,6 +3227,18 @@ class WanModel(torch.nn.Module):
|
|||||||
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)
|
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----#
|
# ---post block----#
|
||||||
|
|
||||||
|
# dual controlnet
|
||||||
|
if dual_control_input is not None and (hasattr(block, "control_blocks_dense") or hasattr(block, "control_blocks_sparse")):
|
||||||
|
if dense_latent is not None and hasattr(block, "control_blocks_dense"):
|
||||||
|
dense = block.control_blocks_dense(dense, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control)
|
||||||
|
if sparse_latent is not None and hasattr(block, "control_blocks_sparse"):
|
||||||
|
sparse = block.control_blocks_sparse(sparse, control_context, control_t_mod, control_freqs, clip_fea=clip_fea_control)
|
||||||
|
|
||||||
|
if prev_latent is not None:
|
||||||
|
x[:, -self.original_seq_len:] += block.control_combine_linears(dense + sparse) * dual_control_input["strength"]
|
||||||
|
else:
|
||||||
|
x += block.control_combine_linears(dense + sparse) * dual_control_input["strength"]
|
||||||
|
|
||||||
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:
|
||||||
@@ -3275,8 +3341,10 @@ 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)
|
||||||
|
|
||||||
|
if prev_latent is not None:
|
||||||
x = x[:, :self.original_seq_len]
|
x = x[:, -self.original_seq_len:]
|
||||||
|
else:
|
||||||
|
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)
|
e_tr=e_token_replace.to(x.device) if use_token_replace else None, tr_start=token_replace_start, tr_num=replace_token_num)
|
||||||
|
|||||||
Reference in New Issue
Block a user