Experimental: Allow HuMo to work with InfiniteTalk
Doesn't work that great, but pushing this anyway for possible future usecases
This commit is contained in:
+25
-3
@@ -204,7 +204,7 @@ class HuMoEmbeds:
|
||||
|
||||
log.info(f"HuMo set to generate {pixel_frame_num} frames")
|
||||
|
||||
audio_emb, _ = get_audio_emb_window(audio_emb, pixel_frame_num, frame0_idx=0)
|
||||
#audio_emb, _ = get_audio_emb_window(audio_emb, pixel_frame_num, frame0_idx=0)
|
||||
|
||||
num_refs = 0
|
||||
if reference_images is not None:
|
||||
@@ -229,8 +229,8 @@ class HuMoEmbeds:
|
||||
if reference_images is not None:
|
||||
mask[:,:-num_refs] = 0
|
||||
image_cond = torch.cat([zero_latents[:, :(target_shape[1]-num_refs)], samples], dim=1)
|
||||
zero_audio_pad = torch.zeros(num_refs, *audio_emb.shape[1:]).to(audio_emb.device)
|
||||
audio_emb = torch.cat([audio_emb, zero_audio_pad], dim=0)
|
||||
#zero_audio_pad = torch.zeros(num_refs, *audio_emb.shape[1:]).to(audio_emb.device)
|
||||
#audio_emb = torch.cat([audio_emb, zero_audio_pad], dim=0)
|
||||
else:
|
||||
image_cond = zero_latents
|
||||
mask = torch.zeros_like(mask)
|
||||
@@ -252,14 +252,36 @@ class HuMoEmbeds:
|
||||
}
|
||||
|
||||
return (embeds, )
|
||||
|
||||
class WanVideoCombineEmbeds:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"embeds_1": ("WANVIDIMAGE_EMBEDS",),
|
||||
"embeds_2": ("WANVIDIMAGE_EMBEDS",),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||
RETURN_NAMES = ("image_embeds",)
|
||||
FUNCTION = "add"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
EXPERIMENTAL = True
|
||||
|
||||
def add(self, embeds_1, embeds_2):
|
||||
# Combine the two sets of embeds
|
||||
combined = {**embeds_1, **embeds_2}
|
||||
return (combined,)
|
||||
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"WhisperModelLoader": WhisperModelLoader,
|
||||
"HuMoEmbeds": HuMoEmbeds,
|
||||
"WanVideoCombineEmbeds": WanVideoCombineEmbeds,
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WhisperModelLoader": "Whisper Model Loader",
|
||||
"HuMoEmbeds": "HuMo Embeds",
|
||||
"WanVideoCombineEmbeds": "WanVideo Combine Embeds",
|
||||
}
|
||||
|
||||
+1
-1
@@ -428,7 +428,7 @@ class WanVideoImageToVideoMultiTalk:
|
||||
image_embeds = {
|
||||
"multitalk_sampling": True,
|
||||
"multitalk_start_image": resized_start_image if start_image is not None else None,
|
||||
"num_frames": num_frames,
|
||||
"frame_window_size": num_frames,
|
||||
"motion_frame": motion_frame,
|
||||
"target_h": H,
|
||||
"target_w": W,
|
||||
|
||||
@@ -2107,15 +2107,24 @@ class WanVideoSampler:
|
||||
|
||||
#HuMo inputs
|
||||
humo_audio = image_embeds.get("humo_audio_emb", None)
|
||||
if humo_audio is not None:
|
||||
humo_audio = humo_audio.to(device, dtype)
|
||||
humo_audio_neg = image_embeds.get("humo_audio_emb_neg", None)
|
||||
humo_reference_count = image_embeds.get("humo_reference_count", 0)
|
||||
num_frames = image_embeds.get("num_frames", 0)
|
||||
if humo_audio is not None:
|
||||
from .HuMo.nodes import get_audio_emb_window
|
||||
if not multitalk_sampling:
|
||||
humo_audio, _ = get_audio_emb_window(humo_audio, num_frames, frame0_idx=0)
|
||||
zero_audio_pad = torch.zeros(humo_reference_count, *humo_audio.shape[1:]).to(humo_audio.device)
|
||||
humo_audio = torch.cat([humo_audio, zero_audio_pad], dim=0)
|
||||
humo_audio_neg = torch.zeros_like(humo_audio, dtype=humo_audio.dtype, device=humo_audio.device)
|
||||
humo_audio = humo_audio.to(device, dtype)
|
||||
|
||||
if humo_audio_neg is not None:
|
||||
humo_audio_neg = humo_audio_neg.to(device, dtype)
|
||||
humo_audio_scale = image_embeds.get("humo_audio_scale", 1.0)
|
||||
humo_image_cond = image_embeds.get("humo_image_cond", None)
|
||||
humo_image_cond_neg = image_embeds.get("humo_image_cond_neg", None)
|
||||
humo_reference_count = image_embeds.get("humo_reference_count", 0)
|
||||
|
||||
humo_audio_cfg_scale = image_embeds.get("humo_audio_cfg_scale", 1.0)
|
||||
humo_start_percent = image_embeds.get("humo_start_percent", 0.0)
|
||||
humo_end_percent = image_embeds.get("humo_end_percent", 1.0)
|
||||
@@ -2178,7 +2187,7 @@ class WanVideoSampler:
|
||||
}
|
||||
|
||||
# FantasyTalking
|
||||
audio_proj = multitalk_audio_embedding = None
|
||||
audio_proj = multitalk_audio_embeds = None
|
||||
audio_scale = 1.0
|
||||
if fantasytalking_embeds is not None:
|
||||
audio_proj = fantasytalking_embeds["audio_proj"].to(device)
|
||||
@@ -2191,13 +2200,13 @@ class WanVideoSampler:
|
||||
# Handle single or multiple speaker embeddings
|
||||
audio_features_in = multitalk_embeds.get("audio_features", None)
|
||||
if audio_features_in is None:
|
||||
multitalk_audio_embedding = None
|
||||
multitalk_audio_embeds = None
|
||||
else:
|
||||
if isinstance(audio_features_in, list):
|
||||
multitalk_audio_embedding = [emb.to(device, dtype) for emb in audio_features_in]
|
||||
multitalk_audio_embeds = [emb.to(device, dtype) for emb in audio_features_in]
|
||||
else:
|
||||
# keep backward-compatibility with single tensor input
|
||||
multitalk_audio_embedding = [audio_features_in.to(device, dtype)]
|
||||
multitalk_audio_embeds = [audio_features_in.to(device, dtype)]
|
||||
|
||||
audio_scale = multitalk_embeds.get("audio_scale", 1.0)
|
||||
audio_cfg_scale = multitalk_embeds.get("audio_cfg_scale", 1.0)
|
||||
@@ -2205,7 +2214,7 @@ class WanVideoSampler:
|
||||
if not isinstance(audio_cfg_scale, list):
|
||||
audio_cfg_scale = [audio_cfg_scale] * (steps + 1)
|
||||
|
||||
shapes = [tuple(e.shape) for e in multitalk_audio_embedding]
|
||||
shapes = [tuple(e.shape) for e in multitalk_audio_embeds]
|
||||
log.info(f"Multitalk audio features shapes (per speaker): {shapes}")
|
||||
|
||||
# FantasyPortrait
|
||||
@@ -2628,7 +2637,8 @@ class WanVideoSampler:
|
||||
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, multitalk_audio_embeds=None, fantasy_portrait_input=None, reverse_time=False,
|
||||
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None):
|
||||
mtv_motion_tokens=None, s2v_audio_input=None, s2v_ref_motion=None, s2v_motion_frames=[1, 0], s2v_pose=None,
|
||||
humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None):
|
||||
nonlocal transformer
|
||||
z = z.to(dtype)
|
||||
autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear)
|
||||
@@ -2746,8 +2756,8 @@ class WanVideoSampler:
|
||||
else:
|
||||
z = torch.cat([z, minimax_latents, minimax_mask_latents], dim=0)
|
||||
|
||||
if not multitalk_sampling and multitalk_audio_embedding is not None:
|
||||
audio_embedding = multitalk_audio_embedding
|
||||
if not multitalk_sampling and multitalk_audio_embeds is not None:
|
||||
audio_embedding = multitalk_audio_embeds
|
||||
audio_embs = []
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
human_num = len(audio_embedding)
|
||||
@@ -2828,8 +2838,8 @@ class WanVideoSampler:
|
||||
"add_cond": add_cond_input, # additional conditioning input
|
||||
"nag_params": text_embeds.get("nag_params", {}), # normalized attention guidance
|
||||
"nag_context": text_embeds.get("nag_prompt_embeds", None), # normalized attention guidance context
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk audio input
|
||||
"ref_target_masks": ref_target_masks if multitalk_audio_embedding is not None else None, # Multi/InfiniteTalk reference target masks
|
||||
"multitalk_audio": multitalk_audio_input if multitalk_audio_embeds is not None else None, # Multi/InfiniteTalk audio input
|
||||
"ref_target_masks": ref_target_masks if multitalk_audio_embeds is not None else None, # Multi/InfiniteTalk reference target masks
|
||||
"inner_t": [shot_len] if shot_len else None, # inner timestep for EchoShot
|
||||
"standin_input": standin_input, # Stand-in reference input
|
||||
"fantasy_portrait_input": fantasy_portrait_input, # Fantasy portrait input
|
||||
@@ -2910,7 +2920,7 @@ class WanVideoSampler:
|
||||
+ cfg_scale * (noise_pred_cond - noise_pred_phantom[0]))
|
||||
return noise_pred, [cache_state_cond, cache_state_uncond, cache_state_phantom]
|
||||
#audio cfg (fantasytalking and multitalk)
|
||||
if (fantasytalking_embeds is not None or multitalk_audio_embedding is not None):
|
||||
if (fantasytalking_embeds is not None or multitalk_audio_embeds is not None):
|
||||
if not math.isclose(audio_cfg_scale[idx], 1.0):
|
||||
if cache_state is not None and len(cache_state) != 3:
|
||||
cache_state.append(None)
|
||||
@@ -3396,7 +3406,8 @@ class WanVideoSampler:
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
partial_timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
|
||||
partial_control_camera_latents, partial_add_cond, current_teacache, context_window=c, fantasy_portrait_input=partial_fantasy_portrait_input,
|
||||
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose)
|
||||
mtv_motion_tokens=partial_mtv_motion_tokens, s2v_audio_input=partial_s2v_audio_input, s2v_motion_frames=[1, 0], s2v_pose=partial_s2v_pose,
|
||||
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg,)
|
||||
|
||||
if cache_args is not None:
|
||||
self.window_tracker.cache_states[window_id] = new_teacache
|
||||
@@ -3416,7 +3427,7 @@ class WanVideoSampler:
|
||||
offload = image_embeds.get("force_offload", False)
|
||||
offloaded = False
|
||||
tiled_vae = image_embeds.get("tiled_vae", False)
|
||||
frame_num = clip_length = image_embeds.get("num_frames", 81)
|
||||
frame_num = clip_length = image_embeds.get("frame_window_size", 81)
|
||||
vae = image_embeds.get("vae", None)
|
||||
clip_embeds = image_embeds.get("clip_context", None)
|
||||
if clip_embeds is not None:
|
||||
@@ -3454,9 +3465,10 @@ class WanVideoSampler:
|
||||
indices = (torch.arange(4 + 1) - 2) * 1
|
||||
current_condframe_index = 0
|
||||
|
||||
audio_embedding = multitalk_audio_embedding
|
||||
audio_embedding = multitalk_audio_embeds
|
||||
human_num = len(audio_embedding)
|
||||
audio_embs = None
|
||||
cond_frame = None
|
||||
|
||||
pcd_data = pcd_data_input = None
|
||||
if uni3c_embeds is not None:
|
||||
@@ -3564,9 +3576,11 @@ class WanVideoSampler:
|
||||
masks = (1-noise_mask.repeat(len(timesteps), 1, 1, 1, 1).to(device)) > thresholds
|
||||
|
||||
# zero padding and vae encode for img cond
|
||||
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)
|
||||
if cond_image is not None or cond_frame is not None:
|
||||
cond_ = cond_image if is_first_clip 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)
|
||||
|
||||
# encode
|
||||
vae.to(device)
|
||||
@@ -3589,6 +3603,25 @@ class WanVideoSampler:
|
||||
y = None
|
||||
latent_motion_frames = noise[:, :1]
|
||||
|
||||
partial_humo_cond_input = partial_humo_cond_neg_input = partial_humo_audio = partial_humo_audio_neg = None
|
||||
if humo_image_cond is not None:
|
||||
partial_humo_cond_input = humo_image_cond[:, :latent_frame_num]
|
||||
partial_humo_cond_neg_input = humo_image_cond_neg[:, :latent_frame_num]
|
||||
if y is not None:
|
||||
partial_humo_cond_input[:, :1] = y[:, :1]
|
||||
if humo_reference_count > 0:
|
||||
partial_humo_cond_input[:, -humo_reference_count:] = humo_image_cond[:, -humo_reference_count:]
|
||||
partial_humo_cond_neg_input[:, -humo_reference_count:] = humo_image_cond_neg[:, -humo_reference_count:]
|
||||
|
||||
if humo_audio is not None:
|
||||
if is_first_clip:
|
||||
audio_embs = None
|
||||
|
||||
partial_humo_audio, _ = get_audio_emb_window(humo_audio, frame_num, frame0_idx=audio_start_idx)
|
||||
#zero_audio_pad = torch.zeros(humo_reference_count, *partial_humo_audio.shape[1:], device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
|
||||
partial_humo_audio[-humo_reference_count:] = 0
|
||||
partial_humo_audio_neg = torch.zeros_like(partial_humo_audio, device=partial_humo_audio.device, dtype=partial_humo_audio.dtype)
|
||||
|
||||
if scheduler == "multitalk":
|
||||
timesteps = list(np.linspace(1000, 1, steps, dtype=np.float32))
|
||||
timesteps.append(0.)
|
||||
@@ -3723,12 +3756,14 @@ class WanVideoSampler:
|
||||
timestep = timesteps[i]
|
||||
latent_model_input = latent.to(device)
|
||||
if mode == "infinitetalk":
|
||||
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent_model_input[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
noise_pred, self.cache_state = predict_with_cfg(
|
||||
latent_model_input, cfg[min(i, len(timesteps)-1)], positive, text_embeds["negative_prompt_embeds"],
|
||||
timestep, i, y, clip_embeds, control_latents, window_vace_data, partial_unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input)
|
||||
cache_state=self.cache_state, multitalk_audio_embeds=audio_embs, fantasy_portrait_input=partial_fantasy_portrait_input,
|
||||
humo_image_cond=partial_humo_cond_input, humo_image_cond_neg=partial_humo_cond_neg_input, humo_audio=partial_humo_audio, humo_audio_neg=partial_humo_audio_neg)
|
||||
|
||||
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)
|
||||
@@ -3761,12 +3796,15 @@ class WanVideoSampler:
|
||||
add_latent = add_noise(latent_motion_frames, motion_add_noise, timesteps[i+1])
|
||||
latent[:, :add_latent.shape[1]] = add_latent
|
||||
else:
|
||||
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
if humo_image_cond is None or not is_first_clip:
|
||||
latent[:, :cur_motion_frames_latent_num] = latent_motion_frames
|
||||
|
||||
del noise, latent_motion_frames
|
||||
if offload:
|
||||
offload_transformer(transformer)
|
||||
offloaded = True
|
||||
if humo_image_cond is not None and humo_reference_count > 0:
|
||||
latent = latent[:,:-humo_reference_count]
|
||||
vae.to(device)
|
||||
videos = vae.decode(latent.unsqueeze(0).to(device, vae.dtype), device=device, tiled=tiled_vae, pbar=False)[0].cpu()
|
||||
vae.model.clear_cache()
|
||||
@@ -3823,7 +3861,7 @@ class WanVideoSampler:
|
||||
|
||||
# Repeat audio emb
|
||||
if multitalk_embeds is not None:
|
||||
audio_start_idx += (frame_num - cur_motion_frames_num)
|
||||
audio_start_idx += (frame_num - cur_motion_frames_num - humo_reference_count)
|
||||
audio_end_idx = audio_start_idx + clip_length
|
||||
if audio_end_idx >= len(audio_embedding[0]):
|
||||
arrive_last_frame = True
|
||||
@@ -4009,7 +4047,8 @@ class WanVideoSampler:
|
||||
cfg[idx], text_embeds["prompt_embeds"],
|
||||
text_embeds["negative_prompt_embeds"],
|
||||
timestep, idx, image_cond, clip_fea, control_latents, vace_data, unianim_data, audio_proj, control_camera_latents, add_cond,
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input)
|
||||
cache_state=self.cache_state, fantasy_portrait_input=fantasy_portrait_input, multitalk_audio_embeds=multitalk_audio_embeds, mtv_motion_tokens=mtv_motion_tokens, s2v_audio_input=s2v_audio_input,
|
||||
humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg)
|
||||
if bidirectional_sampling:
|
||||
noise_pred_flipped, self.cache_state = predict_with_cfg(
|
||||
latent_model_input_flipped,
|
||||
|
||||
+16
-11
@@ -783,15 +783,18 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
for r in reader:
|
||||
all_tensors.extend(r.tensors)
|
||||
for tensor in all_tensors:
|
||||
name = tensor.name
|
||||
if "glob" not in name and "audio_proj" in name:
|
||||
name = name.replace("audio_proj", "multitalk_audio_proj")
|
||||
load_device = device
|
||||
if "vace_blocks." in tensor.name:
|
||||
if "vace_blocks." in name:
|
||||
try:
|
||||
vace_block_idx = int(tensor.name.split("vace_blocks.")[1].split(".")[0])
|
||||
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
vace_block_idx = None
|
||||
elif "blocks." in tensor.name:
|
||||
elif "blocks." in name:
|
||||
try:
|
||||
block_idx = int(tensor.name.split("blocks.")[1].split(".")[0])
|
||||
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
||||
except Exception:
|
||||
block_idx = None
|
||||
|
||||
@@ -805,7 +808,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
||||
|
||||
is_gguf_quant = tensor.tensor_type not in [GGMLQuantizationType.F32, GGMLQuantizationType.F16]
|
||||
weights = torch.from_numpy(tensor.data.copy()).to(load_device)
|
||||
sd[tensor.name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
|
||||
sd[name] = GGUFParameter(weights, quant_type=tensor.tensor_type) if is_gguf_quant else weights
|
||||
sd.update(extra_sd)
|
||||
del all_tensors, extra_sd
|
||||
|
||||
@@ -1298,16 +1301,21 @@ class WanVideoModelLoader:
|
||||
class_interval=4,
|
||||
attention_mode=attention_mode,
|
||||
)
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
transformer.multitalk_audio_proj = multitalk_model["proj_model"]
|
||||
transformer.multitalk_model_type = multitalk_model_type
|
||||
|
||||
extra_model_path = multitalk_model["model_path"]
|
||||
extra_sd = {}
|
||||
if multitalk_model_path.endswith(".gguf"):
|
||||
extra_sd, extra_reader = load_gguf(extra_model_path)
|
||||
extra_sd_temp, extra_reader = load_gguf(extra_model_path)
|
||||
gguf_reader.append(extra_reader)
|
||||
del extra_reader
|
||||
else:
|
||||
extra_sd = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
|
||||
extra_sd_temp = load_torch_file(extra_model_path, device=transformer_load_device, safe_load=True)
|
||||
|
||||
for k, v in extra_sd_temp.items():
|
||||
extra_sd[k.replace("audio_proj.", "multitalk_audio_proj.")] = v
|
||||
|
||||
sd.update(extra_sd)
|
||||
del extra_sd
|
||||
|
||||
@@ -1392,9 +1400,6 @@ class WanVideoModelLoader:
|
||||
from .fp8_optimization import convert_fp8_linear
|
||||
convert_fp8_linear(transformer, base_dtype, params_to_keep, scale_weight_keys=scale_weights)
|
||||
|
||||
if multitalk_model is not None:
|
||||
transformer.audio_proj = multitalk_model["proj_model"]
|
||||
|
||||
if vram_management_args is not None:
|
||||
if gguf:
|
||||
raise ValueError("GGUF models don't support vram management")
|
||||
|
||||
@@ -2279,12 +2279,12 @@ class WanModel(torch.nn.Module):
|
||||
|
||||
# MultiTalk
|
||||
if multitalk_audio is not None:
|
||||
self.audio_proj.to(self.main_device)
|
||||
self.multitalk_audio_proj.to(self.main_device)
|
||||
audio_cond = multitalk_audio.to(device=x.device, dtype=x.dtype)
|
||||
first_frame_audio_emb_s = audio_cond[:, :1, ...]
|
||||
latter_frame_audio_emb = audio_cond[:, 1:, ...]
|
||||
latter_frame_audio_emb = rearrange(latter_frame_audio_emb, "b (n_t n) w s c -> b n_t n w s c", n=4)
|
||||
middle_index = self.audio_proj.seq_len // 2
|
||||
middle_index = self.multitalk_audio_proj.seq_len // 2
|
||||
latter_first_frame_audio_emb = latter_frame_audio_emb[:, :, :1, :middle_index+1, ...]
|
||||
latter_first_frame_audio_emb = rearrange(latter_first_frame_audio_emb, "b n_t n w s c -> b n_t (n w) s c")
|
||||
latter_last_frame_audio_emb = latter_frame_audio_emb[:, :, -1:, middle_index:, ...]
|
||||
@@ -2292,10 +2292,10 @@ class WanModel(torch.nn.Module):
|
||||
latter_middle_frame_audio_emb = latter_frame_audio_emb[:, :, 1:-1, middle_index:middle_index+1, ...]
|
||||
latter_middle_frame_audio_emb = rearrange(latter_middle_frame_audio_emb, "b n_t n w s c -> b n_t (n w) s c")
|
||||
latter_frame_audio_emb_s = torch.concat([latter_first_frame_audio_emb, latter_middle_frame_audio_emb, latter_last_frame_audio_emb], dim=2)
|
||||
multitalk_audio_embedding = self.audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s)
|
||||
multitalk_audio_embedding = self.multitalk_audio_proj(first_frame_audio_emb_s, latter_frame_audio_emb_s)
|
||||
human_num = len(multitalk_audio_embedding)
|
||||
multitalk_audio_embedding = torch.concat(multitalk_audio_embedding.split(1), dim=2).to(x.dtype)
|
||||
self.audio_proj.to(self.offload_device)
|
||||
self.multitalk_audio_proj.to(self.offload_device)
|
||||
|
||||
# convert ref_target_masks to token_ref_target_masks
|
||||
token_ref_target_masks = None
|
||||
|
||||
Reference in New Issue
Block a user