Expose more options, make vid2vid easier
This commit is contained in:
+47
-34
@@ -1,38 +1,42 @@
|
||||
import torch
|
||||
from ..utils import log
|
||||
import comfy.model_management as mm
|
||||
from comfy_api.latest import io
|
||||
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
|
||||
class WanVideoLongCatAvatarExtendEmbeds:
|
||||
class WanVideoLongCatAvatarExtendEmbeds(io.ComfyNode):
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"prev_latents": ("LATENT", {"tooltip": "Previous latents to be used to continue generation"}),
|
||||
"audio_embeds": ("MULTITALK_EMBEDS", {"tooltip": "Full length audio embeddings"}),
|
||||
"num_frames": ("INT", {"default": 93, "min": 1, "max": 256, "step": 1, "tooltip": "Number of new frames to generate" }),
|
||||
"overlap": ("INT", {"default": 13, "min": 0, "max": 16, "step": 1, "tooltip": "Number of overlapping frames from previous latents" }),
|
||||
"frames_processed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Number of frames already processed in the video" }),
|
||||
"if_not_enough_audio": (["pad_with_start", "mirror_from_end"], {"default": "pad_with_start", "tooltip": "What to do if there are not enough frames in pose_images for the window"}),
|
||||
},
|
||||
"optional": {
|
||||
"ref_latent": ("LATENT", {"default": None, "tooltip": "Reference latent for the first frame (used for consistency)"}),
|
||||
}
|
||||
}
|
||||
def define_schema(cls):
|
||||
return io.Schema(
|
||||
node_id="WanVideoLongCatAvatarExtendEmbeds",
|
||||
category="WanVideoWrapper",
|
||||
inputs=[
|
||||
io.Latent.Input("prev_latents", tooltip="Full previous latents to be used to continue generation, continuation frames are selected based on 'overlap' parameter"),
|
||||
io.Custom("MULTITALK_EMBEDS").Input("audio_embeds", tooltip="Full length audio embeddings"),
|
||||
io.Int.Input("num_frames", default=93, min=1, max=256, step=1, tooltip="Number of new frames to generate"),
|
||||
io.Int.Input("overlap", default=13, min=0, max=16, step=1, tooltip="Number of overlapping frames from previous latents for video continuation, set to 0 for T2V"),
|
||||
io.Int.Input("frames_processed", default=0, min=0, max=10000, step=1, tooltip="Number of frames already processed in the video, used to select audio features"),
|
||||
io.Combo.Input("if_not_enough_audio", ["pad_with_start", "mirror_from_end"], default="pad_with_start", tooltip="What to do if there are not enough frames in pose_images for the window"),
|
||||
io.Int.Input("ref_frame_index", default=10, min=0, max=1000, step=1, tooltip="Values between 0 - 24 ensures better consistency, while selecting other ranges (e.g., -10 or 30) helps reduce repeated actions"),
|
||||
io.Int.Input("ref_mask_frame_range", default=3, min=0, max=20, step=1, tooltip="Larger range can further help mitigate repeated actions, but excessively large values may introduce artifacts"),
|
||||
io.Latent.Input("ref_latent", optional=True, tooltip="Reference latent used for consistency, generally should be either the init image, or first latent from first generation"),
|
||||
io.Latent.Input("samples", optional=True, tooltip="For the sampler 'samples' input, used for slicing samples per window for vid2vid"),
|
||||
],
|
||||
outputs=[
|
||||
io.Custom("WANVIDIMAGE_EMBEDS").Output(display_name="image_embeds", tooltip="Embeds for WanVideo LongCat Avatar generation"),
|
||||
io.Latent.Output(display_name="samples_slice", tooltip="Sliced latent samples for the new frames"),
|
||||
],
|
||||
)
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def add(self, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed=0, ref_latent=None):
|
||||
@classmethod
|
||||
def execute(cls, prev_latents, audio_embeds, num_frames, overlap, if_not_enough_audio, frames_processed, ref_frame_index, ref_mask_frame_range, ref_latent=None, samples=None) -> io.NodeOutput:
|
||||
|
||||
new_audio_embed = audio_embeds.copy()
|
||||
|
||||
audio_features = torch.stack(new_audio_embed["audio_features"])
|
||||
print("audio_features shape: ", audio_features.shape)
|
||||
if audio_features.shape[1] < frames_processed + num_frames:
|
||||
deficit = frames_processed + num_frames - audio_features.shape[1]
|
||||
if if_not_enough_audio == "pad_with_start":
|
||||
@@ -47,19 +51,18 @@ class WanVideoLongCatAvatarExtendEmbeds:
|
||||
if ref_target_masks is not None:
|
||||
new_audio_embed["ref_target_masks"] = ref_target_masks[:, frames_processed:frames_processed+num_frames, :]
|
||||
|
||||
latent_overlap = (overlap - 1) // 4 + 1
|
||||
print("prev_latents shape: ", prev_latents["samples"].shape, "latent_overlap: ", latent_overlap)
|
||||
prev_samples = prev_latents["samples"][:, :, -latent_overlap:].clone()
|
||||
prev_samples = prev_latents["samples"].clone()
|
||||
if overlap != 0:
|
||||
latent_overlap = (overlap - 1) // 4 + 1
|
||||
prev_samples = prev_samples[:, :, -latent_overlap:]
|
||||
|
||||
ref_sample = None
|
||||
if ref_latent is not None:
|
||||
ref_sample = ref_latent["samples"][0, :, :1].clone()
|
||||
|
||||
log.info(f"Previous latents shape: {prev_samples.shape}, using last {latent_overlap} latent frames for overlap.")
|
||||
log.info(f"Previous latents shape: {prev_samples.shape}, using last {latent_overlap} latent frames for overlap.")
|
||||
|
||||
new_latent_frames = (num_frames - 1) // 4 + 1
|
||||
target_shape = (16, new_latent_frames, prev_samples.shape[-2], prev_samples.shape[-1])
|
||||
print("target_shape: ", target_shape)
|
||||
|
||||
audio_stride = 2
|
||||
indices = torch.arange(2 * 2 + 1) - 2
|
||||
@@ -72,9 +75,6 @@ class WanVideoLongCatAvatarExtendEmbeds:
|
||||
|
||||
log.info(f"Extracting audio embeddings from index {audio_start_idx} to {audio_end_idx}")
|
||||
|
||||
#center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
|
||||
#center_indices = torch.clamp(center_indices, min=0, max=audio_features.shape[0]-1)
|
||||
#audio_emb = audio_features[center_indices][None,...]
|
||||
audio_embs = []
|
||||
for human_idx in range(len(audio_features)):
|
||||
center_indices = torch.arange(audio_start_idx, audio_end_idx, audio_stride).unsqueeze(1) + indices.unsqueeze(0)
|
||||
@@ -87,15 +87,28 @@ class WanVideoLongCatAvatarExtendEmbeds:
|
||||
new_audio_embed["audio_features"] = None
|
||||
new_audio_embed["audio_emb_slice"] = audio_emb
|
||||
|
||||
longcat_avatar_options = {
|
||||
"longcat_ref_latent": ref_sample,
|
||||
"ref_frame_index": ref_frame_index,
|
||||
"ref_mask_frame_range": ref_mask_frame_range,
|
||||
}
|
||||
|
||||
embeds = {
|
||||
"target_shape": target_shape,
|
||||
"num_frames": num_frames,
|
||||
"extra_latents": [{"samples": prev_samples, "index": 0}],
|
||||
"extra_latents": [{"samples": prev_samples, "index": 0}] if overlap != 0 else None,
|
||||
"multitalk_embeds": new_audio_embed,
|
||||
"longcat_ref_latent": ref_sample,
|
||||
"longcat_avatar_options": longcat_avatar_options,
|
||||
}
|
||||
|
||||
return (embeds,)
|
||||
samples_slice = None
|
||||
if samples is not None:
|
||||
latent_start_index = (frames_processed - 1) // 4 + 1 if frames_processed > 0 else 0
|
||||
latent_end_index = latent_start_index + new_latent_frames
|
||||
samples_slice = samples.copy()
|
||||
samples_slice["samples"] = samples["samples"][:, :, latent_start_index:latent_end_index].clone()
|
||||
|
||||
return io.NodeOutput(embeds, samples_slice)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
@@ -103,4 +116,4 @@ NODE_CLASS_MAPPINGS = {
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoLongCatAvatarExtendEmbeds": "WanVideo LongCat Avatar Extend Embeds",
|
||||
}
|
||||
}
|
||||
+12
-8
@@ -876,14 +876,17 @@ class WanVideoSampler:
|
||||
latent = noise
|
||||
|
||||
# LongCat-Avatar
|
||||
longcat_ref_latent = image_embeds.get("longcat_ref_latent", None)
|
||||
if longcat_ref_latent is not None:
|
||||
latent = torch.cat([longcat_ref_latent.to(latent), latent], dim=1)
|
||||
seq_len = math.ceil((latent.shape[2] * latent.shape[3]) / 4 * latent.shape[1])
|
||||
insert_len = longcat_ref_latent.shape[1]
|
||||
clean_latent_indices = list(range(0, insert_len)) + [i + insert_len for i in clean_latent_indices]
|
||||
latent_video_length += insert_len
|
||||
print("clean_latent_indices:", clean_latent_indices)
|
||||
longcat_ref_latent = None
|
||||
longcat_avatar_options = image_embeds.get("longcat_avatar_options", None)
|
||||
if longcat_avatar_options is not None:
|
||||
longcat_ref_latent = image_embeds.get("longcat_ref_latent", None)
|
||||
if longcat_ref_latent is not None:
|
||||
latent = torch.cat([longcat_ref_latent.to(latent), latent], dim=1)
|
||||
seq_len = math.ceil((latent.shape[2] * latent.shape[3]) / 4 * latent.shape[1])
|
||||
insert_len = longcat_ref_latent.shape[1]
|
||||
clean_latent_indices = list(range(0, insert_len)) + [i + insert_len for i in clean_latent_indices]
|
||||
latent_video_length += insert_len
|
||||
log.info(f"LongCat clean_latent_indices: {clean_latent_indices}")
|
||||
audio_stride = 2 if transformer.is_longcat else 1
|
||||
|
||||
#controlnet
|
||||
@@ -1566,6 +1569,7 @@ class WanVideoSampler:
|
||||
"flashvsr_strength": flashvsr_strength, # FlashVSR strength
|
||||
"longcat_num_cond_latents": len(clean_latent_indices) if transformer.is_longcat else 0,
|
||||
"longcat_num_ref_latents": longcat_ref_latent.shape[1] if longcat_ref_latent is not None else 0,
|
||||
"longcat_avatar_options": longcat_avatar_options, # LongCat avatar attention options
|
||||
"sdancer_input": sdancer_input, # SteadyDancer 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,
|
||||
|
||||
@@ -1003,7 +1003,7 @@ class WanAttentionBlock(nn.Module):
|
||||
humo_audio_input=None, humo_audio_scale=1.0, #humo audio
|
||||
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,
|
||||
longcat_num_cond_latents=0, #longcat image cond amount
|
||||
longcat_num_cond_latents=0, longcat_avatar_options=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
|
||||
):
|
||||
@@ -1212,9 +1212,9 @@ class WanAttentionBlock(nn.Module):
|
||||
# process the noise tokens
|
||||
q_noise = q[:, num_cond_latents_thw:].contiguous()
|
||||
start_noise, end_noise, num_noisy_frames = 0, 0, num_latent_frames - longcat_num_cond_latents
|
||||
mask_frame_range = 3 #todo: make it configurable?
|
||||
ref_img_index = 10 #todo: make it configurable?
|
||||
num_ref_latents = 1 # todo: make it configurable?
|
||||
mask_frame_range = longcat_avatar_options["ref_mask_frame_range"]
|
||||
ref_img_index = longcat_avatar_options["ref_frame_index"]
|
||||
num_ref_latents = 1
|
||||
if mask_frame_range is not None and mask_frame_range > 0:
|
||||
start_noise = ref_img_index - mask_frame_range - longcat_num_cond_latents + num_ref_latents
|
||||
end_noise = ref_img_index + mask_frame_range - longcat_num_cond_latents + num_ref_latents + 1
|
||||
@@ -2313,7 +2313,7 @@ class WanModel(torch.nn.Module):
|
||||
lynx_embeds=None,
|
||||
x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None,
|
||||
flashvsr_LQ_latent=None, flashvsr_strength=1.0,
|
||||
longcat_num_cond_latents=0, longcat_num_ref_latents=0, # for LongCat
|
||||
longcat_num_cond_latents=0, longcat_num_ref_latents=0, longcat_avatar_options=None, # for LongCat
|
||||
add_text_emb=None,
|
||||
sdancer_input=None, # SteadyDancer
|
||||
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
|
||||
@@ -2709,18 +2709,13 @@ class WanModel(torch.nn.Module):
|
||||
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:
|
||||
print("input t shape:", t.shape)
|
||||
print("F:", F)
|
||||
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
|
||||
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
|
||||
time_embed_dtype = self.base_dtype
|
||||
if len(t.shape) == 1:
|
||||
t = t.unsqueeze(1).expand(-1, F) # [B, T]
|
||||
print("t expanded shape:", t.shape)
|
||||
self.time_embedding.to(torch.float32)
|
||||
print("t float shape:", t.float().flatten().shape)
|
||||
e = e0 = self.time_embedding(t.float().flatten(), dtype=torch.float32)#.reshape(1, F, -1)
|
||||
print("e0 shape:", e0.shape)
|
||||
e = e0 = e0.reshape(1, F, -1)
|
||||
|
||||
if self.audio_model is not None:
|
||||
@@ -2861,7 +2856,7 @@ class WanModel(torch.nn.Module):
|
||||
human_num = len(multitalk_audio_embedding)
|
||||
|
||||
# LongCat-Avatar specific
|
||||
print("longcat_num_cond_latents:", longcat_num_cond_latents, "longcat_num_ref_latents:", longcat_num_ref_latents)
|
||||
tqdm.write(f"longcat_num_cond_latents: {longcat_num_cond_latents}, longcat_num_ref_latents: {longcat_num_ref_latents}")
|
||||
|
||||
if longcat_num_ref_latents > 0:
|
||||
audio_start_ref = multitalk_audio_embedding[:, [0], :, :] # padding
|
||||
@@ -3070,6 +3065,7 @@ class WanModel(torch.nn.Module):
|
||||
lynx_ip_scale=lynx_ip_scale,
|
||||
lynx_ref_scale=lynx_ref_scale,
|
||||
longcat_num_cond_latents=longcat_num_cond_latents,
|
||||
longcat_avatar_options=longcat_avatar_options,
|
||||
onetoall_ref_scale=onetoall_ref_scale,
|
||||
e_tr=e0_token_replace if use_token_replace else None,
|
||||
tr_start=token_replace_start,
|
||||
|
||||
Reference in New Issue
Block a user