This commit is contained in:
kijai
2025-11-28 20:32:16 +02:00
parent 772642b4f1
commit e54fa5d059
7 changed files with 400 additions and 63 deletions
+9
View File
@@ -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
View File
@@ -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)
+10
View File
@@ -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
+102
View File
@@ -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)
+60
View File
@@ -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",
}
+138
View File
@@ -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
+64 -62
View File
@@ -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
@@ -2313,7 +2279,11 @@ class WanModel(torch.nn.Module):
freqs = freqs.to(device)
_, 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']:
@@ -2321,7 +2291,7 @@ class WanModel(torch.nn.Module):
if random_ref_emb is not None:
y[0].add_(random_ref_emb, alpha=unianim_data["strength"])
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
#uni3c controlnet
if uni3c_data is not None:
render_latent = uni3c_data["render_latent"].to(self.base_dtype)
@@ -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):