Support SkyReels TalkingAvatar (A2V)
This commit is contained in:
File diff suppressed because it is too large
Load Diff
@@ -148,13 +148,13 @@ class AudioProjModel(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
seq_len=5,
|
||||
seq_len_vf=12,
|
||||
blocks=12,
|
||||
channels=768,
|
||||
seq_len_vf=8,
|
||||
blocks=12,
|
||||
channels=768,
|
||||
intermediate_dim=512,
|
||||
output_dim=768,
|
||||
context_tokens=32,
|
||||
norm_output_audio=False,
|
||||
norm_output_audio=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -278,9 +278,9 @@ class SingleStreamMultiAttention(SingleStreamAttention):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
encoder_hidden_states_dim: int,
|
||||
num_heads: int,
|
||||
qkv_bias: bool,
|
||||
qkv_bias: bool = True,
|
||||
encoder_hidden_states_dim: int = 768,
|
||||
class_range: int = 24,
|
||||
class_interval: int = 4,
|
||||
attention_mode: str = 'sdpa',
|
||||
|
||||
+114
-30
@@ -6,7 +6,7 @@ import numpy as np
|
||||
from ..latent_preview import prepare_callback
|
||||
from ..wanvideo.schedulers import get_scheduler
|
||||
from .multitalk import timestep_transform, add_noise
|
||||
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap
|
||||
from ..utils import log, print_memory, temporal_score_rescaling, offload_transformer, init_blockswap, match_and_blend_colors
|
||||
from comfy.utils import load_torch_file
|
||||
from ..nodes_model_loading import load_weights
|
||||
from ..HuMo.nodes import get_audio_emb_window
|
||||
@@ -48,7 +48,13 @@ def multitalk_loop(self, **kwargs):
|
||||
mode = image_embeds.get("multitalk_mode", "multitalk")
|
||||
if mode == "auto":
|
||||
mode = transformer.multitalk_model_type.lower()
|
||||
elif mode == "skyreelsv3":
|
||||
num_pseudo_frames = 5
|
||||
pseudo_frames = reference_keyframes = None
|
||||
keyframe_index = 0
|
||||
reference_video = image_embeds.get("reference_video", None)
|
||||
log.info(f"Multitalk mode: {mode}")
|
||||
drop_frames = image_embeds.get("drop_frames", 0)
|
||||
cond_frame = None
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
offloaded = False
|
||||
@@ -62,7 +68,9 @@ def multitalk_loop(self, **kwargs):
|
||||
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)
|
||||
original_images = image_embeds.get("multitalk_start_image", None)
|
||||
cond_image = original_images.clone() if original_images is not None else None
|
||||
original_color_reference = cond_image.clone() if cond_image is not None else None
|
||||
if original_images is None:
|
||||
original_images = torch.zeros([noise.shape[0], 1, target_h, target_w], device=device)
|
||||
|
||||
@@ -94,7 +102,6 @@ def multitalk_loop(self, **kwargs):
|
||||
audio_embedding = multitalk_audio_embeds
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
cond_frame = None
|
||||
|
||||
uni3c_data = None
|
||||
if uni3c_embeds is not None:
|
||||
@@ -110,9 +117,56 @@ def multitalk_loop(self, **kwargs):
|
||||
log.warning("No encoded silence file found, padding with end of audio embedding instead.")
|
||||
|
||||
total_frames = len(audio_embedding[0])
|
||||
estimated_iterations = total_frames // (frame_num - motion_frame) + 1
|
||||
estimated_iterations = total_frames // (frame_num - motion_frame - drop_frames) + 1
|
||||
callback = prepare_callback(patcher, estimated_iterations)
|
||||
|
||||
# If reference_video is provided, extract keyframes from it
|
||||
if mode == "skyreelsv3" and reference_video is not None:
|
||||
ref_video_length = reference_video.shape[1] # (C, T, H, W)
|
||||
if colormatch == "reinhard_torch":
|
||||
reference_video = match_and_blend_colors(reference_video, original_color_reference, 1.0)
|
||||
|
||||
if ref_video_length >= total_frames:
|
||||
# Reference is long enough - extract keyframes at the expected positions
|
||||
segment_interval = frame_num - motion_frame - drop_frames
|
||||
generate_idx = []
|
||||
current_idx = frame_num - 1
|
||||
while current_idx < total_frames:
|
||||
generate_idx.append(min(current_idx, ref_video_length - 1))
|
||||
current_idx += segment_interval
|
||||
else:
|
||||
# Calculate target indices then map to reference video
|
||||
audio_length = total_frames
|
||||
generate_idx_target = [0]
|
||||
segment_interval = frame_num - motion_frame - drop_frames
|
||||
current_idx = frame_num - 1
|
||||
while current_idx < audio_length - 1:
|
||||
generate_idx_target.append(current_idx)
|
||||
current_idx += segment_interval
|
||||
if generate_idx_target[-1] != audio_length - 1:
|
||||
generate_idx_target.append(audio_length - 1)
|
||||
|
||||
# Map target indices to reference video
|
||||
generate_idx_target = np.array(generate_idx_target, dtype=np.int16)
|
||||
original_max = generate_idx_target[-1]
|
||||
original_min = generate_idx_target[0]
|
||||
if original_max > original_min:
|
||||
generate_idx_float = (generate_idx_target.astype(np.float64) - original_min) * (ref_video_length - 1) / (original_max - original_min)
|
||||
generate_idx = np.clip(np.round(generate_idx_float), 0, ref_video_length - 1).astype(np.int32).tolist()
|
||||
else:
|
||||
generate_idx = [0]
|
||||
|
||||
generate_idx = generate_idx[1:]
|
||||
log.info(f"Reference video ({ref_video_length} frames) mapped to target ({total_frames} frames). Keyframe indices: {generate_idx}")
|
||||
|
||||
# Extract keyframes from reference video
|
||||
# reference_video shape: (C, T, H, W) from nodes.py processing
|
||||
# Select keyframes and add batch dimension: (C, num_keyframes, H, W) -> (1, C, num_keyframes, H, W)
|
||||
selected_keyframes = reference_video[:, generate_idx] # (C, num_keyframes, H, W)
|
||||
reference_keyframes = selected_keyframes.unsqueeze(0).cpu() # (1, C, num_keyframes, H, W)
|
||||
log.info(f"Extracted {len(generate_idx)} keyframes from provided reference video at indices {generate_idx}, shape: {reference_keyframes.shape}")
|
||||
log.info(f"Reference video total frames: {reference_video.shape[1]}, will generate {total_frames} total frames with {estimated_iterations} windows")
|
||||
|
||||
if frame_num >= total_frames:
|
||||
arrive_last_frame = True
|
||||
estimated_iterations = 1
|
||||
@@ -122,6 +176,14 @@ def multitalk_loop(self, **kwargs):
|
||||
while True: # start video generation iteratively
|
||||
self.cache_state = [None, None]
|
||||
|
||||
if mode == "skyreelsv3" and reference_keyframes is not None:
|
||||
clamped_index = min(keyframe_index, reference_keyframes.shape[2] - 1) # Clamp keyframe_index to reuse last keyframe if we run out
|
||||
pseudo_frames = reference_keyframes[:, :, clamped_index:clamped_index+1].repeat(1, 1, num_pseudo_frames, 1, 1) # Use one keyframe and repeat it 5 times
|
||||
log.info(f"Window {iteration_count}: using keyframe {clamped_index}/{reference_keyframes.shape[2]-1} for pseudo frames.")
|
||||
keyframe_index += 1
|
||||
else:
|
||||
pseudo_frames = None
|
||||
|
||||
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] if cond_image is not None else None
|
||||
@@ -133,15 +195,13 @@ def multitalk_loop(self, **kwargs):
|
||||
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)
|
||||
audio_embs = torch.cat(audio_embs, dim=0).to(dtype)
|
||||
|
||||
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]
|
||||
latent_frame_num = (frame_num - 1) // 4 + 1
|
||||
|
||||
noise = torch.randn(
|
||||
16, latent_frame_num,
|
||||
lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
|
||||
noise = torch.randn(16, latent_frame_num, lat_h, lat_w, dtype=torch.float32, device=torch.device("cpu"), generator=seed_g).to(device)
|
||||
|
||||
# Calculate the correct latent slice based on current iteration
|
||||
if is_first_clip:
|
||||
@@ -198,24 +258,41 @@ def multitalk_loop(self, **kwargs):
|
||||
if cond_image is not None or cond_frame is not None:
|
||||
cond_ = cond_image if (is_first_clip or humo_image_cond is None) else cond_frame
|
||||
cond_frame_num = cond_.shape[2]
|
||||
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
|
||||
padding_frames_pixels_values = torch.concat([cond_.to(device, vae.dtype), video_frames], dim=2)
|
||||
|
||||
# Prepare pseudo frames if enabled and available from reference_video
|
||||
if mode == "skyreelsv3" and pseudo_frames is not None:
|
||||
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num-num_pseudo_frames, target_h, target_w, device=device, dtype=vae.dtype)
|
||||
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames, pseudo_frames.to(device, vae.dtype)], dim=2)
|
||||
else:
|
||||
video_frames = torch.zeros(1, 3, frame_num-cond_frame_num, target_h, target_w, device=device, dtype=vae.dtype)
|
||||
padding_frames_pixels_values = torch.cat([cond_.to(device, vae.dtype), video_frames], dim=2)
|
||||
|
||||
# encode
|
||||
vae.to(device)
|
||||
y = vae.encode(padding_frames_pixels_values, device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
|
||||
|
||||
if mode == "multitalk":
|
||||
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
|
||||
else:
|
||||
if mode == "infinitetalk":
|
||||
cond_ = cond_image if is_first_clip else cond_frame
|
||||
latent_motion_frames = vae.encode(cond_.to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False).to(dtype)[0]
|
||||
else:
|
||||
latent_motion_frames = y[:, :cur_motion_frames_latent_num] # C T H W
|
||||
|
||||
vae.to(offload_device)
|
||||
|
||||
#motion_frame_index = cur_motion_frames_latent_num if mode == "infinitetalk" else 1
|
||||
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
|
||||
msk[:, :1] = 1
|
||||
if mode == "skyreelsv3" and pseudo_frames is not None:
|
||||
# create mask in pixel space, then transform
|
||||
msk_pixel = torch.ones(1, frame_num, lat_h, lat_w, device=device)
|
||||
msk_pixel[:, cur_motion_frames_num : -num_pseudo_frames] = 0
|
||||
msk_pixel = torch.cat([
|
||||
torch.repeat_interleave(msk_pixel[:, 0:1], repeats=4, dim=1),
|
||||
msk_pixel[:, 1:],
|
||||
], dim=1)
|
||||
msk_pixel = msk_pixel.view(1, msk_pixel.shape[1] // 4, 4, lat_h, lat_w)
|
||||
msk = msk_pixel.transpose(1, 2).squeeze(0).to(dtype) # 4 T H W
|
||||
else:
|
||||
msk = torch.zeros(4, latent_frame_num, lat_h, lat_w, device=device, dtype=dtype)
|
||||
msk[:, :1] = 1
|
||||
y = torch.cat([msk, y]) # 4+C T H W
|
||||
mm.soft_empty_cache()
|
||||
else:
|
||||
@@ -258,7 +335,7 @@ def multitalk_loop(self, **kwargs):
|
||||
latent = noise
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip and mode == "multitalk":
|
||||
if not is_first_clip and mode != "infinitetalk":
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[0])
|
||||
@@ -370,12 +447,12 @@ def multitalk_loop(self, **kwargs):
|
||||
latent = image_latent * mask + latent * (1-mask)
|
||||
|
||||
# injecting motion frames
|
||||
if not is_first_clip and mode == "multitalk":
|
||||
if not is_first_clip and mode != "infinitetalk":
|
||||
latent_motion_frames = latent_motion_frames.to(latent.dtype).to(device)
|
||||
motion_add_noise = torch.randn(latent_motion_frames.shape, device=torch.device("cpu"), generator=seed_g).to(device).contiguous()
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
|
||||
latent[:, :add_latent.shape[1]] = add_latent
|
||||
else:
|
||||
elif mode == "infinitetalk":
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
@@ -392,20 +469,27 @@ def multitalk_loop(self, **kwargs):
|
||||
|
||||
sampling_pbar.close()
|
||||
|
||||
# crop drop_frames from end if enabled
|
||||
if mode == "skyreelsv3" and drop_frames > 0 and not arrive_last_frame:
|
||||
videos = videos[:, :-drop_frames]
|
||||
|
||||
# optional color correction (less relevant for InfiniteTalk)
|
||||
if colormatch != "disabled":
|
||||
videos = videos.permute(1, 2, 3, 0).float().numpy()
|
||||
from color_matcher import ColorMatcher
|
||||
cm = ColorMatcher()
|
||||
cm_result_list = []
|
||||
for img in videos:
|
||||
if mode == "multitalk":
|
||||
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
|
||||
else:
|
||||
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
|
||||
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
|
||||
if colormatch == "reinhard_torch":
|
||||
videos = match_and_blend_colors(videos, original_color_reference, 1.0)
|
||||
else:
|
||||
videos = videos.permute(1, 2, 3, 0).float().numpy()
|
||||
from color_matcher import ColorMatcher
|
||||
cm = ColorMatcher()
|
||||
cm_result_list = []
|
||||
for img in videos:
|
||||
if mode == "infinitetalk":
|
||||
cm_result = cm.transfer(src=img, ref=cond_image[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
|
||||
else:
|
||||
cm_result = cm.transfer(src=img, ref=original_images[0].permute(1, 2, 3, 0).squeeze(0).cpu().float().numpy(), method=colormatch)
|
||||
cm_result_list.append(torch.from_numpy(cm_result).to(vae.dtype))
|
||||
|
||||
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
|
||||
videos = torch.stack(cm_result_list, dim=0).permute(3, 0, 1, 2)
|
||||
|
||||
# optionally save generated samples to disk
|
||||
if output_path:
|
||||
@@ -441,7 +525,7 @@ def multitalk_loop(self, **kwargs):
|
||||
|
||||
# Repeat audio emb
|
||||
if multitalk_embeds is not None:
|
||||
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
|
||||
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count - drop_frames)
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
if audio_end_idx >= len(audio_embedding[0]):
|
||||
arrive_last_frame = True
|
||||
|
||||
+89
-1
@@ -461,13 +461,100 @@ class WanVideoImageToVideoMultiTalk:
|
||||
}
|
||||
|
||||
return (image_embeds, output_path)
|
||||
|
||||
|
||||
class WanVideoImageToVideoSkyreelsv3_audio:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"vae": ("WANVAE",),
|
||||
"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": 5, "min": 1, "max": 10000, "step": 1, "tooltip": "Driven frame length used in the long video generation. Basically the overlap length."}),
|
||||
"drop_frames": ("INT", {"default": 12, "min": 0, "max": 10000, "step": 1, "tooltip": "Additional frames to drop when advancing the audio window. Higher values = less overlap = faster generation but potentially less smooth transitions."}),
|
||||
"tiled_vae": ("BOOLEAN", {"default": False, "tooltip": "Use tiled VAE encoding for reduced memory use"}),
|
||||
"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',
|
||||
'reinhard_torch',
|
||||
'mkl',
|
||||
'hm',
|
||||
'reinhard',
|
||||
'mvgd',
|
||||
'hm-mvgd-hm',
|
||||
'hm-mkl-hm',
|
||||
], {
|
||||
"default": 'disabled', "tooltip": "Color matching method to use between the windows"
|
||||
},),
|
||||
},
|
||||
"optional": {
|
||||
"start_image": ("IMAGE", {"tooltip": "Images to encode"}),
|
||||
"reference_video": ("IMAGE", {"tooltip": "Optional: Pre-generated reference video to use for keyframes instead of extracting from first generation. Should be color-matched to source image."}),
|
||||
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
|
||||
"output_path": ("STRING", {"default": "", "tooltip": "If set, will save each window's resulting frames to this folder, also DISABLES returning the final video tensor to save memory"}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "STRING",)
|
||||
RETURN_NAMES = ("image_embeds", "output_path")
|
||||
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, drop_frames, force_offload, colormatch, start_image=None,
|
||||
tiled_vae=False, clip_embeds=None, mode="multitalk", output_path="", reference_video=None):
|
||||
|
||||
H, W = height, width
|
||||
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) // 4 + 1, height // 8, width // 8)
|
||||
|
||||
if output_path:
|
||||
timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
output_path = os.path.join(output_path, f"{timestamp}_{mode}_output")
|
||||
os.makedirs(output_path, exist_ok=True)
|
||||
|
||||
processed_reference_video = None
|
||||
if reference_video is not None:
|
||||
processed_reference_video = common_upscale(reference_video.movedim(-1, 1), W, H, "lanczos", "disabled").movedim(0, 1)
|
||||
processed_reference_video = processed_reference_video * 2 - 1
|
||||
|
||||
image_embeds = {
|
||||
"multitalk_sampling": True,
|
||||
"multitalk_start_image": resized_start_image if start_image is not None else None,
|
||||
"frame_window_size": num_frames,
|
||||
"motion_frame": motion_frame,
|
||||
"drop_frames": drop_frames,
|
||||
"use_pseudo_frames": True,
|
||||
"reference_video": processed_reference_video,
|
||||
"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,
|
||||
"multitalk_mode": "skyreelsv3",
|
||||
"output_path": output_path
|
||||
}
|
||||
|
||||
return (image_embeds, output_path)
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"MultiTalkModelLoader": MultiTalkModelLoader,
|
||||
"MultiTalkWav2VecEmbeds": MultiTalkWav2VecEmbeds,
|
||||
"WanVideoImageToVideoMultiTalk": WanVideoImageToVideoMultiTalk,
|
||||
"Wav2VecModelLoader": Wav2VecModelLoader,
|
||||
"MultiTalkSilentEmbeds": MultiTalkSilentEmbeds,
|
||||
"WanVideoImageToVideoSkyreelsv3_audio": WanVideoImageToVideoSkyreelsv3_audio,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
@@ -476,4 +563,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoImageToVideoMultiTalk": "WanVideo Long I2V Multi/InfiniteTalk",
|
||||
"Wav2VecModelLoader": "Wav2vec2 Model Loader",
|
||||
"MultiTalkSilentEmbeds": "MultiTalk Silent Embeds",
|
||||
"WanVideoImageToVideoSkyreelsv3_audio": "WanVideo Long SkyReelsV3 A2V",
|
||||
}
|
||||
+43
-47
@@ -810,7 +810,6 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
"adapter", "add", "ref_conv", "casual_audio_encoder", "cond_encoder", "frame_packer", "audio_proj_glob", "face_encoder", "fuser_block"}
|
||||
param_count = sum(1 for _ in transformer.named_parameters())
|
||||
pbar = ProgressBar(param_count)
|
||||
cnt = 0
|
||||
block_idx = vace_block_idx = None
|
||||
|
||||
if gguf:
|
||||
@@ -920,9 +919,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
load_device = offload_device
|
||||
# Set tensor to device
|
||||
set_module_tensor_to_device(transformer, name, device=load_device, dtype=dtype_to_use, value=value)
|
||||
cnt += 1
|
||||
if cnt % 100 == 0:
|
||||
pbar.update(100)
|
||||
pbar.update(1)
|
||||
|
||||
#[print(name, param.device, param.dtype) for name, param in transformer.named_parameters()]
|
||||
memory_on_device = get_module_memory_mb_per_device(transformer)
|
||||
@@ -931,6 +928,8 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
for dev, mem_mb in memory_on_device.items():
|
||||
log.info(f"Device: {dev:8s} | Memory: {mem_mb:,.2f} MB")
|
||||
|
||||
if hasattr(pbar, "_last_sent_value"):
|
||||
pbar._last_sent_value = -1
|
||||
pbar.update_absolute(0)
|
||||
|
||||
def patch_control_lora(transformer, device):
|
||||
@@ -1512,7 +1511,45 @@ class WanVideoModelLoader:
|
||||
block.cross_attn.ip_adapter_single_stream_k_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
block.cross_attn.ip_adapter_single_stream_v_proj = nn.Linear(context_dim, dim, bias=False)
|
||||
|
||||
if multitalk_model is not None:
|
||||
# LongCat Avatar
|
||||
if "multitalk_audio_proj.proj1.weight" in sd and "blocks.0.audio_cross_attn.q_norm.weight" in sd:
|
||||
log.info("MultiTalk/InfiniteTalk model detected, patching model...")
|
||||
from .multitalk.multitalk import AudioProjModel
|
||||
from .wanvideo.modules.model import WanLayerNorm
|
||||
from .LongCat.layers import SingleStreamAttention
|
||||
|
||||
|
||||
for block in transformer.blocks:
|
||||
with init_empty_weights():
|
||||
if "blocks.0.audio_modulation.1.weight" in sd:
|
||||
block.audio_modulation = nn.Sequential(nn.SiLU(), nn.Linear(512, 3 * dim, bias=True))
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
|
||||
block.audio_cross_attn = SingleStreamAttention(
|
||||
dim=dim,
|
||||
encoder_hidden_states_dim=768,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
qk_norm=True,
|
||||
class_range=24,
|
||||
class_interval=4,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
multitalk_proj_model = AudioProjModel()
|
||||
transformer.multitalk_audio_proj = multitalk_proj_model
|
||||
# SkyreelsV3
|
||||
elif "blocks.1.audio_cross_attn.kv_linear.weight" in sd and "audio_proj.proj1.weight" in sd:
|
||||
sd = {k.replace("audio_proj", "multitalk_audio_proj"): v for k, v in sd.items()}
|
||||
# init audio module
|
||||
from .multitalk.multitalk import SingleStreamMultiAttention, AudioProjModel
|
||||
from .wanvideo.modules.model import WanLayerNorm
|
||||
|
||||
for block in transformer.blocks:
|
||||
with init_empty_weights():
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
|
||||
block.audio_cross_attn = SingleStreamMultiAttention(dim=dim, num_heads=num_heads, attention_mode=attention_mode)
|
||||
|
||||
transformer.multitalk_audio_proj = AudioProjModel()
|
||||
elif multitalk_model is not None:
|
||||
multitalk_model_type = multitalk_model.get("model_type", "MultiTalk")
|
||||
log.info(f"{multitalk_model_type} detected, patching model...")
|
||||
|
||||
@@ -1529,15 +1566,7 @@ class WanVideoModelLoader:
|
||||
for block in transformer.blocks:
|
||||
with init_empty_weights():
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
|
||||
block.audio_cross_attn = SingleStreamMultiAttention(
|
||||
dim=dim,
|
||||
encoder_hidden_states_dim=768,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
class_range=24,
|
||||
class_interval=4,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
block.audio_cross_attn = SingleStreamMultiAttention(dim=dim, num_heads=num_heads, attention_mode=attention_mode)
|
||||
transformer.multitalk_audio_proj = multitalk_model["proj_model"]
|
||||
transformer.multitalk_model_type = multitalk_model_type
|
||||
|
||||
@@ -1555,39 +1584,6 @@ class WanVideoModelLoader:
|
||||
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
elif "multitalk_audio_proj.proj1.weight" in sd:
|
||||
log.info("MultiTalk/InfiniteTalk model detected, patching model...")
|
||||
from .multitalk.multitalk import AudioProjModel
|
||||
from .wanvideo.modules.model import WanLayerNorm
|
||||
from .LongCat.layers import SingleStreamAttention
|
||||
|
||||
audio_window = 5
|
||||
vae_scale = 4
|
||||
|
||||
for block in transformer.blocks:
|
||||
with init_empty_weights():
|
||||
if "blocks.0.audio_modulation.1.weight" in sd:
|
||||
block.audio_modulation = nn.Sequential(nn.SiLU(), nn.Linear(512, 3 * dim, bias=True))
|
||||
block.norm_x = WanLayerNorm(dim, transformer.eps, elementwise_affine=True)
|
||||
block.audio_cross_attn = SingleStreamAttention(
|
||||
dim=dim,
|
||||
encoder_hidden_states_dim=768,
|
||||
num_heads=num_heads,
|
||||
qkv_bias=True,
|
||||
qk_norm=True,
|
||||
class_range=24,
|
||||
class_interval=4,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
multitalk_proj_model = AudioProjModel(
|
||||
seq_len=audio_window,
|
||||
seq_len_vf=audio_window+vae_scale-1,
|
||||
intermediate_dim=512,
|
||||
output_dim=768,
|
||||
context_tokens=32,
|
||||
norm_output_audio=True,
|
||||
)
|
||||
transformer.multitalk_audio_proj = multitalk_proj_model
|
||||
|
||||
sd = {k.replace(".weight_scale", ".scale_weight"): v for k, v in sd.items()}
|
||||
|
||||
|
||||
+6
-5
@@ -185,11 +185,12 @@ class WanVideoSampler:
|
||||
|
||||
is_pusa = "pusa" in sample_scheduler.__class__.__name__.lower()
|
||||
|
||||
scheduler_step_args = {"generator": seed_g}
|
||||
step_sig = inspect.signature(sample_scheduler.step)
|
||||
for arg in list(scheduler_step_args.keys()):
|
||||
if arg not in step_sig.parameters:
|
||||
scheduler_step_args.pop(arg)
|
||||
if scheduler != "multitalk":
|
||||
scheduler_step_args = {"generator": seed_g}
|
||||
step_sig = inspect.signature(sample_scheduler.step)
|
||||
for arg in list(scheduler_step_args.keys()):
|
||||
if arg not in step_sig.parameters:
|
||||
scheduler_step_args.pop(arg)
|
||||
|
||||
# Ovi
|
||||
if transformer.audio_model is not None: # temporary workaround (...nothing more permanent)
|
||||
|
||||
@@ -718,3 +718,55 @@ def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.
|
||||
if not t == 1.0:
|
||||
model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t)
|
||||
return model_output
|
||||
|
||||
def match_and_blend_colors(
|
||||
source_chunk: torch.Tensor, # (C, T, H, W), range [-1, 1]
|
||||
reference_image: torch.Tensor, # (C, 1, H, W), range [-1, 1]
|
||||
strength: float,
|
||||
) -> torch.Tensor:
|
||||
import kornia
|
||||
if strength == 0.0:
|
||||
return source_chunk
|
||||
source_chunk = source_chunk.unsqueeze(0) # (1, C, T, H, W)
|
||||
|
||||
# shapes
|
||||
B, C, T, H, W = source_chunk.shape
|
||||
input_dtype = source_chunk.dtype
|
||||
|
||||
# [-1,1] -> [0,1]
|
||||
src_01 = (source_chunk + 1.0) * 0.5
|
||||
ref_01 = (reference_image + 1.0) * 0.5
|
||||
|
||||
src32 = src_01.to(torch.float32)
|
||||
ref32 = ref_01.to(torch.float32)
|
||||
|
||||
# (B, C, T, H, W) -> (B*T, C, H, W)
|
||||
src_bt = src32.permute(0, 2, 1, 3, 4).contiguous().view(B * T, C, H, W)
|
||||
ref_bchw = ref32[:, :, 0, :, :].contiguous()
|
||||
|
||||
# RGB->Lab
|
||||
src_lab = kornia.color.rgb_to_lab(src_bt) # (B*T, C, H, W)
|
||||
ref_lab = kornia.color.rgb_to_lab(ref_bchw) # (B, C, H, W)
|
||||
|
||||
src_lab_flat = src_lab.view(B * T, C, -1) # (B*T, C, HW)
|
||||
ref_lab_flat = ref_lab.view(B, C, -1) # (B, C, HW)
|
||||
src_std, src_mean = torch.std_mean(src_lab_flat, dim=-1, keepdim=True, unbiased=False)
|
||||
ref_std, ref_mean = torch.std_mean(ref_lab_flat, dim=-1, keepdim=True, unbiased=False)
|
||||
src_std = src_std.clamp_min_(1e-6)
|
||||
|
||||
ref_mean_bt = ref_mean.repeat_interleave(T, dim=0) # (B*T, C, 1)
|
||||
ref_std_bt = ref_std.repeat_interleave(T, dim=0) # (B*T, C, 1)
|
||||
|
||||
corrected_lab_flat = (src_lab_flat - src_mean) * (ref_std_bt / src_std) + ref_mean_bt
|
||||
corrected_lab = corrected_lab_flat.view(B * T, C, H, W)
|
||||
|
||||
# Lab->RGB
|
||||
corrected_rgb_01 = kornia.color.lab_to_rgb(corrected_lab) # (B*T, C, H, W)
|
||||
|
||||
blended_rgb_01 = (1.0 - strength) * src_bt + strength * corrected_rgb_01
|
||||
|
||||
# (B, C, T, H, W)
|
||||
blended_rgb_01 = blended_rgb_01.view(B, T, C, H, W).permute(0, 2, 1, 3, 4).contiguous()
|
||||
|
||||
# [0,1] -> [-1,1]
|
||||
return (blended_rgb_01 * 2.0 - 1.0)[0].to(dtype=input_dtype)
|
||||
|
||||
Reference in New Issue
Block a user