realisdance
This commit is contained in:
@@ -680,6 +680,19 @@ class WanVideoModelLoader:
|
||||
block.cross_attn.v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
sd.update(fantasytalking_model["sd"])
|
||||
|
||||
# RealisDance-DiT
|
||||
if "add_conv_in.weight" in sd:
|
||||
def zero_module(module):
|
||||
for p in module.parameters():
|
||||
torch.nn.init.zeros_(p)
|
||||
return module
|
||||
inner_dim = sd["add_conv_in.weight"].shape[0]
|
||||
add_cond_in_dim = sd["add_conv_in.weight"].shape[1]
|
||||
attn_cond_in_dim = sd["attn_conv_in.weight"].shape[1]
|
||||
transformer.add_conv_in = torch.nn.Conv3d(add_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
|
||||
transformer.add_proj = zero_module(torch.nn.Linear(inner_dim, inner_dim))
|
||||
transformer.attn_conv_in = torch.nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
|
||||
|
||||
comfy_model = WanVideoModel(
|
||||
WanVideoModelConfig(base_dtype),
|
||||
model_type=comfy.model_base.ModelType.FLOW,
|
||||
@@ -702,7 +715,7 @@ class WanVideoModelLoader:
|
||||
dtype = torch.float8_e5m2
|
||||
else:
|
||||
dtype = base_dtype
|
||||
params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter"}
|
||||
params_to_keep = {"norm", "head", "bias", "time_in", "vector_in", "patch_embedding", "time_", "img_emb", "modulation", "text_embedding", "adapter", "add"}
|
||||
#if lora is not None:
|
||||
# transformer_load_device = device
|
||||
if not lora_low_mem_load:
|
||||
@@ -1553,6 +1566,38 @@ class WanVideoClipVisionEncode:
|
||||
}
|
||||
|
||||
return (clip_embeds_dict,)
|
||||
|
||||
class WanVideoRealisDanceLatents:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"ref_latent": ("LATENT", {"tooltip": "Reference image to encode"}),
|
||||
"smpl_latent": ("LATENT", {"tooltip": "SMPL pose image to encode"}),
|
||||
},
|
||||
"optional": {
|
||||
"hamer_latent": ("LATENT", {"tooltip": "Hamer hand pose image to encode"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("REALISDANCELATENTS",)
|
||||
RETURN_NAMES = ("realisdance_latents",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, ref_latent, smpl_latent, hamer_latent=None):
|
||||
if hamer_latent is None:
|
||||
hamer = torch.zeros_like(smpl_latent["samples"])
|
||||
else:
|
||||
hamer = hamer_latent["samples"]
|
||||
|
||||
pose_latent = torch.cat((smpl_latent["samples"], hamer), dim=1)
|
||||
|
||||
realisdance_latents = {
|
||||
"ref_latent": ref_latent["samples"],
|
||||
"pose_latent": pose_latent,
|
||||
}
|
||||
|
||||
return (realisdance_latents,)
|
||||
|
||||
class WanVideoImageToVideoEncode:
|
||||
@classmethod
|
||||
@@ -1576,7 +1621,7 @@ class WanVideoImageToVideoEncode:
|
||||
"temporal_mask": ("MASK", {"tooltip": "mask"}),
|
||||
"extra_latents": ("LATENT", {"tooltip": "Extra latents to add to the input front, used for Skyreels A2 reference images"}),
|
||||
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
|
||||
|
||||
"realisdance_latents": ("REALISDANCELATENTS", {"tooltip": "RealisDance latents"}),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1587,7 +1632,7 @@ class WanVideoImageToVideoEncode:
|
||||
|
||||
def process(self, vae, width, height, num_frames, force_offload, noise_aug_strength,
|
||||
start_latent_strength, end_latent_strength, start_image=None, end_image=None, control_embeds=None, fun_or_fl2v_model=False,
|
||||
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False):
|
||||
temporal_mask=None, extra_latents=None, clip_embeds=None, tiled_vae=False, realisdance_latents=None):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
@@ -1690,6 +1735,9 @@ class WanVideoImageToVideoEncode:
|
||||
frames_per_stride = (num_frames - 1) // 4 + (2 if end_image is not None and not fun_or_fl2v_model else 1)
|
||||
max_seq_len = frames_per_stride * patches_per_frame
|
||||
|
||||
if realisdance_latents is not None:
|
||||
realisdance_latents["ref_latent_neg"] = vae.encode(torch.zeros(1, 3, 1, H, W, device=device, dtype=vae.dtype), device)
|
||||
|
||||
vae.model.clear_cache()
|
||||
if force_offload:
|
||||
vae.model.to(offload_device)
|
||||
@@ -1708,6 +1756,7 @@ class WanVideoImageToVideoEncode:
|
||||
"end_image": resized_end_image if end_image is not None else None,
|
||||
"fun_or_fl2v_model": fun_or_fl2v_model,
|
||||
"has_ref": has_ref,
|
||||
"realisdance_latents": realisdance_latents
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
@@ -2416,6 +2465,7 @@ class WanVideoSampler:
|
||||
|
||||
image_cond = image_embeds.get("image_embeds", None)
|
||||
ATI_tracks = None
|
||||
add_cond = attn_cond = attn_cond_neg = None
|
||||
|
||||
if image_cond is not None:
|
||||
log.info(f"image_cond shape: {image_cond.shape}")
|
||||
@@ -2430,6 +2480,12 @@ class WanVideoSampler:
|
||||
ati_end_percent = transformer_options.get("ati_end_percentage", 1.0)
|
||||
image_cond_ati = patch_motion(ATI_tracks.to(image_cond.device, image_cond.dtype), image_cond, topk=topk, temperature=temperature)
|
||||
log.info(f"ATI tracks shape: {ATI_tracks.shape}")
|
||||
|
||||
realisdance_latents = image_embeds.get("realisdance_latents", None)
|
||||
if realisdance_latents is not None:
|
||||
add_cond = realisdance_latents["pose_latent"]
|
||||
attn_cond = realisdance_latents["ref_latent"]
|
||||
attn_cond_neg = realisdance_latents["ref_latent_neg"]
|
||||
|
||||
end_image = image_embeds.get("end_image", None)
|
||||
lat_h = image_embeds.get("lat_h", None)
|
||||
@@ -2973,7 +3029,8 @@ class WanVideoSampler:
|
||||
'audio_context_lens': audio_context_lens if fantasytalking_embeds is not None else None,
|
||||
'audio_scale': audio_scale if fantasytalking_embeds is not None else None,
|
||||
"pcd_data": pcd_data,
|
||||
"controlnet": controlnet
|
||||
"controlnet": controlnet,
|
||||
"add_cond": add_cond,
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -2987,7 +3044,7 @@ class WanVideoSampler:
|
||||
[z_pos], context=positive_embeds, y=[image_cond_input] if image_cond_input is not None else None,
|
||||
clip_fea=clip_fea, is_uncond=False, current_step_percentage=current_step_percentage,
|
||||
pred_id=teacache_state[0] if teacache_state else None,
|
||||
vace_data=vace_data,
|
||||
vace_data=vace_data, attn_cond=attn_cond,
|
||||
**base_params
|
||||
)
|
||||
noise_pred_cond = noise_pred_cond[0].to(intermediate_device)
|
||||
@@ -3009,7 +3066,7 @@ class WanVideoSampler:
|
||||
y=[image_cond_input] if image_cond_input is not None else None,
|
||||
is_uncond=True, current_step_percentage=current_step_percentage,
|
||||
pred_id=teacache_state[1] if teacache_state else None,
|
||||
vace_data=vace_data,
|
||||
vace_data=vace_data, attn_cond=attn_cond_neg,
|
||||
**base_params
|
||||
)
|
||||
noise_pred_uncond = noise_pred_uncond[0].to(intermediate_device)
|
||||
@@ -3673,7 +3730,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoVACEStartToEndFrame": WanVideoVACEStartToEndFrame,
|
||||
"WanVideoVACEModelSelect": WanVideoVACEModelSelect,
|
||||
"WanVideoPhantomEmbeds": WanVideoPhantomEmbeds,
|
||||
"CreateCFGScheduleFloatList": CreateCFGScheduleFloatList
|
||||
"CreateCFGScheduleFloatList": CreateCFGScheduleFloatList,
|
||||
"WanVideoRealisDanceLatents": WanVideoRealisDanceLatents
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoSampler": "WanVideo Sampler",
|
||||
@@ -3710,5 +3768,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoVACEStartToEndFrame": "WanVideo VACE Start To End Frame",
|
||||
"WanVideoVACEModelSelect": "WanVideo VACE Model Select",
|
||||
"WanVideoPhantomEmbeds": "WanVideo Phantom Embeds",
|
||||
"CreateCFGScheduleFloatList": "WanVideo CFG Schedule Float List"
|
||||
"CreateCFGScheduleFloatList": "WanVideo CFG Schedule Float List",
|
||||
"WanVideoRealisDanceLatents": "WanVideo RealisDance Latents",
|
||||
}
|
||||
|
||||
@@ -1163,6 +1163,8 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
audio_scale=1.0,
|
||||
pcd_data=None,
|
||||
controlnet=None,
|
||||
add_cond=None,
|
||||
attn_cond=None
|
||||
|
||||
):
|
||||
r"""
|
||||
@@ -1237,6 +1239,19 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
|
||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||
|
||||
if add_cond is not None:
|
||||
add_cond = self.add_conv_in(add_cond.to(self.add_conv_in.weight.dtype)).to(x[0].dtype)
|
||||
add_cond = add_cond.flatten(2).transpose(1, 2)
|
||||
x[0] = x[0] + self.add_proj(add_cond)
|
||||
if attn_cond is not None:
|
||||
F_cond, H_cond, W_cond = attn_cond.shape[2], attn_cond.shape[3], attn_cond.shape[4]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
attn_cond = self.attn_conv_in(attn_cond.to(self.attn_conv_in.weight.dtype)).to(x[0].dtype)
|
||||
attn_cond = attn_cond.flatten(2).transpose(1, 2)
|
||||
x_len = x[0].shape[1]
|
||||
x[0] = torch.cat([x[0], attn_cond], dim=1)
|
||||
seq_len += attn_cond.size(1)
|
||||
|
||||
if self.ref_conv is not None and fun_ref is not None:
|
||||
fun_ref = self.ref_conv(fun_ref).flatten(2).transpose(1, 2)
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
@@ -1260,18 +1275,37 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(0, f_len - 1, steps=f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(0, h_len - 1, steps=h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(0, w_len - 1, steps=w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
|
||||
freqs = self.rope_embedder(img_ids).movedim(1, 2)
|
||||
if attn_cond is not None:
|
||||
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=x.device, dtype=x.dtype)
|
||||
|
||||
#shift
|
||||
shift_f_size = 81 # Default value
|
||||
shift_f = False
|
||||
if shift_f:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
else:
|
||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=x.device, dtype=x.dtype).reshape(-1, 1, 1)
|
||||
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=x.device, dtype=x.dtype).reshape(1, -1, 1)
|
||||
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=x.device, dtype=x.dtype).reshape(1, 1, -1)
|
||||
|
||||
# Combine original and conditional position ids
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
|
||||
|
||||
# Generate RoPE frequencies for the combined positions
|
||||
freqs = self.rope_embedder(combined_img_ids).movedim(1, 2)
|
||||
else:
|
||||
img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
||||
freqs = self.rope_embedder(img_ids).movedim(1, 2)
|
||||
else:
|
||||
rope_func = "default"
|
||||
|
||||
# time embeddings
|
||||
|
||||
# e = self.time_embedding(
|
||||
# sinusoidal_embedding_1d(self.freq_dim, t).float())
|
||||
# e0 = self.time_projection(e).unflatten(1, (6, self.dim))
|
||||
# assert e.dtype == torch.float32 and e0.dtype == torch.float32
|
||||
if t.dim() == 2:
|
||||
b, f = t.shape
|
||||
_flag_df = True
|
||||
@@ -1468,6 +1502,10 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
x = x[:, full_ref_length:]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
if attn_cond is not None:
|
||||
x = x[:, :x_len]
|
||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||
|
||||
x = self.head(x, e.to(x.device))
|
||||
x = self.unpatchify(x, grid_sizes) # type: ignore[arg-type]
|
||||
x = [u.float() for u in x]
|
||||
|
||||
Reference in New Issue
Block a user