Merge branch 'main' into dev
This commit is contained in:
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -20,8 +20,8 @@ class DownloadAndLoadWav2VecModel:
|
||||
"required": {
|
||||
"model": (
|
||||
[
|
||||
"facebook/wav2vec2-base-960h",
|
||||
"TencentGameMate/chinese-wav2vec2-base"
|
||||
"TencentGameMate/chinese-wav2vec2-base",
|
||||
"facebook/wav2vec2-base-960h"
|
||||
],
|
||||
),
|
||||
|
||||
|
||||
+39
-18
@@ -52,7 +52,8 @@ class MultiTalkModelLoader:
|
||||
multitalk = {
|
||||
"proj_model": multitalk_proj_model,
|
||||
"sd": sd,
|
||||
"is_gguf": model_path.endswith(".gguf")
|
||||
"is_gguf": model_path.endswith(".gguf"),
|
||||
"model_type": "InfiniteTalk" if "infinite" in model.lower() else "MultiTalk",
|
||||
}
|
||||
|
||||
return (multitalk,)
|
||||
@@ -75,8 +76,8 @@ class MultiTalkWav2VecEmbeds:
|
||||
return {"required": {
|
||||
"wav2vec_model": ("WAV2VECMODEL",),
|
||||
"audio_1": ("AUDIO",),
|
||||
"normalize_loudness": ("BOOLEAN", {"default": True}),
|
||||
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 1}),
|
||||
"normalize_loudness": ("BOOLEAN", {"default": True, "tooltip": "Normalize the audio loudness to -23 LUFS"}),
|
||||
"num_frames": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 1, "tooltip": "The total frame count to generate."}),
|
||||
"fps": ("FLOAT", {"default": 25.0, "min": 1.0, "max": 60.0, "step": 0.1}),
|
||||
"audio_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "Strength of the audio conditioning"}),
|
||||
"audio_cfg_scale": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.1, "tooltip": "When not 1.0, an extra model pass without audio conditioning is done: slower inference but more motion is allowed"}),
|
||||
@@ -90,8 +91,8 @@ class MultiTalkWav2VecEmbeds:
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", )
|
||||
RETURN_NAMES = ("multitalk_embeds", "audio", )
|
||||
RETURN_TYPES = ("MULTITALK_EMBEDS", "AUDIO", "INT", )
|
||||
RETURN_NAMES = ("multitalk_embeds", "audio", "num_frames", )
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
@@ -230,12 +231,30 @@ class MultiTalkWav2VecEmbeds:
|
||||
offset += w.shape[-1]
|
||||
out_audio = {"waveform": mixed, "sample_rate": sr}
|
||||
|
||||
# Calculate actual frames based on audio duration
|
||||
actual_num_frames = num_frames
|
||||
if len(audio_outputs) > 0:
|
||||
if multi_audio_type == "para":
|
||||
# For parallel mode, use the longest audio duration
|
||||
max_audio_duration = max([ao["waveform"].shape[-1] / sr for ao in audio_outputs])
|
||||
actual_frames_from_audio = int(max_audio_duration * fps)
|
||||
else: # "add"
|
||||
# For sequential mode, use the total audio duration
|
||||
total_audio_duration = sum([ao["waveform"].shape[-1] / sr for ao in audio_outputs])
|
||||
actual_frames_from_audio = int(total_audio_duration * fps)
|
||||
|
||||
# Use the smaller of requested frames or actual audio frames
|
||||
actual_num_frames = min(num_frames, actual_frames_from_audio)
|
||||
|
||||
if actual_frames_from_audio < num_frames:
|
||||
log.info(f"[MultiTalk] Audio duration ({actual_frames_from_audio} frames) is shorter than requested ({num_frames} frames). Using {actual_num_frames} frames.")
|
||||
|
||||
# Debug: log final mixed audio length and mode
|
||||
total_samples_raw = sum([ao["waveform"].shape[-1] for ao in audio_outputs])
|
||||
log.info(f"[MultiTalk] total raw duration = {total_samples_raw/sr:.3f}s")
|
||||
log.info(f"[MultiTalk] multi_audio_type={multi_audio_type} | final waveform shape={out_audio['waveform'].shape} | length={out_audio['waveform'].shape[-1]} samples | seconds={out_audio['waveform'].shape[-1]/sr:.3f}s (expected {'sum' if multi_audio_type=='add' else 'max'} of raw)")
|
||||
|
||||
return (multitalk_embeds, out_audio)
|
||||
return (multitalk_embeds, out_audio, actual_num_frames)
|
||||
|
||||
|
||||
class WanVideoImageToVideoMultiTalk:
|
||||
@@ -243,11 +262,11 @@ class WanVideoImageToVideoMultiTalk:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"vae": ("WANVAE",),
|
||||
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}),
|
||||
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}),
|
||||
"frame_window_size": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "Number of frames to encode"}),
|
||||
"motion_frame": ("INT", {"default": 25, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation."}),
|
||||
"force_offload": ("BOOLEAN", {"default": True}),
|
||||
"width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the generation"}),
|
||||
"height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the generation"}),
|
||||
"frame_window_size": ("INT", {"default": 81, "min": 1, "max": 10000, "step": 4, "tooltip": "The number of frames to process at once, should be a value the model is generally good at."}),
|
||||
"motion_frame": ("INT", {"default": 25, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation. Basically the overlap length."}),
|
||||
"force_offload": ("BOOLEAN", {"default": False, "tooltip": "Whether to force offload the model within the loop for VAE operations, enable if you encounter memory issues."}),
|
||||
"colormatch": (
|
||||
[
|
||||
'disabled',
|
||||
@@ -258,17 +277,18 @@ class WanVideoImageToVideoMultiTalk:
|
||||
'hm-mvgd-hm',
|
||||
'hm-mkl-hm',
|
||||
], {
|
||||
"default": 'disabled'
|
||||
}),
|
||||
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
|
||||
},),
|
||||
},
|
||||
"optional": {
|
||||
"start_image": ("IMAGE", {"tooltip": "Image to encode"}),
|
||||
"start_image": ("IMAGE", {"tooltip": "Images to encode"}),
|
||||
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
|
||||
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
|
||||
"mode": ([
|
||||
"auto",
|
||||
"multitalk",
|
||||
"infinitetalk"
|
||||
], {"default": "multitalk", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
|
||||
], {"default": "auto", "tooltip": "The sampling strategy to use in the long video generation loop, should match the model used"})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -276,6 +296,7 @@ class WanVideoImageToVideoMultiTalk:
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Enables Multi/InfiniteTalk long video generation sampling method, the video is created in windows with overlapping frames. Not compatible or necessary to be used with context windows and many other features besides Multi/InfiniteTalk."
|
||||
|
||||
def process(self, vae, width, height, frame_window_size, motion_frame, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None, mode="multitalk"):
|
||||
|
||||
@@ -320,7 +341,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MultiTalkModelLoader": "MultiTalk Model Loader",
|
||||
"MultiTalkWav2VecEmbeds": "MultiTalk Wav2Vec Embeds",
|
||||
"WanVideoImageToVideoMultiTalk": "WanVideo Image To Video MultiTalk"
|
||||
"MultiTalkModelLoader": "Multi/InfiniteTalk Model Loader",
|
||||
"MultiTalkWav2VecEmbeds": "Multi/InfiniteTalk Wav2Vec Embeds",
|
||||
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk"
|
||||
}
|
||||
@@ -874,8 +874,8 @@ class WanVideoImageToVideoEncode:
|
||||
H = height
|
||||
W = width
|
||||
|
||||
lat_h = H // 8
|
||||
lat_w = W // 8
|
||||
lat_h = H // vae.upsampling_factor
|
||||
lat_w = W // vae.upsampling_factor
|
||||
|
||||
num_frames = ((num_frames - 1) // 4) * 4 + 1
|
||||
two_ref_images = start_image is not None and end_image is not None
|
||||
@@ -1280,8 +1280,6 @@ class WanVideoVACEEncode:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, vae, width, height, num_frames, strength, vace_start_percent, vace_end_percent, input_frames=None, ref_images=None, input_masks=None, prev_vace_embeds=None, tiled_vae=False):
|
||||
vae = vae.to(device)
|
||||
|
||||
width = (width // 16) * 16
|
||||
height = (height // 16) * 16
|
||||
|
||||
@@ -1292,8 +1290,8 @@ class WanVideoVACEEncode:
|
||||
if input_frames is None:
|
||||
input_frames = torch.zeros((1, 3, num_frames, height, width), device=device, dtype=vae.dtype)
|
||||
else:
|
||||
input_frames = input_frames[:num_frames]
|
||||
input_frames = common_upscale(input_frames.clone().movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
|
||||
input_frames = input_frames.clone()[:num_frames, :, :, :3]
|
||||
input_frames = common_upscale(input_frames.movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
|
||||
input_frames = input_frames.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3) # B, C, T, H, W
|
||||
input_frames = input_frames * 2 - 1
|
||||
if input_masks is None:
|
||||
@@ -1306,6 +1304,7 @@ class WanVideoVACEEncode:
|
||||
input_masks = input_masks.unsqueeze(-1).unsqueeze(0).permute(0, 4, 1, 2, 3).repeat(1, 3, 1, 1, 1) # B, C, T, H, W
|
||||
|
||||
if ref_images is not None:
|
||||
ref_images = ref_images.clone()[..., :3]
|
||||
# Create padded image
|
||||
if ref_images.shape[0] > 1:
|
||||
ref_images = torch.cat([ref_images[i] for i in range(ref_images.shape[0])], dim=1).unsqueeze(0)
|
||||
@@ -1332,11 +1331,11 @@ class WanVideoVACEEncode:
|
||||
ref_images = ref_images.to(vae.dtype).to(device).unsqueeze(0).permute(0, 4, 1, 2, 3).unsqueeze(0)
|
||||
ref_images = ref_images * 2 - 1
|
||||
|
||||
vae = vae.to(device)
|
||||
z0 = self.vace_encode_frames(vae, input_frames, ref_images, masks=input_masks, tiled_vae=tiled_vae)
|
||||
vae.model.clear_cache()
|
||||
m0 = self.vace_encode_masks(input_masks, ref_images)
|
||||
z = self.vace_latent(z0, m0)
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
vace_input = {
|
||||
@@ -1357,6 +1356,7 @@ class WanVideoVACEEncode:
|
||||
vace_input["additional_vace_inputs"].append(prev_vace_embeds)
|
||||
|
||||
return (vace_input,)
|
||||
|
||||
def vace_encode_frames(self, vae, frames, ref_images, masks=None, tiled_vae=False):
|
||||
if ref_images is None:
|
||||
ref_images = [None] * len(frames)
|
||||
@@ -1675,6 +1675,10 @@ class WanVideoSampler:
|
||||
transformer = compile_model(transformer, model["compile_args"])
|
||||
|
||||
multitalk_sampling = image_embeds.get("multitalk_sampling", False)
|
||||
|
||||
if multitalk_sampling and context_options is not None:
|
||||
raise Exception("context_options are not compatible or necessary with 'WanVideoImageToVideoMultiTalk' node, since it's already an alternative method that creates the video in a loop.")
|
||||
|
||||
if not multitalk_sampling and scheduler == "multitalk":
|
||||
raise Exception("multitalk scheduler is only for multitalk sampling when using ImagetoVideoMultiTalk -node")
|
||||
|
||||
@@ -1803,7 +1807,7 @@ class WanVideoSampler:
|
||||
|
||||
control_embeds = image_embeds.get("control_embeds", None)
|
||||
if control_embeds is not None:
|
||||
if transformer.in_dim not in [52, 48, 36, 32]:
|
||||
if transformer.in_dim not in [148, 52, 48, 36, 32]:
|
||||
raise ValueError("Control signal only works with Fun-Control model")
|
||||
|
||||
control_latents = control_embeds.get("control_images", None)
|
||||
@@ -1900,10 +1904,10 @@ class WanVideoSampler:
|
||||
patcher = apply_lora(patcher, device, device, low_mem_load=False, control_lora=True)
|
||||
patcher.model.is_patched = True
|
||||
else:
|
||||
if transformer.in_dim not in [48, 36, 32, 52]:
|
||||
if transformer.in_dim not in [148, 48, 36, 32, 52]:
|
||||
raise ValueError("Control signal only works with Fun-Control model")
|
||||
image_cond = torch.zeros_like(noise).to(device) #fun control
|
||||
if transformer.in_dim == 52 or transformer.control_adapter is not None: #fun 2.2 control
|
||||
if transformer.in_dim in [148, 52] or transformer.control_adapter is not None: #fun 2.2 control
|
||||
mask_latents = torch.tile(
|
||||
torch.zeros_like(noise[:1]), [4, 1, 1, 1]
|
||||
)
|
||||
@@ -1917,7 +1921,7 @@ class WanVideoSampler:
|
||||
control_start_percent = control_embeds.get("start_percent", 0.0)
|
||||
control_end_percent = control_embeds.get("end_percent", 1.0)
|
||||
else:
|
||||
if transformer.in_dim == 36: #fun inp
|
||||
if transformer.in_dim in [148, 52]: #fun inp
|
||||
mask_latents = torch.tile(
|
||||
torch.zeros_like(noise[:1]), [4, 1, 1, 1]
|
||||
)
|
||||
@@ -2113,7 +2117,7 @@ class WanVideoSampler:
|
||||
mtv_freqs = mtv_freqs.to(device, dtype)
|
||||
|
||||
# vid2vid
|
||||
if samples is not None:
|
||||
if samples is not None and not multitalk_sampling:
|
||||
saved_generator_state = samples.get("generator_state", None)
|
||||
if saved_generator_state is not None:
|
||||
seed_g.set_state(saved_generator_state)
|
||||
@@ -2127,29 +2131,29 @@ class WanVideoSampler:
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
mask = samples.get("noise_mask", None)
|
||||
if mask is not None:
|
||||
log.info(f"Latent mask shape: {mask.shape}")
|
||||
noise_mask = samples.get("noise_mask", None)
|
||||
if noise_mask is not None:
|
||||
log.info(f"Latent noise_mask shape: {noise_mask.shape}")
|
||||
original_image = input_samples.to(device)
|
||||
if len(mask.shape) == 4:
|
||||
mask = mask.squeeze(1)
|
||||
if len(noise_mask.shape) == 4:
|
||||
noise_mask = noise_mask.squeeze(1)
|
||||
|
||||
mask = torch.nn.functional.interpolate(
|
||||
mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
noise_mask = torch.nn.functional.interpolate(
|
||||
noise_mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
|
||||
size=(noise.shape[1], noise.shape[2], noise.shape[3]),
|
||||
mode='trilinear',
|
||||
align_corners=False
|
||||
).squeeze(0) # Remove batch dim, keep channel dim
|
||||
|
||||
# Add batch & channel dims for final output
|
||||
mask = mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1)
|
||||
noise_mask = noise_mask.unsqueeze(0).repeat(1, noise.shape[0], 1, 1, 1)
|
||||
|
||||
if mask.shape[2] != noise.shape[1]:
|
||||
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
|
||||
if noise_mask.shape[2] != noise.shape[1]:
|
||||
noise_mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - noise_mask.shape[2], noise.shape[2], noise.shape[3]), noise_mask], dim=2)
|
||||
|
||||
# extra latents (Pusa) and 5b
|
||||
latents_to_insert = add_index = None
|
||||
if (extra_latents := image_embeds.get("extra_latents", None)) is not None:
|
||||
if (extra_latents := image_embeds.get("extra_latents", None)) is not None and transformer.multitalk_model_type.lower() != "infinitetalk":
|
||||
all_indices = []
|
||||
for entry in extra_latents:
|
||||
add_index = entry["index"]
|
||||
@@ -2178,7 +2182,7 @@ class WanVideoSampler:
|
||||
if uni3c_embeds is not None:
|
||||
transformer.controlnet = uni3c_embeds["controlnet"]
|
||||
pcd_data = {
|
||||
"render_latent": uni3c_embeds["render_latent"].to(dtype),
|
||||
"render_latent": uni3c_embeds["render_latent"],
|
||||
"render_mask": uni3c_embeds["render_mask"],
|
||||
"camera_embedding": uni3c_embeds["camera_embedding"],
|
||||
"controlnet_weight": uni3c_embeds["controlnet_weight"],
|
||||
@@ -2701,6 +2705,7 @@ class WanVideoSampler:
|
||||
from .latent_preview import prepare_callback #custom for tiny VAE previews
|
||||
callback = prepare_callback(patcher, len(timesteps))
|
||||
|
||||
if not multitalk_sampling:
|
||||
log.info(f"Input sequence length: {seq_len}")
|
||||
log.info(f"Sampling {(latent_video_length-1) * 4 + 1} frames at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
|
||||
@@ -2708,11 +2713,11 @@ class WanVideoSampler:
|
||||
|
||||
# diff diff prep
|
||||
masks = None
|
||||
if samples is not None and mask is not None:
|
||||
mask = 1 - mask
|
||||
if not multitalk_sampling and samples is not None and noise_mask is not None:
|
||||
noise_mask = 1 - noise_mask
|
||||
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
|
||||
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
|
||||
masks = mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
|
||||
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
|
||||
masks = masks > thresholds
|
||||
|
||||
latent_shift_loop = False
|
||||
@@ -2789,27 +2794,24 @@ class WanVideoSampler:
|
||||
try:
|
||||
pbar = ProgressBar(len(timesteps))
|
||||
#region main loop start
|
||||
for idx, t in enumerate(tqdm(timesteps)):
|
||||
for idx, t in enumerate(tqdm(timesteps, disable=multitalk_sampling)):
|
||||
if flowedit_args is not None:
|
||||
if idx < skip_steps:
|
||||
continue
|
||||
|
||||
# diff diff
|
||||
if masks is not None:
|
||||
if idx < len(timesteps) - 1:
|
||||
noise_timestep = timesteps[idx+1]
|
||||
image_latent = sample_scheduler.scale_noise(
|
||||
original_image, torch.tensor([noise_timestep]), noise.to(device)
|
||||
)
|
||||
mask = masks[idx]
|
||||
mask = mask.to(latent)
|
||||
latent = image_latent * mask + latent * (1-mask)
|
||||
# end diff diff
|
||||
|
||||
if bidirectional_sampling:
|
||||
latent_flipped = torch.flip(latent, dims=[1])
|
||||
latent_model_input_flipped = latent_flipped.to(device)
|
||||
|
||||
#InfiniteTalk first frame handling
|
||||
if (extra_latents is not None
|
||||
and not multitalk_sampling
|
||||
and transformer.multitalk_model_type=="InfiniteTalk"):
|
||||
for entry in extra_latents:
|
||||
add_index = entry["index"]
|
||||
num_extra_frames = entry["samples"].shape[2]
|
||||
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
|
||||
|
||||
latent_model_input = latent.to(device)
|
||||
|
||||
current_step_percentage = idx / len(timesteps)
|
||||
@@ -3097,8 +3099,10 @@ class WanVideoSampler:
|
||||
#region multitalk
|
||||
elif multitalk_sampling:
|
||||
mode = image_embeds.get("multitalk_mode", "multitalk")
|
||||
if mode == "auto":
|
||||
mode = transformer.multitalk_model_type.lower()
|
||||
log.info(f"Multitalk mode: {mode}")
|
||||
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
|
||||
cond_frame = None
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
tiled_vae = image_embeds.get("tiled_vae", False)
|
||||
frame_num = clip_length = image_embeds.get("num_frames", 81)
|
||||
@@ -3110,23 +3114,20 @@ class WanVideoSampler:
|
||||
motion_frame = image_embeds.get("motion_frame", 25)
|
||||
target_w = image_embeds.get("target_w", None)
|
||||
target_h = image_embeds.get("target_h", None)
|
||||
original_images = cond_image = image_embeds.get("multitalk_start_image", None)
|
||||
if original_images is None:
|
||||
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
|
||||
|
||||
if len(multitalk_embeds['audio_features'])==2 and (multitalk_embeds['ref_target_masks'] is None):
|
||||
face_scale = 0.1
|
||||
x_min, x_max = int(target_h * face_scale), int(target_h * (1 - face_scale))
|
||||
background_mask = torch.zeros([target_h, target_w])
|
||||
background_mask = torch.zeros([target_h, target_w])
|
||||
human_mask1 = torch.zeros([target_h, target_w])
|
||||
human_mask2 = torch.zeros([target_h, target_w])
|
||||
lefty_min, lefty_max = int((target_w//2) * face_scale), int((target_w//2) * (1 - face_scale))
|
||||
righty_min, righty_max = int((target_w//2) * face_scale + (target_w//2)), int((target_w//2) * (1 - face_scale) + (target_w//2))
|
||||
human_mask1, human_mask2 = (torch.zeros([target_h, target_w]) for _ in range(2))
|
||||
human_mask1[x_min:x_max, lefty_min:lefty_max] = 1
|
||||
human_mask2[x_min:x_max, righty_min:righty_max] = 1
|
||||
background_mask += human_mask1
|
||||
background_mask += human_mask2
|
||||
human_masks = [human_mask1, human_mask2]
|
||||
background_mask = torch.where(background_mask > 0, torch.tensor(0), torch.tensor(1))
|
||||
human_masks.append(background_mask)
|
||||
background_mask = torch.where((human_mask1 + human_mask2) > 0, torch.tensor(0), torch.tensor(1))
|
||||
human_masks = [human_mask1, human_mask2, background_mask]
|
||||
ref_target_masks = torch.stack(human_masks, dim=0)
|
||||
multitalk_embeds['ref_target_masks'] = ref_target_masks
|
||||
|
||||
@@ -3134,7 +3135,7 @@ class WanVideoSampler:
|
||||
is_first_clip = True
|
||||
arrive_last_frame = False
|
||||
cur_motion_frames_num = 1
|
||||
audio_start_idx = iteration_count = 0
|
||||
audio_start_idx = iteration_count = step_iteration_count= 0
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
current_condframe_index = 0
|
||||
@@ -3146,7 +3147,7 @@ class WanVideoSampler:
|
||||
if uni3c_embeds is not None:
|
||||
transformer.controlnet = uni3c_embeds["controlnet"]
|
||||
pcd_data = {
|
||||
"render_latent": uni3c_embeds["render_latent"].to(dtype),
|
||||
"render_latent": uni3c_embeds["render_latent"],
|
||||
"render_mask": uni3c_embeds["render_mask"],
|
||||
"camera_embedding": uni3c_embeds["camera_embedding"],
|
||||
"controlnet_weight": uni3c_embeds["controlnet_weight"],
|
||||
@@ -3155,19 +3156,19 @@ class WanVideoSampler:
|
||||
}
|
||||
|
||||
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
|
||||
log.info(f"Total frames: {total_frames}, frame_num: {frame_num}, motion_frame: {motion_frame}")
|
||||
log.info(f"Estimated iterations: {estimated_iterations}")
|
||||
loop_pbar = tqdm(total=estimated_iterations, desc="Generating video clips")
|
||||
loop_pbar = tqdm(total=estimated_iterations, desc="Total progress", position=1, leave=True)
|
||||
callback = prepare_callback(patcher, estimated_iterations)
|
||||
|
||||
audio_embedding = multitalk_audio_embedding
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
|
||||
log.info(f"Sampling {total_frames} frames in {estimated_iterations} windows, at {latent.shape[3]*vae_upscale_factor}x{latent.shape[2]*vae_upscale_factor} with {steps} steps")
|
||||
|
||||
while True: # start video generation iteratively
|
||||
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
|
||||
if mode == "infinitetalk":
|
||||
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1]
|
||||
log.info(f"current_condframe_index: {current_condframe_index}")
|
||||
log.info(f"audio_start_idx: {audio_start_idx}")
|
||||
cond_image = original_images[:, :, current_condframe_index:current_condframe_index+1] if cond_image is not None else None
|
||||
if multitalk_embeds is not None:
|
||||
audio_embs = []
|
||||
# split audio with window size
|
||||
@@ -3180,10 +3181,22 @@ class WanVideoSampler:
|
||||
|
||||
if uni3c_embeds is not None:
|
||||
vae.to(device)
|
||||
render_latent = vae.encode(original_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
|
||||
# Pad original_images if needed
|
||||
num_frames = original_images.shape[2]
|
||||
required_frames = audio_end_idx - audio_start_idx
|
||||
if audio_end_idx > num_frames:
|
||||
pad_len = audio_end_idx - num_frames
|
||||
last_frame = original_images[:, :, -1:].repeat(1, 1, pad_len, 1, 1)
|
||||
padded_images = torch.cat([original_images, last_frame], dim=2)
|
||||
else:
|
||||
padded_images = original_images
|
||||
render_latent = vae.encode(
|
||||
padded_images[:, :, audio_start_idx:audio_end_idx].to(device, vae.dtype),
|
||||
device=device, tiled=tiled_vae
|
||||
).to(dtype)
|
||||
pcd_data['render_latent'] = render_latent
|
||||
|
||||
h, w = cond_image.shape[-2], cond_image.shape[-1]
|
||||
h, w = (cond_image.shape[-2], cond_image.shape[-1]) if cond_image is not None else (target_h, target_w)
|
||||
lat_h, lat_w = h // VAE_STRIDE[1], w // VAE_STRIDE[2]
|
||||
seq_len = ((frame_num - 1) // VAE_STRIDE[0] + 1) * lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
|
||||
|
||||
@@ -3191,7 +3204,60 @@ class WanVideoSampler:
|
||||
16, (frame_num - 1) // 4 + 1,
|
||||
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
|
||||
|
||||
# get mask
|
||||
# Calculate the correct latent slice based on current iteration
|
||||
if is_first_clip:
|
||||
latent_start_idx = 0
|
||||
latent_end_idx = noise.shape[1]
|
||||
else:
|
||||
new_frames_per_iteration = frame_num - motion_frame
|
||||
new_latent_frames_per_iteration = ((new_frames_per_iteration - 1) // 4 + 1)
|
||||
latent_start_idx = iteration_count * new_latent_frames_per_iteration
|
||||
latent_end_idx = latent_start_idx + noise.shape[1]
|
||||
|
||||
if samples is not None:
|
||||
input_samples = samples["samples"].squeeze(0).to(noise)
|
||||
# Check if we have enough frames in input_samples
|
||||
if latent_end_idx > input_samples.shape[1]:
|
||||
# We need more frames than available - pad the input_samples at the end
|
||||
pad_length = latent_end_idx - input_samples.shape[1]
|
||||
last_frame = input_samples[:, -1:].repeat(1, pad_length, 1, 1)
|
||||
input_samples = torch.cat([input_samples, last_frame], dim=1)
|
||||
input_samples = input_samples[:, latent_start_idx:latent_end_idx]
|
||||
|
||||
assert input_samples.shape[1] == noise.shape[1], f"Slice mismatch: {input_samples.shape[1]} vs {noise.shape[1]}"
|
||||
|
||||
if add_noise_to_samples:
|
||||
latent_timestep = timesteps[0]
|
||||
noise = noise * latent_timestep / 1000 + (1 - latent_timestep / 1000) * input_samples
|
||||
else:
|
||||
noise = input_samples
|
||||
|
||||
# diff diff prep
|
||||
masks = None
|
||||
if noise_mask is not None:
|
||||
noise_mask = 1 - noise_mask
|
||||
thresholds = torch.arange(len(timesteps), dtype=original_image.dtype) / len(timesteps)
|
||||
thresholds = thresholds.unsqueeze(1).unsqueeze(1).unsqueeze(1).unsqueeze(1).to(device)
|
||||
masks = noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)
|
||||
masks = masks > thresholds
|
||||
|
||||
window_vace_data = None
|
||||
if vace_data is not None:
|
||||
window_vace_data = []
|
||||
for vace_entry in vace_data:
|
||||
partial_context = vace_entry["context"][0][:, latent_start_idx:latent_end_idx]
|
||||
if has_ref:
|
||||
partial_context[:, 0] = vace_entry["context"][0][:, 0]
|
||||
|
||||
window_vace_data.append({
|
||||
"context": [partial_context],
|
||||
"scale": vace_entry["scale"],
|
||||
"start": vace_entry["start"],
|
||||
"end": vace_entry["end"],
|
||||
"seq_len": vace_entry["seq_len"]
|
||||
})
|
||||
|
||||
# get image cond mask
|
||||
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
|
||||
if mode == "multitalk":
|
||||
msk[:, cur_motion_frames_num:] = 0
|
||||
@@ -3206,26 +3272,28 @@ class WanVideoSampler:
|
||||
mm.soft_empty_cache()
|
||||
|
||||
# zero padding and vae encode
|
||||
video_frames = torch.zeros(1, cond_image.shape[1], frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
|
||||
if cond_image is not None:
|
||||
video_frames = torch.zeros(1, 3, frame_num-cond_image.shape[2], target_h, target_w, device=device, dtype=vae.dtype)
|
||||
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
|
||||
|
||||
# encode
|
||||
vae.to(device)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae).to(dtype)
|
||||
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)
|
||||
|
||||
if mode == "multitalk":
|
||||
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
|
||||
else:
|
||||
if is_first_clip:
|
||||
latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
|
||||
latent_motion_frames = vae.encode(cond_image.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
|
||||
else:
|
||||
latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae).to(dtype)
|
||||
latent_motion_frames = vae.encode(cond_frame.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)
|
||||
latent_motion_frames = latent_motion_frames[0]
|
||||
vae.to(offload_device)
|
||||
|
||||
y = torch.concat([msk, y], dim=1) # B 4+C T H W
|
||||
y = torch.concat([msk, y], dim=1).squeeze(0) # 4+C T H W
|
||||
mm.soft_empty_cache()
|
||||
else:
|
||||
y = None
|
||||
latent_motion_frames = noise[:, :1]
|
||||
|
||||
if scheduler == "multitalk":
|
||||
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
|
||||
@@ -3235,6 +3303,24 @@ class WanVideoSampler:
|
||||
else:
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, transformer.dim, flowedit_args, denoise_strength, sigmas=sigmas)
|
||||
|
||||
steps = len(timesteps)
|
||||
if end_step != -1 and start_step >= end_step:
|
||||
raise ValueError("start_step must be less than end_step")
|
||||
if denoise_strength < 1.0:
|
||||
if start_step != 0:
|
||||
raise ValueError("start_step must be 0 when denoise_strength is used")
|
||||
start_step = steps - int(steps * denoise_strength) - 1
|
||||
if (end_step != -1 or end_step >= steps):
|
||||
timesteps = timesteps[:end_step]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[:end_step+1]
|
||||
if start_step > 0:
|
||||
timesteps = timesteps[start_step:]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_step:]
|
||||
|
||||
if sample_scheduler is not None:
|
||||
if hasattr(sample_scheduler, 'timesteps'):
|
||||
sample_scheduler.timesteps = timesteps
|
||||
|
||||
transformed_timesteps = []
|
||||
for t in timesteps:
|
||||
t_tensor = torch.tensor([t.item()], device=device)
|
||||
@@ -3287,8 +3373,16 @@ class WanVideoSampler:
|
||||
else:
|
||||
transformer.to(device)
|
||||
|
||||
comfy_pbar = ProgressBar(len(timesteps)-1)
|
||||
for i in tqdm(range(len(timesteps)-1)):
|
||||
# Use the appropriate prompt for this section
|
||||
if len(text_embeds["prompt_embeds"]) > 1:
|
||||
prompt_index = min(iteration_count, len(text_embeds["prompt_embeds"]) - 1)
|
||||
positive = [text_embeds["prompt_embeds"][prompt_index]]
|
||||
log.info(f"Using prompt index: {prompt_index}")
|
||||
else:
|
||||
positive = text_embeds["prompt_embeds"]
|
||||
|
||||
sampling_pbar = tqdm(total=len(timesteps)-1, desc=f"Sampling audio indices {audio_start_idx}-{audio_end_idx}", position=0, leave=True)
|
||||
for i in range(len(timesteps)-1):
|
||||
timestep = timesteps[i]
|
||||
latent_model_input = latent.to(device)
|
||||
if mode == "infinitetalk":
|
||||
@@ -3297,14 +3391,18 @@ class WanVideoSampler:
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
cfg[idx],
|
||||
text_embeds["prompt_embeds"],
|
||||
positive,
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, y.squeeze(0), clip_embeds, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
timestep, idx, y, clip_embeds, control_latents, window_vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs)
|
||||
|
||||
sampling_pbar.update(1)
|
||||
|
||||
if callback is not None:
|
||||
callback_latent = (latent_model_input.to(device) - noise_pred.to(device) * t.to(device) / 1000).detach().permute(1,0,2,3)
|
||||
callback(iteration_count, callback_latent, None, estimated_iterations)
|
||||
callback(step_iteration_count, callback_latent, None, estimated_iterations*(len(timesteps)-1))
|
||||
|
||||
step_iteration_count += 1
|
||||
|
||||
# update latent
|
||||
if scheduler == "multitalk":
|
||||
@@ -3314,7 +3412,6 @@ class WanVideoSampler:
|
||||
latent = latent + noise_pred * dt[:, None, None, None]
|
||||
else:
|
||||
latent = latent.to(intermediate_device)
|
||||
|
||||
temp_x0 = sample_scheduler.step(
|
||||
noise_pred.unsqueeze(0),
|
||||
timestep,
|
||||
@@ -3322,6 +3419,16 @@ class WanVideoSampler:
|
||||
**scheduler_step_args)[0]
|
||||
latent = temp_x0.squeeze(0)
|
||||
|
||||
# differential diffusion inpaint
|
||||
if masks is not None:
|
||||
if idx < len(timesteps) - 1:
|
||||
noise_timestep = timesteps[idx+1]
|
||||
image_latent = sample_scheduler.scale_noise(
|
||||
original_image, torch.tensor([noise_timestep]), noise.to(device)
|
||||
)
|
||||
mask = masks[idx].to(latent)
|
||||
latent = image_latent * mask + latent * (1-mask)
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip and mode == "multitalk":
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
@@ -3334,14 +3441,15 @@ class WanVideoSampler:
|
||||
|
||||
x0 = latent.to(device)
|
||||
del latent_model_input, timestep
|
||||
comfy_pbar.update(1)
|
||||
|
||||
if offload:
|
||||
offload_transformer(transformer)
|
||||
vae.to(device)
|
||||
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae)
|
||||
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device, tiled=tiled_vae, pbar=False)
|
||||
vae.to(offload_device)
|
||||
|
||||
sampling_pbar.close()
|
||||
|
||||
# cache generated samples
|
||||
videos = torch.stack(videos).cpu() # B C T H W
|
||||
if colormatch != "disabled":
|
||||
@@ -3411,9 +3519,7 @@ class WanVideoSampler:
|
||||
del noise, latent
|
||||
if force_offload:
|
||||
if not model["auto_cpu_offload"]:
|
||||
transformer.to(offload_device)
|
||||
mm.soft_empty_cache()
|
||||
gc.collect()
|
||||
offload_transformer(transformer)
|
||||
try:
|
||||
print_memory(device)
|
||||
torch.cuda.reset_peak_memory_stats(device)
|
||||
@@ -3486,6 +3592,25 @@ class WanVideoSampler:
|
||||
latent_backwards = torch.flip(latent_backwards, dims=[1])
|
||||
latent = latent * 0.5 + latent_backwards * 0.5
|
||||
|
||||
#InfiniteTalk first frame handling
|
||||
if (extra_latents is not None
|
||||
and not multitalk_sampling
|
||||
and transformer.multitalk_model_type=="InfiniteTalk"):
|
||||
for entry in extra_latents:
|
||||
add_index = entry["index"]
|
||||
num_extra_frames = entry["samples"].shape[2]
|
||||
latent[:, add_index:add_index+num_extra_frames] = entry["samples"].to(latent)
|
||||
|
||||
# differential diffusion inpaint
|
||||
if masks is not None:
|
||||
if idx < len(timesteps) - 1:
|
||||
noise_timestep = timesteps[idx+1]
|
||||
image_latent = sample_scheduler.scale_noise(
|
||||
original_image, torch.tensor([noise_timestep]), noise.to(device)
|
||||
)
|
||||
mask = masks[idx].to(latent)
|
||||
latent = image_latent * mask + latent * (1-mask)
|
||||
|
||||
if freeinit_args is not None:
|
||||
current_latent = latent.clone()
|
||||
|
||||
|
||||
+7
-11
@@ -503,7 +503,7 @@ class WanVideoVACEModelSelect:
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"vace_model": (folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' VACE model to use when not using model that has it included"}),
|
||||
"vace_model": (folder_paths.get_filename_list("unet_gguf") + folder_paths.get_filename_list("diffusion_models"), {"tooltip": "These models are loaded from the 'ComfyUI/models/diffusion_models' VACE model to use when not using model that has it included"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1021,15 +1021,9 @@ class WanVideoModelLoader:
|
||||
|
||||
if extra_model is not None:
|
||||
if gguf:
|
||||
if not extra_model["path"].endswith(".gguf"):
|
||||
raise ValueError("With GGUF main model the extra model must also be a GGUF quantized, if the main model already has extra included, you can disconnect the extra model loader")
|
||||
from diffusers.models.model_loading_utils import load_gguf_checkpoint
|
||||
extra_sd = load_gguf_checkpoint(extra_model["path"])
|
||||
new_keys = {}
|
||||
for k, v in extra_sd.items():
|
||||
if "vace" in k:
|
||||
new_keys[k] = v
|
||||
extra_sd = new_keys
|
||||
if not vace_model["path"].endswith(".gguf"):
|
||||
raise ValueError("With GGUF main model the VACE module must also be a GGUF quantized, if the main model already has VACE included, you can disconnect the VACE module loader")
|
||||
vace_sd = load_gguf_checkpoint(vace_model["path"])
|
||||
else:
|
||||
extra_sd = load_torch_file(extra_model["path"], device=transformer_load_device, safe_load=True)
|
||||
sd.update(extra_sd)
|
||||
@@ -1209,6 +1203,7 @@ class WanVideoModelLoader:
|
||||
if multitalk_model is not None:
|
||||
if multitalk_model["is_gguf"] and not gguf:
|
||||
raise ValueError("Multitalk/InfiniteTalk model is a GGUF model, main model also has to be a GGUF model.")
|
||||
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
|
||||
# init audio module
|
||||
from .multitalk.multitalk import SingleStreamMultiAttention
|
||||
from .wanvideo.modules.model import WanRMSNorm, WanLayerNorm
|
||||
@@ -1228,8 +1223,9 @@ class WanVideoModelLoader:
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True) if norm_input_visual else nn.Identity()
|
||||
log.info("MultiTalk model detected, patching model...")
|
||||
log.info(f"{multitalk_model_type} detected, patching model...")
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
transformer.multitalk_model_type = multitalk_model_type
|
||||
sd.update(multitalk_model["sd"])
|
||||
|
||||
# Additional cond latents
|
||||
|
||||
+9
-9
@@ -21,7 +21,7 @@ class WanVideoUni3C_ControlnetLoader:
|
||||
"model": (folder_paths.get_filename_list("controlnet"), {"tooltip": "These models are loaded from the 'ComfyUI/models/controlnet' -folder",}),
|
||||
|
||||
"base_precision": (["fp32", "bf16", "fp16"], {"default": "fp16"}),
|
||||
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e4m3fn_fast', 'fp8_e5m2', 'fp8_e4m3fn_fast_no_ffn'], {"default": 'disabled', "tooltip": "optional quantization method"}),
|
||||
"quantization": (['disabled', 'fp8_e4m3fn', 'fp8_e5m2'], {"default": 'disabled', "tooltip": "optional quantization method"}),
|
||||
"load_device": (["main_device", "offload_device"], {"default": "main_device", "tooltip": "Initial device to load the model to, NOT recommended with the larger models unless you have 48GB+ VRAM"}),
|
||||
"attention_mode": ([
|
||||
"sdpa",
|
||||
@@ -138,12 +138,12 @@ class WanVideoUni3C_embeds:
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"controlnet": ("WANVIDEOCONTROLNET",),
|
||||
"render_latent": ("LATENT",),
|
||||
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.001}),
|
||||
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply the controlnet"}),
|
||||
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply the controlnet"}),
|
||||
},
|
||||
"optional": {
|
||||
"render_latent": ("LATENT",),
|
||||
"render_mask": ("MASK", {"tooltip": "NOT IMPLEMENTED!"}),
|
||||
},
|
||||
}
|
||||
@@ -153,16 +153,16 @@ class WanVideoUni3C_embeds:
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, controlnet, render_latent, strength, start_percent, end_percent, render_mask=None):
|
||||
def process(self, controlnet, strength, start_percent, end_percent, render_latent=None, render_mask=None):
|
||||
|
||||
device = mm.get_torch_device()
|
||||
|
||||
latent_mask = None
|
||||
|
||||
latent_mask = latents = None
|
||||
if render_latent is not None:
|
||||
latents = render_latent["samples"]
|
||||
nframe = latents.shape[2] * 4
|
||||
height = latents.shape[3] * 8
|
||||
width = latents.shape[4] * 8
|
||||
# nframe = latents.shape[2] * 4
|
||||
# height = latents.shape[3] * 8
|
||||
# width = latents.shape[4] * 8
|
||||
|
||||
if render_mask is not None:
|
||||
raise NotImplementedError("render_mask is not implemented at this time")
|
||||
@@ -224,7 +224,7 @@ class WanVideoUni3C_embeds:
|
||||
"controlnet_weight": strength,
|
||||
"start": start_percent,
|
||||
"end": end_percent,
|
||||
"render_latent": latents.to(device),
|
||||
"render_latent": latents,
|
||||
"render_mask": latent_mask,
|
||||
"camera_embedding": None
|
||||
}
|
||||
|
||||
@@ -1305,6 +1305,8 @@ class WanModel(torch.nn.Module):
|
||||
self.video_attention_split_steps = []
|
||||
self.lora_scheduling_enabled = False
|
||||
|
||||
self.multitalk_model_type = "none"
|
||||
|
||||
# embeddings
|
||||
self.patch_embedding = nn.Conv3d(
|
||||
in_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||
@@ -1666,7 +1668,7 @@ class WanModel(torch.nn.Module):
|
||||
#uni3c controlnet
|
||||
if pcd_data is not None:
|
||||
hidden_states = x[0].unsqueeze(0).clone().float()
|
||||
render_latent = torch.cat([hidden_states[:, :20], pcd_data["render_latent"]], dim=1)
|
||||
render_latent = torch.cat([hidden_states[:, :20], pcd_data["render_latent"].to(x[0].dtype)], dim=1)
|
||||
|
||||
# embeddings
|
||||
if control_lora_enabled:
|
||||
|
||||
+40
-24
@@ -1019,12 +1019,13 @@ class VideoVAE_(nn.Module):
|
||||
return mu
|
||||
|
||||
|
||||
def encode(self, x):
|
||||
def encode(self, x, pbar=True):
|
||||
self.clear_cache()
|
||||
## cache
|
||||
pbar = ProgressBar(x.shape[2])
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
|
||||
for i in range(iter_):
|
||||
self._enc_conv_idx = [0]
|
||||
@@ -1037,10 +1038,13 @@ class VideoVAE_(nn.Module):
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx)
|
||||
out = torch.cat([out, out_], 2)
|
||||
pbar.update(iter_)
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
if pbar:
|
||||
pbar.update_absolute(0)
|
||||
|
||||
return mu
|
||||
|
||||
@@ -1076,14 +1080,13 @@ class VideoVAE_(nn.Module):
|
||||
|
||||
|
||||
|
||||
def decode(self, z):
|
||||
def decode(self, z, pbar=True):
|
||||
self.clear_cache()
|
||||
# z: [b,c,t,h,w]
|
||||
pbar = ProgressBar(z.shape[2])
|
||||
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
|
||||
iter_ = z.shape[2]
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
x = self.conv2(z)
|
||||
for i in range(iter_):
|
||||
self._conv_idx = [0]
|
||||
@@ -1096,7 +1099,11 @@ class VideoVAE_(nn.Module):
|
||||
feat_cache=self._feat_map,
|
||||
feat_idx=self._conv_idx)
|
||||
out = torch.cat([out, out_], 2) # may add tensor offload
|
||||
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
if pbar:
|
||||
pbar.update_absolute(0)
|
||||
return out
|
||||
|
||||
def reparameterize(self, mu, log_var):
|
||||
@@ -1167,7 +1174,7 @@ class WanVideoVAE(nn.Module):
|
||||
return mask
|
||||
|
||||
|
||||
def tiled_decode(self, hidden_states, device, tile_size, tile_stride):
|
||||
def tiled_decode(self, hidden_states, device, tile_size, tile_stride, pbar=True):
|
||||
_, _, T, H, W = hidden_states.shape
|
||||
size_h, size_w = tile_size
|
||||
stride_h, stride_w = tile_stride
|
||||
@@ -1187,7 +1194,7 @@ class WanVideoVAE(nn.Module):
|
||||
out_T = T * 4 - 3
|
||||
weight = torch.zeros((1, 1, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
values = torch.zeros((1, 3, out_T, H * self.upsampling_factor, W * self.upsampling_factor), dtype=hidden_states.dtype, device=data_device)
|
||||
|
||||
if pbar:
|
||||
pbar = ProgressBar(len(tasks))
|
||||
for h, h_, w, w_ in tqdm(tasks, desc="VAE decoding"):
|
||||
hidden_states_batch = hidden_states[:, :, :, h:h_, w:w_].to(computation_device)
|
||||
@@ -1215,13 +1222,14 @@ class WanVideoVAE(nn.Module):
|
||||
target_h: target_h + hidden_states_batch.shape[3],
|
||||
target_w: target_w + hidden_states_batch.shape[4],
|
||||
] += mask
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
values = values / weight
|
||||
values = values.float().clamp_(-1, 1)
|
||||
return values
|
||||
|
||||
|
||||
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False):
|
||||
def tiled_encode(self, video, device, tile_size, tile_stride, end_=False, pbar=True):
|
||||
_, _, T, H, W = video.shape
|
||||
|
||||
if tile_size is None and tile_stride is None:
|
||||
@@ -1248,7 +1256,7 @@ class WanVideoVAE(nn.Module):
|
||||
out_T += 1
|
||||
weight = torch.zeros((1, 1, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
|
||||
values = torch.zeros((1, self.z_dim, out_T, H // self.upsampling_factor, W // self.upsampling_factor), dtype=video.dtype, device=data_device)
|
||||
|
||||
if pbar:
|
||||
pbar = ProgressBar(len(tasks))
|
||||
for h, h_, w, w_ in tqdm(tasks, desc="VAE encoding"):
|
||||
hidden_states_batch = video[:, :, :, h:h_, w:w_].to(computation_device)
|
||||
@@ -1279,21 +1287,22 @@ class WanVideoVAE(nn.Module):
|
||||
target_h: target_h + hidden_states_batch.shape[3],
|
||||
target_w: target_w + hidden_states_batch.shape[4],
|
||||
] += mask
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
values = values / weight
|
||||
values = values.float()
|
||||
return values
|
||||
|
||||
|
||||
def single_encode(self, video, device):
|
||||
def single_encode(self, video, device, pbar=True):
|
||||
video = video.to(device)
|
||||
x = self.model.encode(video)
|
||||
x = self.model.encode(video, pbar=pbar)
|
||||
return x.float()
|
||||
|
||||
|
||||
def single_decode(self, hidden_state, device):
|
||||
def single_decode(self, hidden_state, device, pbar=True):
|
||||
hidden_state = hidden_state.to(device)
|
||||
video = self.model.decode(hidden_state)
|
||||
video = self.model.decode(hidden_state, pbar=pbar)
|
||||
return video
|
||||
|
||||
def double_encode(self, video, device):
|
||||
@@ -1308,36 +1317,36 @@ class WanVideoVAE(nn.Module):
|
||||
video = self.model.decode_2(hidden_state)
|
||||
return video
|
||||
|
||||
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None):
|
||||
def encode(self, videos, device, tiled=False,end_=False, tile_size=None, tile_stride=None, pbar=True):
|
||||
videos = [video.to("cpu") for video in videos]
|
||||
hidden_states = []
|
||||
for video in videos:
|
||||
video = video.unsqueeze(0)
|
||||
if tiled:
|
||||
hidden_state = self.tiled_encode(video, device, tile_size, tile_stride, end_=end_)
|
||||
hidden_state = self.tiled_encode(video, device, tile_size, tile_stride, end_=end_, pbar=pbar)
|
||||
else:
|
||||
if end_:
|
||||
hidden_state = self.double_encode(video, device)
|
||||
else:
|
||||
hidden_state = self.single_encode(video, device)
|
||||
hidden_state = self.single_encode(video, device, pbar=pbar)
|
||||
hidden_state = hidden_state.squeeze(0)
|
||||
hidden_states.append(hidden_state)
|
||||
hidden_states = torch.stack(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16)):
|
||||
def decode(self, hidden_states, device, tiled=False, end_=False, tile_size=(34, 34), tile_stride=(18, 16), pbar=True):
|
||||
hidden_states = [hidden_state.to("cpu") for hidden_state in hidden_states]
|
||||
videos = []
|
||||
for hidden_state in hidden_states:
|
||||
hidden_state = hidden_state.unsqueeze(0)
|
||||
if tiled:
|
||||
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride)
|
||||
video = self.tiled_decode(hidden_state, device, tile_size, tile_stride, pbar=pbar)
|
||||
else:
|
||||
if end_:
|
||||
video = self.double_decode(hidden_state, device)
|
||||
else:
|
||||
video = self.single_decode(hidden_state, device)
|
||||
video = self.single_decode(hidden_state, device, pbar=pbar)
|
||||
video = video.squeeze(0)
|
||||
videos.append(video)
|
||||
return videos
|
||||
@@ -1397,11 +1406,13 @@ class VideoVAE38_(VideoVAE_):
|
||||
attn_scales, self.temperal_upsample, dropout)
|
||||
|
||||
|
||||
def encode(self, x):
|
||||
def encode(self, x, pbar=True):
|
||||
self.clear_cache()
|
||||
x = patchify(x, patch_size=2)
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
for i in range(iter_):
|
||||
self._enc_conv_idx = [0]
|
||||
if i == 0:
|
||||
@@ -1413,6 +1424,8 @@ class VideoVAE38_(VideoVAE_):
|
||||
feat_cache=self._enc_feat_map,
|
||||
feat_idx=self._enc_conv_idx)
|
||||
out = torch.cat([out, out_], 2)
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
mu = self.conv1(out).chunk(2, dim=1)[0]
|
||||
|
||||
mu = (mu - self.mean.to(mu)) * self.inv_std.to(mu)
|
||||
@@ -1421,12 +1434,13 @@ class VideoVAE38_(VideoVAE_):
|
||||
return mu
|
||||
|
||||
|
||||
def decode(self, z):
|
||||
def decode(self, z, pbar=True):
|
||||
self.clear_cache()
|
||||
|
||||
z = z / self.inv_std.to(z) + self.mean.to(z)
|
||||
|
||||
iter_ = z.shape[2]
|
||||
if pbar:
|
||||
pbar = ProgressBar(iter_)
|
||||
x = self.conv2(z)
|
||||
for i in range(iter_):
|
||||
self._conv_idx = [0]
|
||||
@@ -1440,6 +1454,8 @@ class VideoVAE38_(VideoVAE_):
|
||||
feat_cache=self._feat_map,
|
||||
feat_idx=self._conv_idx)
|
||||
out = torch.cat([out, out_], 2)
|
||||
if pbar:
|
||||
pbar.update(1)
|
||||
out = unpatchify(out, patch_size=2)
|
||||
self.clear_cache()
|
||||
return out
|
||||
|
||||
Reference in New Issue
Block a user