Init
This commit is contained in:
@@ -77,6 +77,13 @@ except Exception as e:
|
||||
OVI_NODE_CLASS_MAPPINGS = {}
|
||||
OVI_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
try:
|
||||
from .steadydancer.nodes import NODE_CLASS_MAPPINGS as STEADYDANCER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except Exception as e:
|
||||
log.warning(f"WanVideoWrapper WARNING: SteadyDancer nodes not available due to error in importing them: {e}")
|
||||
STEADYDANCER_NODE_CLASS_MAPPINGS = {}
|
||||
STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
|
||||
@@ -100,6 +107,7 @@ NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
|
||||
NODE_CLASS_MAPPINGS.update(STEADYDANCER_NODE_CLASS_MAPPINGS)
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
@@ -124,5 +132,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
NODE_DISPLAY_NAME_MAPPINGS.update(STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
+17
-1
@@ -872,6 +872,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
|
||||
key = name.replace("_orig_mod.", "")
|
||||
value=sd[key]
|
||||
keep_fp32 = ["patch_embedding", "motion_encoder", "condition_embedding"]
|
||||
|
||||
if gguf:
|
||||
dtype_to_use = torch.float32 if "patch_embedding" in name or "motion_encoder" in name else base_dtype
|
||||
@@ -883,7 +884,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
dtype_to_use = value.dtype
|
||||
if "bias" in name or "img_emb" in name:
|
||||
dtype_to_use = base_dtype
|
||||
if "patch_embedding" in name or "motion_encoder" in name:
|
||||
if any(k in name for k in keep_fp32):
|
||||
dtype_to_use = torch.float32
|
||||
if "modulation" in name or "norm" in name:
|
||||
dtype_to_use = value.dtype if value.dtype == torch.float32 else base_dtype
|
||||
@@ -1499,6 +1500,21 @@ class WanVideoModelLoader:
|
||||
device=device,
|
||||
)
|
||||
|
||||
# SteadyDancer
|
||||
if "condition_embedding_align.cross_attn.in_proj_bias" in sd:
|
||||
from .steadydancer.mobilenetv2_dcd import DYModule
|
||||
from .steadydancer.small_archs import PoseRefNetNoBNV3, FactorConv3d
|
||||
in_dim_c = 16
|
||||
transformer.patch_embedding_fuse = nn.Conv3d(in_channels + in_dim_c + in_dim_c, dim, kernel_size=patch_size, stride=patch_size) # x, fused pose, aligned pose
|
||||
transformer.patch_embedding_ref_c = nn.Conv3d(in_dim_c, dim, kernel_size=patch_size, stride=patch_size) # ref_c
|
||||
transformer.condition_embedding_spatial = DYModule(inp=in_dim_c, oup=in_dim_c) # Spatial Structure Adaptive Extractor
|
||||
transformer.condition_embedding_temporal = nn.Sequential( # Temporal Motion Coherence Module
|
||||
FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU(),
|
||||
FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU(),
|
||||
FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU())
|
||||
transformer.condition_embedding_align = PoseRefNetNoBNV3(in_channels_x=16, in_channels_c=16, hidden_dim=128, num_heads=8) # Frame-wise Attention Alignment Unit
|
||||
|
||||
|
||||
comfy_model.diffusion_model = transformer
|
||||
comfy_model.load_device = transformer_load_device
|
||||
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
|
||||
|
||||
@@ -1171,6 +1171,15 @@ class WanVideoSampler:
|
||||
|
||||
latent = add_noise(ttm_reference_latents, noise, timesteps[ttm_start_step].to(noise.device)).to(latent)
|
||||
|
||||
# SteadyDancer
|
||||
sdance_embeds = image_embeds.get("sdance_embeds", None)
|
||||
sdancer_input = None
|
||||
if sdance_embeds is not None:
|
||||
print("Using SteadyDancer embeddings")
|
||||
print(f"SteadyDancer embeds keys: {list(sdance_embeds.keys())}")
|
||||
sdancer_input = sdance_embeds.copy()
|
||||
sdancer_input = dict_to_device(sdancer_input, device, dtype)
|
||||
|
||||
#region model pred
|
||||
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
|
||||
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
|
||||
@@ -1450,6 +1459,7 @@ class WanVideoSampler:
|
||||
"flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling
|
||||
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
|
||||
"num_cond_latents": len(all_indices) if transformer.is_longcat else None,
|
||||
"sdancer_input": sdancer_input, # SteadyDancer input
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
# Modify from https://github.com/liyunsheng13/dcd/blob/main/models/imagenet/mobilenetv2_dcd.py
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
|
||||
class Hsigmoid(nn.Module):
|
||||
def __init__(self, inplace=True):
|
||||
super(Hsigmoid, self).__init__()
|
||||
self.inplace = inplace
|
||||
|
||||
def forward(self, x):
|
||||
return F.relu6(x + 3., inplace=self.inplace) / 3.
|
||||
|
||||
|
||||
class DYModule(nn.Module):
|
||||
def __init__(self, inp, oup, fc_squeeze=8):
|
||||
super(DYModule, self).__init__()
|
||||
self.conv = nn.Conv2d(inp, oup, 1, 1, 0, bias=False)
|
||||
if inp < oup:
|
||||
self.mul = 4
|
||||
reduction = 8
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(2)
|
||||
else:
|
||||
self.mul = 1
|
||||
reduction = 2
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
|
||||
self.dim = min((inp * self.mul) // reduction, oup // reduction)
|
||||
while self.dim ** 2 > inp * self.mul * 2:
|
||||
reduction *= 2
|
||||
self.dim = min((inp * self.mul) // reduction, oup // reduction)
|
||||
if self.dim < 4:
|
||||
self.dim = 4
|
||||
|
||||
squeeze = max(inp * self.mul, self.dim ** 2) // fc_squeeze
|
||||
if squeeze < 4:
|
||||
squeeze = 4
|
||||
self.conv_q = nn.Conv2d(inp, self.dim, 1, 1, 0, bias=False)
|
||||
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(inp * self.mul, squeeze, bias=False),
|
||||
SEModule_small(squeeze),
|
||||
)
|
||||
self.fc_phi = nn.Linear(squeeze, self.dim ** 2, bias=False)
|
||||
self.fc_scale = nn.Linear(squeeze, oup, bias=False)
|
||||
self.hs = Hsigmoid()
|
||||
self.conv_p = nn.Conv2d(self.dim, oup, 1, 1, 0, bias=False)
|
||||
# self.bn1 = nn.BatchNorm2d(self.dim)
|
||||
self.bn1 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
|
||||
# self.bn2 = nn.BatchNorm1d(self.dim)
|
||||
self.bn2 = nn.GroupNorm(num_groups=4, num_channels=self.dim)
|
||||
|
||||
def forward(self, x):
|
||||
r = self.conv(x)
|
||||
|
||||
b, c, h, w = x.size()
|
||||
y = self.avg_pool(x).view(b, c * self.mul)
|
||||
y = self.fc(y)
|
||||
dy_phi = self.fc_phi(y).view(b, self.dim, self.dim)
|
||||
dy_scale = self.hs(self.fc_scale(y)).view(b, -1, 1, 1)
|
||||
r = dy_scale.expand_as(r) * r
|
||||
|
||||
x = self.conv_q(x)
|
||||
x = self.bn1(x)
|
||||
x = x.view(b, -1, h * w)
|
||||
x = self.bn2(torch.matmul(dy_phi, x)) + x
|
||||
x = x.view(b, -1, h, w)
|
||||
x = self.conv_p(x)
|
||||
return x + r
|
||||
|
||||
|
||||
class SEModule_small(nn.Module):
|
||||
def __init__(self, channel):
|
||||
super(SEModule_small, self).__init__()
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(channel, channel, bias=False),
|
||||
Hsigmoid()
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
y = self.fc(x)
|
||||
return x * y
|
||||
|
||||
|
||||
class SEModule(nn.Module):
|
||||
def __init__(self, channel, reduction=4):
|
||||
super(SEModule, self).__init__()
|
||||
self.avg_pool = nn.AdaptiveAvgPool2d(1)
|
||||
self.fc = nn.Sequential(
|
||||
nn.Linear(channel, channel // reduction, bias=False),
|
||||
nn.ReLU(inplace=True),
|
||||
nn.Linear(channel // reduction, channel, bias=False),
|
||||
Hsigmoid()
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
b, c, _, _ = x.size()
|
||||
y = self.avg_pool(x).view(b, c)
|
||||
y = self.fc(y).view(b, c, 1, 1)
|
||||
return x * y.expand_as(x)
|
||||
@@ -0,0 +1,60 @@
|
||||
import os
|
||||
import torch
|
||||
import numpy as np
|
||||
from ..utils import log
|
||||
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import load_torch_file, ProgressBar
|
||||
import folder_paths
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
|
||||
class WanVideoAddSteadyDancerEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||
"pose_latents_positive": ("LATENT",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the portrait 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": {
|
||||
"pose_latents_negative": ("LATENT",),
|
||||
"clip_vision_embeds": ("WANVIDIMAGE_CLIPEMBEDS",),
|
||||
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, embeds, pose_latents_positive, strength, start_percent=0.0, end_percent=1.0, pose_latents_negative=None, clip_vision_embeds=None):
|
||||
sdance_embeds = {
|
||||
"cond_pos": pose_latents_positive["samples"][0],
|
||||
"cond_neg": pose_latents_negative["samples"][0] if pose_latents_negative else None,
|
||||
"strength": strength,
|
||||
"start_percent": start_percent,
|
||||
"end_percent": end_percent,
|
||||
"clip_fea": clip_vision_embeds,
|
||||
}
|
||||
|
||||
updated = dict(embeds)
|
||||
updated["sdance_embeds"] = sdance_embeds
|
||||
return (updated,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoAddSteadyDancerEmbeds": WanVideoAddSteadyDancerEmbeds,
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoAddSteadyDancerEmbeds": "WanVideo Add SteadyDancer Embeds",
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class FactorConv3d(nn.Module):
|
||||
"""
|
||||
(2+1)D decomposition of 3D convolution: 1xHxW spatial convolution → Swish → Tx1x1 temporal convolution
|
||||
"""
|
||||
def __init__(self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
kernel_size,
|
||||
stride: int = 1,
|
||||
dilation: int = 1):
|
||||
super().__init__()
|
||||
|
||||
if isinstance(kernel_size, int):
|
||||
k_t, k_h, k_w = kernel_size, kernel_size, kernel_size
|
||||
else:
|
||||
k_t, k_h, k_w = kernel_size
|
||||
|
||||
pad_t = (k_t - 1) * dilation // 2
|
||||
pad_hw = (k_h - 1) * dilation // 2
|
||||
|
||||
self.spatial = nn.Conv3d(
|
||||
in_channels, in_channels,
|
||||
kernel_size=(1, k_h, k_w),
|
||||
stride=(1, stride, stride),
|
||||
padding=(0, pad_hw, pad_hw),
|
||||
dilation=(1, dilation, dilation),
|
||||
groups=in_channels,
|
||||
bias=False
|
||||
)
|
||||
|
||||
self.temporal = nn.Conv3d(
|
||||
in_channels, out_channels,
|
||||
kernel_size=(k_t, 1, 1),
|
||||
stride=(stride, 1, 1),
|
||||
padding=(pad_t, 0, 0),
|
||||
dilation=(dilation, 1, 1),
|
||||
bias=True
|
||||
)
|
||||
|
||||
self.act = nn.SiLU()
|
||||
|
||||
def forward(self, x):
|
||||
x = self.spatial(x)
|
||||
x = self.act(x)
|
||||
x = self.temporal(x)
|
||||
return x
|
||||
|
||||
|
||||
class LayerNorm2D(nn.Module):
|
||||
"""
|
||||
LayerNorm over C for a 4-D tensor (B, C, H, W)
|
||||
"""
|
||||
def __init__(self, num_channels, eps=1e-5, affine=True):
|
||||
super().__init__()
|
||||
self.num_channels = num_channels
|
||||
self.eps = eps
|
||||
self.affine = affine
|
||||
if affine:
|
||||
self.weight = nn.Parameter(torch.ones(1, num_channels, 1, 1))
|
||||
self.bias = nn.Parameter(torch.zeros(1, num_channels, 1, 1))
|
||||
|
||||
def forward(self, x):
|
||||
# x: (B, C, H, W)
|
||||
mean = x.mean(dim=1, keepdim=True) # (B, 1, H, W)
|
||||
var = x.var (dim=1, keepdim=True, unbiased=False)
|
||||
x = (x - mean) / torch.sqrt(var + self.eps)
|
||||
if self.affine:
|
||||
x = x * self.weight + self.bias
|
||||
return x
|
||||
|
||||
|
||||
class PoseRefNetNoBNV3(nn.Module):
|
||||
def __init__(self,
|
||||
in_channels_c: int,
|
||||
in_channels_x: int,
|
||||
hidden_dim: int = 256,
|
||||
num_heads: int = 8,
|
||||
dropout: float = 0.1):
|
||||
super().__init__()
|
||||
self.d_model = hidden_dim
|
||||
self.nhead = num_heads
|
||||
|
||||
self.proj_p = nn.Conv2d(in_channels_c, hidden_dim, kernel_size=1)
|
||||
self.proj_r = nn.Conv2d(in_channels_x, hidden_dim, kernel_size=1)
|
||||
|
||||
self.proj_p_back = nn.Conv2d(hidden_dim, in_channels_c, kernel_size=1)
|
||||
|
||||
self.cross_attn = nn.MultiheadAttention(hidden_dim,
|
||||
num_heads=num_heads,
|
||||
dropout=dropout)
|
||||
|
||||
self.ffn_pose = nn.Sequential(
|
||||
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1)
|
||||
)
|
||||
|
||||
self.norm1 = LayerNorm2D(hidden_dim)
|
||||
self.norm2 = LayerNorm2D(hidden_dim)
|
||||
|
||||
def forward(self, pose, ref, mask=None):
|
||||
"""
|
||||
pose : (B, C1, T, H, W)
|
||||
ref : (B, C2, T, H, W)
|
||||
mask : (B, T*H*W) optional key_padding_mask
|
||||
return: (B, d_model, T, H, W)
|
||||
"""
|
||||
B, _, T, H, W = pose.shape
|
||||
L = H * W
|
||||
|
||||
p_trans = pose.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
|
||||
r_trans = ref.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1)
|
||||
|
||||
p_trans = self.proj_p(p_trans)
|
||||
r_trans = self.proj_r(r_trans)
|
||||
|
||||
p_trans = p_trans.flatten(2).transpose(1, 2)
|
||||
r_trans = r_trans.flatten(2).transpose(1, 2)
|
||||
|
||||
out = self.cross_attn(query=r_trans,
|
||||
key=p_trans,
|
||||
value=p_trans,
|
||||
key_padding_mask=mask)[0]
|
||||
|
||||
out = out.transpose(1, 2).contiguous().view(B*T, -1, H, W)
|
||||
out = self.norm1(out)
|
||||
|
||||
ffn_out = self.ffn_pose(out)
|
||||
out = out + ffn_out
|
||||
out = self.norm2(out)
|
||||
out = self.proj_p_back(out)
|
||||
out = out.view(B, T, -1, H, W).contiguous().transpose(1, 2)
|
||||
|
||||
return out
|
||||
+63
-61
@@ -1623,55 +1623,22 @@ class AudioInjector_WAN(nn.Module):
|
||||
class WanModel(torch.nn.Module):
|
||||
def __init__(self,
|
||||
model_type='t2v',
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=512,
|
||||
in_dim=16,
|
||||
dim=2048,
|
||||
in_features=5120,
|
||||
out_features=5120,
|
||||
ffn_dim=8192,
|
||||
ffn2_dim=8192,
|
||||
freq_dim=256,
|
||||
text_dim=4096,
|
||||
out_dim=16,
|
||||
num_heads=16,
|
||||
num_layers=32,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func='comfy',
|
||||
rms_norm_function='default',
|
||||
main_device=torch.device('cuda'),
|
||||
offload_device=torch.device('cpu'),
|
||||
dtype=torch.float16,
|
||||
teacache_coefficients=[],
|
||||
magcache_ratios=[],
|
||||
vace_layers=None,
|
||||
vace_in_dim=None,
|
||||
inject_sample_info=False,
|
||||
add_ref_conv=False,
|
||||
in_dim_ref_conv=16,
|
||||
add_control_adapter=False,
|
||||
in_dim_control_adapter=24,
|
||||
use_motion_attn=False,
|
||||
patch_size=(1, 2, 2), text_len=512,
|
||||
in_dim=16, dim=2048, in_features=5120, out_features=5120, ffn_dim=8192, ffn2_dim=8192,
|
||||
freq_dim=256, text_dim=4096, out_dim=16, 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,
|
||||
teacache_coefficients=[], magcache_ratios=[], vace_layers=None, vace_in_dim=None,
|
||||
inject_sample_info=False, add_ref_conv=False, in_dim_ref_conv=16, add_control_adapter=False,
|
||||
in_dim_control_adapter=24, use_motion_attn=False,
|
||||
#s2v
|
||||
cond_dim=0,
|
||||
audio_dim=1024,
|
||||
num_audio_token=4,
|
||||
enable_adain=False,
|
||||
adain_mode="attn_norm",
|
||||
audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
|
||||
zero_timestep=False,
|
||||
humo_audio=False,
|
||||
cond_dim=0, audio_dim=1024, num_audio_token=4, enable_adain=False, zero_timestep=False, humo_audio=False,
|
||||
adain_mode="attn_norm", audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39],
|
||||
# WanAnimate
|
||||
is_wananimate=False,
|
||||
motion_encoder_dim=512,
|
||||
is_wananimate=False, motion_encoder_dim=512,
|
||||
# lynx
|
||||
lynx_ip_layers=None,
|
||||
lynx_ref_layers=None,
|
||||
# ovi
|
||||
is_ovi_audio_model=False,
|
||||
lynx_ip_layers=None, lynx_ref_layers=None,
|
||||
# LongCat
|
||||
is_longcat=False,
|
||||
):
|
||||
@@ -2190,8 +2157,7 @@ class WanModel(torch.nn.Module):
|
||||
self, x, t, context, seq_len,
|
||||
is_uncond=False,
|
||||
current_step_percentage=0.0, current_step=0, last_step=0, total_steps=50,
|
||||
clip_fea=None,
|
||||
y=None,
|
||||
clip_fea=None, y=None,
|
||||
device=torch.device('cuda'),
|
||||
freqs=None,
|
||||
enhance_enabled=False,
|
||||
@@ -2203,8 +2169,7 @@ class WanModel(torch.nn.Module):
|
||||
fps_embeds=None,
|
||||
fun_ref=None, fun_camera=None,
|
||||
audio_proj=None, audio_scale=1.0,
|
||||
uni3c_data=None,
|
||||
controlnet=None,
|
||||
uni3c_data=None, controlnet=None,
|
||||
add_cond=None, attn_cond=None,
|
||||
nag_params={}, nag_context=None,
|
||||
multitalk_audio=None,
|
||||
@@ -2227,6 +2192,7 @@ class WanModel(torch.nn.Module):
|
||||
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
|
||||
num_cond_latents=None,
|
||||
add_text_emb=None,
|
||||
sdancer_input=None # SteadyDancer
|
||||
):
|
||||
r"""
|
||||
Forward pass through the diffusion model
|
||||
@@ -2314,6 +2280,10 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
_, F, H, W = x[0].shape
|
||||
|
||||
if sdancer_input is not None:
|
||||
x_noise_clone = torch.stack(x)
|
||||
|
||||
# I2V
|
||||
if y is not None:
|
||||
if hasattr(self, "randomref_embedding_pose") and unianim_data is not None:
|
||||
if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']:
|
||||
@@ -2330,13 +2300,27 @@ class WanModel(torch.nn.Module):
|
||||
hidden_states = torch.cat([hidden_states, torch.zeros_like(hidden_states[:, :4])], dim=1)
|
||||
render_latent = torch.cat([hidden_states[:, :20], render_latent], dim=1)
|
||||
|
||||
# patch embed
|
||||
if control_lora_enabled:
|
||||
self.expanded_patch_embedding.to(self.main_device)
|
||||
x = [self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x]
|
||||
# SteadyDancer
|
||||
if sdancer_input is not None:
|
||||
sdancer_cond = sdancer_input["cond_pos"] if not is_uncond else sdancer_input["cond_neg"]
|
||||
condition_temporal = [self.condition_embedding_temporal(c.unsqueeze(0).float()).to(self.base_dtype) for c in [sdancer_cond]] # Temporal Motion Coherence Module.
|
||||
sdancer_cond = sdancer_cond.unsqueeze(0)
|
||||
bs, _, time_steps, _, _ = sdancer_cond.shape
|
||||
condition_reshape = rearrange(sdancer_cond, 'b c t h w -> (b t) c h w')
|
||||
condition_spatial = self.condition_embedding_spatial(condition_reshape.float()).to(self.base_dtype) # Spatial Structure Adaptive Extractor.
|
||||
condition_spatial = rearrange(condition_spatial, '(b t) c h w -> b c t h w', t=time_steps, b=bs)
|
||||
condition_fused = sdancer_cond + condition_temporal[0] + condition_spatial # Hierarchical Aggregation (1): condition, temporal condition, spatial condition
|
||||
condition_aligned = self.condition_embedding_align(condition_fused.float(), x_noise_clone).to(self.base_dtype) # Frame-wise Attention Alignment Unit.
|
||||
else:
|
||||
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]
|
||||
# patch embed
|
||||
if control_lora_enabled:
|
||||
self.expanded_patch_embedding.to(self.main_device)
|
||||
x = [self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x]
|
||||
else:
|
||||
self.original_patch_embedding.to(self.main_device)
|
||||
x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x]
|
||||
|
||||
orig_frames = x[0].shape[1]
|
||||
|
||||
# ovi audio model
|
||||
if self.audio_model is not None:
|
||||
@@ -2365,13 +2349,29 @@ class WanModel(torch.nn.Module):
|
||||
fun_camera = self.control_adapter(fun_camera)
|
||||
x = [u + v for u, v in zip(x, fun_camera)]
|
||||
|
||||
# SteadyDancer
|
||||
if sdancer_input is not None:
|
||||
ref_x = y[0][4:, :1] # reuse I2V input as reference, slice mask off
|
||||
msk = torch.ones(4, 1, H, W, device=ref_x.device) # new mask goes in middle
|
||||
ref_x = [torch.concat([ref_x, msk, ref_x])]
|
||||
ref_c = sdancer_cond[0][:, :1]
|
||||
ref_c = [torch.concat([ref_c, msk * 0, ref_c])] # zero mask for cond ref
|
||||
# Condition Fusion/Injection, Hierarchical Aggregation (2): x, fused condition, aligned condition
|
||||
x = [self.patch_embedding_fuse(torch.cat([u[None], c[None], a[None]], 1)) for u, c, a in zip(x, condition_fused, condition_aligned)]
|
||||
# Condition Augmentation: x_cond, ref_x, ref_c
|
||||
ref_x = [self.patch_embedding(r.unsqueeze(0).float()).to(self.base_dtype) for r in ref_x]
|
||||
ref_c = [self.patch_embedding_ref_c(r[:16].unsqueeze(0).float()).to(self.base_dtype) for r in ref_c]
|
||||
F += ref_x[0].shape[2] + ref_c[0].shape[2] # update frame count for rope
|
||||
x = [torch.cat([r, u, v], dim=2) for r, u, v in zip(x, ref_x, ref_c)]
|
||||
seq_len = torch.tensor([u.flatten(2).transpose(1, 2).size(1) for u in x], dtype=torch.int32).max() # update seq len
|
||||
|
||||
# grid sizes and seq len
|
||||
grid_sizes = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x])
|
||||
original_grid_sizes = grid_sizes.clone()
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
self.original_seq_len = x[0].shape[1]
|
||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
|
||||
assert seq_lens.max() <= seq_len
|
||||
assert seq_lens.max() <= seq_len, f"max seq len {seq_lens.max()} exceeds provided seq_len {seq_len}"
|
||||
|
||||
cond_mask_weight = None
|
||||
if self.trainable_cond_mask is not None:
|
||||
@@ -2570,11 +2570,13 @@ class WanModel(torch.nn.Module):
|
||||
# clip vision embedding
|
||||
clip_embed = None
|
||||
if clip_fea is not None and hasattr(self, "img_emb"):
|
||||
clip_fea = clip_fea.to(self.main_device)
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.main_device)
|
||||
clip_embed = self.img_emb(clip_fea) # bs x 257 x dim
|
||||
#context = torch.concat([context_clip, context], dim=1)
|
||||
clip_embed = self.img_emb(clip_fea.to(self.main_device)) # bs x 257 x dim
|
||||
if sdancer_input is not None:
|
||||
clip_fea_c = sdancer_input.get("clip_fea_c", None)
|
||||
if clip_fea_c is not None:
|
||||
clip_embed += self.img_emb(clip_fea_c.to(self.main_device))
|
||||
if self.offload_img_emb:
|
||||
self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking)
|
||||
|
||||
@@ -3028,7 +3030,7 @@ class WanModel(torch.nn.Module):
|
||||
x_ovi = [u.float() for u in x_ovi]
|
||||
|
||||
x = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type]
|
||||
x = [u.float() for u in x]
|
||||
x = [u[:, :orig_frames, ...].float() for u in x]
|
||||
return (x, x_ovi, pred_id) if pred_id is not None else (x, x_ovi, None)
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
|
||||
Reference in New Issue
Block a user