MultiTalk sampling
New node (WanVideoImageToVideoMultiTalk) that also enables the continuous sampling method from the original code. Compared to context windows this works better for shorter clips, but degrades longer it goes.
This commit is contained in:
+7
-1
@@ -1,6 +1,5 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .recammaster.nodes import NODE_CLASS_MAPPINGS as RECAM_MASTER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .skyreels.nodes import NODE_CLASS_MAPPINGS as SKYREELS_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as SKYREELS_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .fantasytalking.nodes import NODE_CLASS_MAPPINGS as FANTASYTALKING_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FANTASYTALKING_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .fun_camera.nodes import NODE_CLASS_MAPPINGS as FUN_CAMERA_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as FUN_CAMERA_NODE_DISPLAY_NAME_MAPPINGS
|
||||
@@ -9,6 +8,13 @@ from .controlnet.nodes import NODE_CLASS_MAPPINGS as CONTROLNET_NODE_CLASS_MAPPI
|
||||
from .ATI.nodes import NODE_CLASS_MAPPINGS as ATI_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ATI_NODE_DISPLAY_NAME_MAPPINGS
|
||||
from .multitalk.nodes import NODE_CLASS_MAPPINGS as MULTITALK_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as MULTITALK_NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
try:
|
||||
from .unianimate.nodes import NODE_CLASS_MAPPINGS as UNIANIMATE_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS
|
||||
except ImportError:
|
||||
print("UniAnimate not available due to missing dependencies")
|
||||
UNIANIMATE_NODE_CLASS_MAPPINGS = {}
|
||||
UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
|
||||
#from .causvid.nodes import NODE_CLASS_MAPPINGS as CAUSVID_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as CAUSVID_NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
|
||||
|
||||
@@ -7,6 +7,30 @@ from ..wanvideo.modules.attention import attention
|
||||
|
||||
from comfy import model_management as mm
|
||||
|
||||
def timestep_transform(
|
||||
t,
|
||||
shift=5.0,
|
||||
num_timesteps=1000,
|
||||
):
|
||||
t = t / num_timesteps
|
||||
# shift the timestep based on ratio
|
||||
new_t = shift * t / (1 + (shift - 1) * t)
|
||||
new_t = new_t * num_timesteps
|
||||
return new_t
|
||||
|
||||
def add_noise(
|
||||
original_samples: torch.FloatTensor,
|
||||
noise: torch.FloatTensor,
|
||||
timesteps: torch.IntTensor,
|
||||
) -> torch.FloatTensor:
|
||||
"""
|
||||
compatible with diffusers add_noise()
|
||||
"""
|
||||
timesteps = timesteps.float() / 1000
|
||||
timesteps = timesteps.view(timesteps.shape + (1,) * (len(noise.shape)-1))
|
||||
|
||||
return (1 - timesteps) * original_samples + timesteps * noise
|
||||
|
||||
def normalize_and_scale(column, source_range, target_range, epsilon=1e-8):
|
||||
|
||||
source_min, source_max = source_range
|
||||
|
||||
+74
-1
@@ -1,10 +1,11 @@
|
||||
import folder_paths
|
||||
from comfy import model_management as mm
|
||||
from comfy.utils import load_torch_file
|
||||
from comfy.utils import load_torch_file, common_upscale
|
||||
from accelerate import init_empty_weights
|
||||
from accelerate.utils import set_module_tensor_to_device
|
||||
import torch
|
||||
|
||||
|
||||
class MultiTalkModelLoader:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
@@ -161,14 +162,86 @@ class MultiTalkWav2VecEmbeds:
|
||||
}
|
||||
|
||||
return (multitalk_embeds, audio_output)
|
||||
|
||||
class WanVideoImageToVideoMultiTalk:
|
||||
@classmethod
|
||||
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}),
|
||||
"colormatch": (
|
||||
[
|
||||
'disabled',
|
||||
'mkl',
|
||||
'hm',
|
||||
'reinhard',
|
||||
'mvgd',
|
||||
'hm-mvgd-hm',
|
||||
'hm-mkl-hm',
|
||||
], {
|
||||
"default": 'disabled'
|
||||
}),
|
||||
},
|
||||
"optional": {
|
||||
"start_image": ("IMAGE", {"tooltip": "Image 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"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "process"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, vae, width, height, frame_window_size, motion_frame, force_offload, colormatch, start_image=None, tiled_vae=False, clip_embeds=None):
|
||||
|
||||
H = height
|
||||
W = width
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
|
||||
num_frames = ((frame_window_size - 1) // 4) * 4 + 1
|
||||
|
||||
# Resize and rearrange the input image dimensions
|
||||
if start_image is not None:
|
||||
resized_start_image = common_upscale(start_image.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
|
||||
resized_start_image = resized_start_image * 2 - 1
|
||||
resized_start_image = resized_start_image.unsqueeze(0)
|
||||
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
height // VAE_STRIDE[1],
|
||||
width // VAE_STRIDE[2])
|
||||
|
||||
image_embeds = {
|
||||
"multitalk_sampling": True,
|
||||
"multitalk_start_image": resized_start_image if start_image is not None else None,
|
||||
"num_frames": num_frames,
|
||||
"motion_frame": motion_frame,
|
||||
"target_h": H,
|
||||
"target_w": W,
|
||||
"tiled_vae": tiled_vae,
|
||||
"force_offload": force_offload,
|
||||
"vae": vae,
|
||||
"target_shape": target_shape,
|
||||
"clip_context": clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None,
|
||||
"colormatch": colormatch
|
||||
}
|
||||
|
||||
return (image_embeds,)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"MultiTalkModelLoader": MultiTalkModelLoader,
|
||||
"MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds,
|
||||
"WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk
|
||||
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"MultiTalkModelLoader": "MultiTalk Model Loader",
|
||||
"MultiTalkWav2VecEmbeds": "MultiTalk Wav2Vec Embeds",
|
||||
"WanVideoImageToVideoMultiTalk": "WanVideo Image To Video MultiTalk"
|
||||
}
|
||||
@@ -17,6 +17,8 @@ from .wanvideo.utils.basic_flowmatch import FlowMatchScheduler
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler
|
||||
from .wanvideo.utils.scheduling_flow_match_lcm import FlowMatchLCMScheduler
|
||||
|
||||
from .multitalk.multitalk import timestep_transform, add_noise
|
||||
|
||||
from .enhance_a_video.globals import enable_enhance, disable_enhance, set_enhance_weight, set_num_frames
|
||||
from .taehv import TAEHV
|
||||
|
||||
@@ -28,13 +30,15 @@ import folder_paths
|
||||
import comfy.model_management as mm
|
||||
from comfy.utils import load_torch_file, ProgressBar, common_upscale
|
||||
import comfy.model_base
|
||||
import comfy.latent_formats
|
||||
from comfy.clip_vision import clip_preprocess, ClipVisionModel
|
||||
from comfy.sd import load_lora_for_models
|
||||
from comfy.cli_args import args, LatentPreviewMethod
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
|
||||
VAE_STRIDE = (4, 8, 8)
|
||||
PATCH_SIZE = (1, 2, 2)
|
||||
|
||||
def add_noise_to_reference_video(image, ratio=None):
|
||||
sigma = torch.ones((image.shape[0],)).to(image.device, image.dtype) * ratio
|
||||
image_noise = torch.randn_like(image) * sigma[:, None, None, None]
|
||||
@@ -1555,8 +1559,6 @@ class WanVideoImageClipEncode:
|
||||
|
||||
self.image_mean = [0.48145466, 0.4578275, 0.40821073]
|
||||
self.image_std = [0.26862954, 0.26130258, 0.27577711]
|
||||
patch_size = (1, 2, 2)
|
||||
vae_stride = (4, 8, 8)
|
||||
|
||||
H, W = image.shape[1], image.shape[2]
|
||||
max_area = generation_width * generation_height
|
||||
@@ -1579,13 +1581,13 @@ class WanVideoImageClipEncode:
|
||||
if adjust_resolution:
|
||||
aspect_ratio = H / W
|
||||
lat_h = round(
|
||||
np.sqrt(max_area * aspect_ratio) // vae_stride[1] //
|
||||
patch_size[1] * patch_size[1])
|
||||
np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] //
|
||||
PATCH_SIZE[1] * PATCH_SIZE[1])
|
||||
lat_w = round(
|
||||
np.sqrt(max_area / aspect_ratio) // vae_stride[2] //
|
||||
patch_size[2] * patch_size[2])
|
||||
h = lat_h * vae_stride[1]
|
||||
w = lat_w * vae_stride[2]
|
||||
np.sqrt(max_area / aspect_ratio) // VAE_STRIDE[2] //
|
||||
PATCH_SIZE[2] * PATCH_SIZE[2])
|
||||
h = lat_h * VAE_STRIDE[1]
|
||||
w = lat_w * VAE_STRIDE[2]
|
||||
else:
|
||||
h = generation_height
|
||||
w = generation_width
|
||||
@@ -1607,8 +1609,8 @@ class WanVideoImageClipEncode:
|
||||
mask = mask.transpose(1, 2)[0]
|
||||
|
||||
# Calculate maximum sequence length
|
||||
frames_per_stride = (num_frames - 1) // vae_stride[0] + 1
|
||||
patches_per_frame = lat_h * lat_w // (patch_size[1] * patch_size[2])
|
||||
frames_per_stride = (num_frames - 1) // VAE_STRIDE[0] + 1
|
||||
patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
|
||||
max_seq_len = frames_per_stride * patches_per_frame
|
||||
|
||||
vae.to(device)
|
||||
@@ -1665,9 +1667,6 @@ class WanVideoImageResizeToClosest:
|
||||
DESCRIPTION = "Resizes image to the closest supported resolution based on aspect ratio and max pixels, according to the original code"
|
||||
|
||||
def process(self, image, generation_width, generation_height, aspect_ratio_preservation ):
|
||||
|
||||
patch_size = (1, 2, 2)
|
||||
vae_stride = (4, 8, 8)
|
||||
|
||||
H, W = image.shape[1], image.shape[2]
|
||||
max_area = generation_width * generation_height
|
||||
@@ -1682,13 +1681,13 @@ class WanVideoImageResizeToClosest:
|
||||
crop = "center"
|
||||
|
||||
lat_h = round(
|
||||
np.sqrt(max_area * aspect_ratio) // vae_stride[1] //
|
||||
patch_size[1] * patch_size[1])
|
||||
np.sqrt(max_area * aspect_ratio) // VAE_STRIDE[1] //
|
||||
PATCH_SIZE[1] * PATCH_SIZE[1])
|
||||
lat_w = round(
|
||||
np.sqrt(max_area / aspect_ratio) // vae_stride[2] //
|
||||
patch_size[2] * patch_size[2])
|
||||
h = lat_h * vae_stride[1]
|
||||
w = lat_w * vae_stride[2]
|
||||
np.sqrt(max_area / aspect_ratio) // VAE_STRIDE[2] //
|
||||
PATCH_SIZE[2] * PATCH_SIZE[2])
|
||||
h = lat_h * VAE_STRIDE[1]
|
||||
w = lat_w * VAE_STRIDE[2]
|
||||
|
||||
resized_image = common_upscale(image.movedim(-1, 1), w, h, "lanczos", crop).movedim(1, -1)
|
||||
|
||||
@@ -1867,8 +1866,6 @@ class WanVideoImageToVideoEncode:
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
patch_size = (1, 2, 2)
|
||||
|
||||
H = height
|
||||
W = width
|
||||
|
||||
@@ -1961,7 +1958,7 @@ class WanVideoImageToVideoEncode:
|
||||
y[:, 1:-1] = 0 # doesn't seem to work anyway though...
|
||||
|
||||
# Calculate maximum sequence length
|
||||
patches_per_frame = lat_h * lat_w // (patch_size[1] * patch_size[2])
|
||||
patches_per_frame = lat_h * lat_w // (PATCH_SIZE[1] * PATCH_SIZE[2])
|
||||
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
|
||||
|
||||
@@ -2010,11 +2007,9 @@ class WanVideoEmptyEmbeds:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, num_frames, width, height, control_embeds=None):
|
||||
vae_stride = (4, 8, 8)
|
||||
|
||||
target_shape = (16, (num_frames - 1) // vae_stride[0] + 1,
|
||||
height // vae_stride[1],
|
||||
width // vae_stride[2])
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
height // VAE_STRIDE[1],
|
||||
width // VAE_STRIDE[2])
|
||||
|
||||
embeds = {
|
||||
"target_shape": target_shape,
|
||||
@@ -2042,11 +2037,9 @@ class WanVideoMiniMaxRemoverEmbeds:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, num_frames, width, height, latents, mask_latents):
|
||||
vae_stride = (4, 8, 8)
|
||||
|
||||
target_shape = (16, (num_frames - 1) // vae_stride[0] + 1,
|
||||
height // vae_stride[1],
|
||||
width // vae_stride[2])
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
height // VAE_STRIDE[1],
|
||||
width // VAE_STRIDE[2])
|
||||
|
||||
embeds = {
|
||||
"target_shape": target_shape,
|
||||
@@ -2083,7 +2076,6 @@ class WanVideoPhantomEmbeds:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def process(self, num_frames, phantom_cfg_scale, phantom_start_percent, phantom_end_percent, phantom_latent_1, phantom_latent_2=None, phantom_latent_3=None, phantom_latent_4=None, vace_embeds=None):
|
||||
vae_stride = (4, 8, 8)
|
||||
samples = phantom_latent_1["samples"].squeeze(0)
|
||||
if phantom_latent_2 is not None:
|
||||
samples = torch.cat([samples, phantom_latent_2["samples"].squeeze(0)], dim=1)
|
||||
@@ -2095,9 +2087,9 @@ class WanVideoPhantomEmbeds:
|
||||
|
||||
log.info(f"Phantom latents shape: {samples.shape}")
|
||||
|
||||
target_shape = (16, (num_frames - 1) // vae_stride[0] + 1 + T,
|
||||
H * 8 // vae_stride[1],
|
||||
W * 8 // vae_stride[2])
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1 + T,
|
||||
H * 8 // VAE_STRIDE[1],
|
||||
W * 8 // VAE_STRIDE[2])
|
||||
|
||||
embeds = {
|
||||
"target_shape": target_shape,
|
||||
@@ -2219,14 +2211,13 @@ class WanVideoVACEEncode:
|
||||
self.device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
self.vae = vae.to(self.device)
|
||||
self.vae_stride = (4, 8, 8)
|
||||
|
||||
width = (width // 16) * 16
|
||||
height = (height // 16) * 16
|
||||
|
||||
target_shape = (16, (num_frames - 1) // self.vae_stride[0] + 1,
|
||||
height // self.vae_stride[1],
|
||||
width // self.vae_stride[2])
|
||||
target_shape = (16, (num_frames - 1) // VAE_STRIDE[0] + 1,
|
||||
height // VAE_STRIDE[1],
|
||||
width // VAE_STRIDE[2])
|
||||
# vace context encode
|
||||
if input_frames is None:
|
||||
input_frames = torch.zeros((1, 3, num_frames, height, width), device=self.device, dtype=self.vae.dtype)
|
||||
@@ -2334,18 +2325,18 @@ class WanVideoVACEEncode:
|
||||
result_masks = []
|
||||
for mask, refs in zip(masks, ref_images):
|
||||
c, depth, height, width = mask.shape
|
||||
new_depth = int((depth + 3) // self.vae_stride[0])
|
||||
height = 2 * (int(height) // (self.vae_stride[1] * 2))
|
||||
width = 2 * (int(width) // (self.vae_stride[2] * 2))
|
||||
new_depth = int((depth + 3) // VAE_STRIDE[0])
|
||||
height = 2 * (int(height) // (VAE_STRIDE[1] * 2))
|
||||
width = 2 * (int(width) // (VAE_STRIDE[2] * 2))
|
||||
|
||||
# reshape
|
||||
mask = mask[0, :, :, :]
|
||||
mask = mask.view(
|
||||
depth, height, self.vae_stride[1], width, self.vae_stride[1]
|
||||
depth, height, VAE_STRIDE[1], width, VAE_STRIDE[1]
|
||||
) # depth, height, 8, width, 8
|
||||
mask = mask.permute(2, 4, 0, 1, 3) # 8, 8, depth, height, width
|
||||
mask = mask.reshape(
|
||||
self.vae_stride[1] * self.vae_stride[2], depth, height, width
|
||||
VAE_STRIDE[1] * VAE_STRIDE[2], depth, height, width
|
||||
) # 8*8, depth, height, width
|
||||
|
||||
# interpolation
|
||||
@@ -2599,7 +2590,7 @@ class WanVideoSampler:
|
||||
"shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}),
|
||||
"seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}),
|
||||
"force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}),
|
||||
"scheduler": (["unipc", "unipc/beta", "dpm++", "dpm++/beta","dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", "euler/accvideo", "deis", "lcm", "lcm/beta", "flowmatch_causvid", "flowmatch_distill"],
|
||||
"scheduler": (["unipc", "unipc/beta", "dpm++", "dpm++/beta","dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", "euler/accvideo", "deis", "lcm", "lcm/beta", "flowmatch_causvid", "flowmatch_distill", "multitalk"],
|
||||
{
|
||||
"default": 'unipc'
|
||||
}),
|
||||
@@ -2644,6 +2635,10 @@ class WanVideoSampler:
|
||||
dtype = model["dtype"]
|
||||
control_lora = model["control_lora"]
|
||||
|
||||
multitalk_sampling = image_embeds.get("multitalk_sampling", False)
|
||||
if not multitalk_sampling and scheduler == "multitalk":
|
||||
raise Exception("multitalk scheduler is only for multitalk sampling when using ImagetoVideoMultiTalk -node")
|
||||
|
||||
transformer_options = patcher.model_options.get("transformer_options", None)
|
||||
|
||||
device = mm.get_torch_device()
|
||||
@@ -2662,82 +2657,89 @@ class WanVideoSampler:
|
||||
log.info(f"Received {len(cfg)} cfg values, but only {steps} steps. Setting step count to match.")
|
||||
steps = len(cfg)
|
||||
|
||||
timesteps = None
|
||||
if 'unipc' in scheduler:
|
||||
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
|
||||
else:
|
||||
sample_scheduler.sigmas = sigmas.to(device)
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
|
||||
elif scheduler in ['euler/beta', 'euler']:
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
|
||||
if flowedit_args: #seems to work better
|
||||
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
|
||||
else:
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
elif scheduler in ['euler/accvideo']:
|
||||
if steps != 50:
|
||||
raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used")
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list]
|
||||
steps = len(start_latent_list) - 1
|
||||
sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:steps]]
|
||||
elif 'dpm++' in scheduler:
|
||||
if 'sde' in scheduler:
|
||||
algorithm_type = "sde-dpmsolver++"
|
||||
else:
|
||||
algorithm_type = "dpmsolver++"
|
||||
sample_scheduler = FlowDPMSolverMultistepScheduler(shift=shift, algorithm_type=algorithm_type)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler))
|
||||
else:
|
||||
sample_scheduler.sigmas = sigmas.to(device)
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
elif scheduler == 'deis':
|
||||
sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift)
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
sample_scheduler.sigmas[-1] = 1e-6
|
||||
elif 'lcm' in scheduler:
|
||||
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
elif 'flowmatch_causvid' in scheduler:
|
||||
if transformer.dim == 5120:
|
||||
denoising_list = [999, 934, 862, 756, 603, 410, 250, 140, 74]
|
||||
else:
|
||||
if steps != 4:
|
||||
raise ValueError("CausVid 1.3B schedule is only for 4 steps")
|
||||
denoising_list = [1000, 750, 500, 250]
|
||||
sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True)
|
||||
sample_scheduler.timesteps = torch.tensor(denoising_list)[:steps].to(device)
|
||||
sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
|
||||
elif 'flowmatch_distill' in scheduler:
|
||||
sample_scheduler = FlowMatchScheduler(
|
||||
shift=shift, sigma_min=0.0, extra_one_step=True
|
||||
)
|
||||
sample_scheduler.set_timesteps(1000, training=True)
|
||||
|
||||
denoising_step_list = torch.tensor([999, 750, 500, 250] , dtype=torch.long)
|
||||
temp_timesteps = torch.cat((sample_scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
|
||||
denoising_step_list = temp_timesteps[1000 - denoising_step_list]
|
||||
print("denoising_step_list: ", denoising_step_list)
|
||||
|
||||
|
||||
#denoising_step_list = [999, 750, 500, 250]
|
||||
if steps != 4:
|
||||
raise ValueError("This scheduler is only for 4 steps")
|
||||
#sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True)
|
||||
sample_scheduler.timesteps = torch.tensor(denoising_step_list)[:steps].to(device)
|
||||
sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
|
||||
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
log.info(f"timesteps: {timesteps}")
|
||||
def get_scheduler(scheduler, steps, shift, device, sigmas=None):
|
||||
timesteps = None
|
||||
if 'unipc' in scheduler:
|
||||
sample_scheduler = FlowUniPCMultistepScheduler(shift=shift)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device, shift=shift, use_beta_sigmas=('beta' in scheduler))
|
||||
else:
|
||||
sample_scheduler.sigmas = sigmas.to(device)
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
|
||||
elif scheduler in ['euler/beta', 'euler']:
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
|
||||
if flowedit_args: #seems to work better
|
||||
timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift))
|
||||
else:
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
elif scheduler in ['euler/accvideo']:
|
||||
if steps != 50:
|
||||
raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used")
|
||||
sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta'))
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50]
|
||||
sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list]
|
||||
steps = len(start_latent_list) - 1
|
||||
sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:steps]]
|
||||
elif 'dpm++' in scheduler:
|
||||
if 'sde' in scheduler:
|
||||
algorithm_type = "sde-dpmsolver++"
|
||||
else:
|
||||
algorithm_type = "dpmsolver++"
|
||||
sample_scheduler = FlowDPMSolverMultistepScheduler(shift=shift, algorithm_type=algorithm_type)
|
||||
if sigmas is None:
|
||||
sample_scheduler.set_timesteps(steps, device=device, use_beta_sigmas=('beta' in scheduler))
|
||||
else:
|
||||
sample_scheduler.sigmas = sigmas.to(device)
|
||||
sample_scheduler.timesteps = (sample_scheduler.sigmas[:-1] * 1000).to(torch.int64).to(device)
|
||||
sample_scheduler.num_inference_steps = len(sample_scheduler.timesteps)
|
||||
elif scheduler == 'deis':
|
||||
sample_scheduler = DEISMultistepScheduler(use_flow_sigmas=True, prediction_type="flow_prediction", flow_shift=shift)
|
||||
sample_scheduler.set_timesteps(steps, device=device)
|
||||
sample_scheduler.sigmas[-1] = 1e-6
|
||||
elif 'lcm' in scheduler:
|
||||
sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta'))
|
||||
sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None)
|
||||
elif 'flowmatch_causvid' in scheduler:
|
||||
if transformer.dim == 5120:
|
||||
denoising_list = [999, 934, 862, 756, 603, 410, 250, 140, 74]
|
||||
else:
|
||||
if steps != 4:
|
||||
raise ValueError("CausVid 1.3B schedule is only for 4 steps")
|
||||
denoising_list = [1000, 750, 500, 250]
|
||||
sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True)
|
||||
sample_scheduler.timesteps = torch.tensor(denoising_list)[:steps].to(device)
|
||||
sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
|
||||
elif 'flowmatch_distill' in scheduler:
|
||||
sample_scheduler = FlowMatchScheduler(
|
||||
shift=shift, sigma_min=0.0, extra_one_step=True
|
||||
)
|
||||
sample_scheduler.set_timesteps(1000, training=True)
|
||||
|
||||
denoising_step_list = torch.tensor([999, 750, 500, 250] , dtype=torch.long)
|
||||
temp_timesteps = torch.cat((sample_scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
|
||||
denoising_step_list = temp_timesteps[1000 - denoising_step_list]
|
||||
print("denoising_step_list: ", denoising_step_list)
|
||||
|
||||
|
||||
#denoising_step_list = [999, 750, 500, 250]
|
||||
if steps != 4:
|
||||
raise ValueError("This scheduler is only for 4 steps")
|
||||
#sample_scheduler = FlowMatchScheduler(num_inference_steps=steps, shift=shift, sigma_min=0, extra_one_step=True)
|
||||
sample_scheduler.timesteps = torch.tensor(denoising_step_list)[:steps].to(device)
|
||||
sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)])
|
||||
return sample_scheduler, timesteps
|
||||
|
||||
if scheduler != "multitalk":
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas)
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
log.info(f"timesteps: {timesteps}")
|
||||
else:
|
||||
timesteps = torch.tensor([1000, 750, 500, 250], device=device)
|
||||
|
||||
if denoise_strength < 1.0:
|
||||
steps = int(steps * denoise_strength)
|
||||
@@ -3250,7 +3252,7 @@ class WanVideoSampler:
|
||||
#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,
|
||||
add_cond=None, cache_state=None, context_window=None):
|
||||
add_cond=None, cache_state=None, context_window=None, multitalk_audio_embeds=None):
|
||||
z = z.to(dtype)
|
||||
with torch.autocast(device_type=mm.get_autocast_device(device), dtype=dtype, enabled=("fp8" in model["quantization"])):
|
||||
|
||||
@@ -3345,7 +3347,7 @@ class WanVideoSampler:
|
||||
if minimax_latents is not None:
|
||||
z_pos = z_neg = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
|
||||
|
||||
if multitalk_audio_embedding is not None:
|
||||
if not multitalk_sampling and multitalk_audio_embedding is not None:
|
||||
audio_embedding = [multitalk_audio_embedding]
|
||||
audio_embs = []
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
@@ -3371,7 +3373,10 @@ class WanVideoSampler:
|
||||
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0] - 1)
|
||||
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
|
||||
audio_embs.append(audio_emb)
|
||||
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
|
||||
multitalk_audio_input = torch.concat(audio_embs, dim=0).to(dtype)
|
||||
|
||||
elif multitalk_sampling and multitalk_audio_embeds is not None:
|
||||
multitalk_audio_input = multitalk_audio_embeds
|
||||
|
||||
if context_window is not None and pcd_data is not None and pcd_data["render_latent"].shape[2] != context_frames:
|
||||
pcd_data_input = {"render_latent": pcd_data["render_latent"][:, :, context_window]}
|
||||
@@ -3401,7 +3406,7 @@ class WanVideoSampler:
|
||||
"add_cond": add_cond_input,
|
||||
"nag_params": text_embeds.get("nag_params", {}),
|
||||
"nag_context": text_embeds.get("nag_prompt_embeds", None),
|
||||
"multitalk_audio": audio_embs if multitalk_audio_embedding is not None else None,
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None,
|
||||
}
|
||||
|
||||
batch_size = 1
|
||||
@@ -3810,6 +3815,251 @@ class WanVideoSampler:
|
||||
counter[:, c] += window_mask
|
||||
context_pbar.update(1)
|
||||
noise_pred /= counter
|
||||
#region multitalk
|
||||
elif multitalk_sampling:
|
||||
original_image = cond_image = image_embeds.get("multitalk_start_image", None)
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
frame_num = clip_length = image_embeds.get("num_frames", 81)
|
||||
vae = image_embeds.get("vae", None)
|
||||
clip_embeds = image_embeds.get("clip_context", None)
|
||||
colormatch = image_embeds.get("colormatch", "disabled")
|
||||
motion_frame = image_embeds.get("motion_frame", 25)
|
||||
target_w = image_embeds.get("target_w", None)
|
||||
target_h = image_embeds.get("target_h", None)
|
||||
|
||||
max_frames_num = 1000
|
||||
gen_video_list = []
|
||||
is_first_clip = True
|
||||
arrive_last_frame = False
|
||||
cur_motion_frames_num = 1
|
||||
audio_start_idx = 0
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
|
||||
# start video generation iteratively
|
||||
audio_embedding = [multitalk_audio_embedding]
|
||||
while True:
|
||||
audio_embs = []
|
||||
# split audio with window size
|
||||
for human_idx in range(1):
|
||||
center_indices = torch.arange(audio_start_idx, audio_end_idx, 1).unsqueeze(1) + indices.unsqueeze(0)
|
||||
center_indices = torch.clamp(center_indices, min=0, max=audio_embedding[human_idx].shape[0]-1)
|
||||
audio_emb = audio_embedding[human_idx][center_indices].unsqueeze(0).to(device)
|
||||
audio_embs.append(audio_emb)
|
||||
audio_embs = torch.concat(audio_embs, dim=0).to(dtype)
|
||||
|
||||
h, w = cond_image.shape[-2], cond_image.shape[-1]
|
||||
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])
|
||||
|
||||
noise = torch.randn(
|
||||
16, (frame_num - 1) // 4 + 1,
|
||||
lat_h,
|
||||
lat_w,
|
||||
dtype=torch.float32,
|
||||
device=device)
|
||||
|
||||
# get mask
|
||||
msk = torch.ones(1, frame_num, lat_h, lat_w, device=device)
|
||||
msk[:, cur_motion_frames_num:] = 0
|
||||
msk = torch.concat([
|
||||
torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]
|
||||
],
|
||||
dim=1)
|
||||
msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w)
|
||||
msk = msk.transpose(1, 2).to(dtype) # B 4 T H W
|
||||
|
||||
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)
|
||||
padding_frames_pixels_values = torch.concat([cond_image.to(device, vae.dtype), video_frames], dim=2)
|
||||
|
||||
vae.to(device)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device).to(dtype)
|
||||
vae.to(offload_device)
|
||||
|
||||
cur_motion_frames_latent_num = int(1 + (cur_motion_frames_num-1) // 4)
|
||||
latent_motion_frames = y[:, :, :cur_motion_frames_latent_num][0] # C T H W
|
||||
y = torch.concat([msk, y], dim=1) # B 4+C T H W
|
||||
mm.soft_empty_cache()
|
||||
|
||||
if scheduler == "multitalk":
|
||||
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
|
||||
timesteps.append(0.)
|
||||
timesteps = [torch.tensor([t], device=device) for t in timesteps]
|
||||
timesteps = [timestep_transform(t, shift=shift, num_timesteps=1000) for t in timesteps]
|
||||
else:
|
||||
sample_scheduler, timesteps = get_scheduler(scheduler, steps, shift, device, sigmas=sigmas)
|
||||
if timesteps is None:
|
||||
timesteps = sample_scheduler.timesteps
|
||||
|
||||
transformed_timesteps = []
|
||||
for t in timesteps:
|
||||
t_tensor = torch.tensor([t.item()], device=device)
|
||||
transformed_timesteps.append(t_tensor)
|
||||
|
||||
transformed_timesteps.append(torch.tensor([0.], device=device))
|
||||
timesteps = transformed_timesteps
|
||||
|
||||
# sample videos
|
||||
latent = noise
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip:
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn_like(latent_motion_frames).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
|
||||
_, T_m, _, _ = add_latent.shape
|
||||
latent[:, :T_m] = add_latent
|
||||
|
||||
if offload:
|
||||
#blockswap init
|
||||
if transformer_options is not None:
|
||||
block_swap_args = transformer_options.get("block_swap_args", None)
|
||||
|
||||
if block_swap_args is not None:
|
||||
transformer.use_non_blocking = block_swap_args.get("use_non_blocking", True)
|
||||
for name, param in transformer.named_parameters():
|
||||
if "block" not in name:
|
||||
param.data = param.data.to(device)
|
||||
if "control_adapter" in name:
|
||||
param.data = param.data.to(device)
|
||||
elif block_swap_args["offload_txt_emb"] and "txt_emb" in name:
|
||||
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
|
||||
elif block_swap_args["offload_img_emb"] and "img_emb" in name:
|
||||
param.data = param.data.to(offload_device, non_blocking=transformer.use_non_blocking)
|
||||
|
||||
transformer.block_swap(
|
||||
block_swap_args["blocks_to_swap"] - 1 ,
|
||||
block_swap_args["offload_txt_emb"],
|
||||
block_swap_args["offload_img_emb"],
|
||||
vace_blocks_to_swap = block_swap_args.get("vace_blocks_to_swap", None),
|
||||
)
|
||||
|
||||
elif model["auto_cpu_offload"]:
|
||||
for module in transformer.modules():
|
||||
if hasattr(module, "offload"):
|
||||
module.offload()
|
||||
if hasattr(module, "onload"):
|
||||
module.onload()
|
||||
elif model["manual_offloading"]:
|
||||
transformer.to(device)
|
||||
|
||||
for i in tqdm(range(len(timesteps)-1)):
|
||||
timestep = timesteps[i]
|
||||
latent_model_input = latent.to(device)
|
||||
print("latent_model_input shape: ", latent_model_input.shape)
|
||||
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input,
|
||||
cfg[idx],
|
||||
text_embeds["prompt_embeds"],
|
||||
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,
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs)
|
||||
|
||||
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(idx, callback_latent, None, steps)
|
||||
else:
|
||||
pbar.update(1)
|
||||
|
||||
# update latent
|
||||
if scheduler == "multitalk":
|
||||
noise_pred = -noise_pred
|
||||
dt = timesteps[i] - timesteps[i + 1]
|
||||
dt = dt / 1000
|
||||
latent = latent + noise_pred * dt[:, None, None, None]
|
||||
else:
|
||||
latent = latent.to(intermediate_device)
|
||||
step_args = {
|
||||
"generator": seed_g,
|
||||
}
|
||||
if isinstance(sample_scheduler, DEISMultistepScheduler) or isinstance(sample_scheduler, FlowMatchScheduler):
|
||||
step_args.pop("generator", None)
|
||||
temp_x0 = sample_scheduler.step(
|
||||
noise_pred.unsqueeze(0),
|
||||
timestep,
|
||||
latent.unsqueeze(0),
|
||||
#return_dict=False,
|
||||
**step_args)[0]
|
||||
latent = temp_x0.squeeze(0)
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip:
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn_like(latent_motion_frames).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
|
||||
_, T_m, _, _ = add_latent.shape
|
||||
latent[:, :T_m] = add_latent
|
||||
|
||||
x0 = latent.to(device)
|
||||
del latent_model_input, timestep
|
||||
|
||||
if offload:
|
||||
transformer.to(offload_device)
|
||||
vae.to(device)
|
||||
videos = vae.decode(x0.unsqueeze(0).to(vae.dtype), device=device)
|
||||
vae.to(offload_device)
|
||||
|
||||
# cache generated samples
|
||||
videos = torch.stack(videos).cpu() # B C T H W
|
||||
if colormatch != "disabled":
|
||||
videos = videos[0].permute(1, 2, 3, 0).cpu().numpy()
|
||||
from color_matcher import ColorMatcher
|
||||
cm = ColorMatcher()
|
||||
cm_result_list = []
|
||||
for img in videos:
|
||||
cm_result = cm.transfer(src=img, ref=original_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().numpy(), method=colormatch)
|
||||
cm_result_list.append(torch.from_numpy(cm_result))
|
||||
|
||||
videos = torch.stack(cm_result_list, dim=0).to(torch.float32).permute(3, 0, 1, 2).unsqueeze(0)
|
||||
|
||||
if is_first_clip:
|
||||
gen_video_list.append(videos)
|
||||
else:
|
||||
gen_video_list.append(videos[:, :, cur_motion_frames_num:])
|
||||
|
||||
# decide whether is done
|
||||
if arrive_last_frame: break
|
||||
|
||||
# update next condition frames
|
||||
is_first_clip = False
|
||||
cur_motion_frames_num = motion_frame
|
||||
|
||||
cond_image = videos[:, :, -cur_motion_frames_num:].to(torch.float32).to(device)
|
||||
audio_start_idx += (frame_num - cur_motion_frames_num)
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
|
||||
# Repeat audio emb
|
||||
if audio_end_idx >= min(max_frames_num, len(audio_embedding[0])):
|
||||
arrive_last_frame = True
|
||||
miss_lengths = []
|
||||
source_frames = []
|
||||
for human_inx in range(1):
|
||||
source_frame = len(audio_embedding[human_inx])
|
||||
source_frames.append(source_frame)
|
||||
if audio_end_idx >= len(audio_embedding[human_inx]):
|
||||
miss_length = audio_end_idx - len(audio_embedding[human_inx]) + 3
|
||||
add_audio_emb = torch.flip(audio_embedding[human_inx][-1*miss_length:], dims=[0])
|
||||
audio_embedding[human_inx] = torch.cat([audio_embedding[human_inx], add_audio_emb], dim=0)
|
||||
miss_lengths.append(miss_length)
|
||||
else:
|
||||
miss_lengths.append(0)
|
||||
|
||||
if max_frames_num <= frame_num: break
|
||||
|
||||
|
||||
gen_video_samples = torch.cat(gen_video_list, dim=2)[:, :, :int(max_frames_num)]
|
||||
gen_video_samples = gen_video_samples.to(torch.float32)
|
||||
if max_frames_num > frame_num and sum(miss_lengths) > 0:
|
||||
# split video frames
|
||||
gen_video_samples = gen_video_samples[:, :, :-1*miss_lengths[0]]
|
||||
|
||||
del noise, latent
|
||||
return {"video": gen_video_samples[0].permute(1, 2, 3, 0).cpu()},
|
||||
|
||||
#region normal inference
|
||||
else:
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
@@ -3959,6 +4209,11 @@ class WanVideoDecode:
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
mm.soft_empty_cache()
|
||||
video = samples.get("video", None)
|
||||
if video is not None:
|
||||
video = torch.clamp(video, -1.0, 1.0)
|
||||
video = (video + 1.0) / 2.0
|
||||
return video.cpu(),
|
||||
latents = samples["samples"]
|
||||
end_image = samples.get("end_image", None)
|
||||
has_ref = samples.get("has_ref", False)
|
||||
|
||||
Reference in New Issue
Block a user