From 36b564c1abe7e2becf9702ac39b86aaa187615a0 Mon Sep 17 00:00:00 2001 From: smthemex <138738845+smthemex@users.noreply.github.com> Date: Thu, 21 Aug 2025 15:52:46 +0800 Subject: [PATCH] Update wan_inference_long_pipeline.py --- .../pipeline/wan_inference_long_pipeline.py | 74 ++++++++++++++----- 1 file changed, 55 insertions(+), 19 deletions(-) diff --git a/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py b/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py index 94463b2..d91ea6e 100644 --- a/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py +++ b/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py @@ -233,6 +233,22 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): do_convert_grayscale=True ) + def enhance_temporal_consistency(self, latents, prev_latents, weight=0.2): + """ + 增强时间一致性 + """ + # 计算时间差异并应用一致性约束 + if prev_latents is not None: + + if latents.shape[2] < prev_latents.shape[2]: + prev_latents_cropped = prev_latents[:, :, :latents.shape[2], :, :] + diff = latents - prev_latents_cropped + else: + diff = latents - prev_latents + # 应用平滑约束 + consistency_correction = diff * weight + return latents - consistency_correction + return latents def _get_t5_prompt_embeds( @@ -369,6 +385,7 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): f" size of {batch_size}. Make sure the batch size matches the length of the generators." ) + shape = ( batch_size, num_channels_latents, @@ -645,7 +662,7 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): audio_token_per_frame = int(sr / fps) max_audio_index = vocal_input_values.shape[0] total_frames = int(max_audio_index / audio_token_per_frame) - #print("total_frames" ,total_frames) #total_frames 122 + print("total_frames is" ,total_frames) #total_frames 122 frames_per_batch = 21 if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler): @@ -682,7 +699,7 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): ) #print(f"latents.shape: {latents.shape}")#latents.shape: torch.Size([1, 16, 31, 64, 64]) infer_length = latents.size()[2] - #print(f"infer_length: {infer_length}") #31 + #print(f"audio infer length is : {infer_length}") #31 num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0) latents_all = latents.clone() @@ -739,10 +756,9 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): for i, t in enumerate(timesteps): if isinstance(self.scheduler, FlowUniPCMultistepScheduler): self.scheduler.set_timesteps(num_inference_steps, device=device) - # elif isinstance(self.scheduler, WanStepDistillScheduler): - # self.scheduler.set_denoising_timesteps(device=device) + elif isinstance(self.scheduler, WanStepDistillScheduler): + self.scheduler.set_denoising_timesteps(device=device) - if self.interrupt: continue @@ -752,42 +768,51 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): overlap_window_length = overlap_window_length # 5 # [5, 7, 10] longer length --> higher quality index_end = index_start + frames_per_batch # 0+21 index_previous_end = index_end # 0+21 - - + while index_end <= infer_length: #31 - if hasattr(self.scheduler, '_step_index'): - self.scheduler._step_index = None - + # if hasattr(self.scheduler, '_step_index'): + # self.scheduler._step_index = None + + idx_list = [ii for ii in range(index_start, min(index_end, infer_length))] - #print(f"idx_list: {idx_list}") - # 修复:正确计算音频索引 + start_audio_idx = int(index_start * 4 * audio_token_per_frame) if index_end == infer_length: - # 最后一段,使用剩余的所有音频 + idx_list_audio = list(range(start_audio_idx, max_audio_index)) else: - # 中间段,计算对应的音频范围 + end_audio_idx = int(index_end * 4 * audio_token_per_frame) end_audio_idx = min(end_audio_idx, max_audio_index) idx_list_audio = list(range(start_audio_idx, end_audio_idx)) ## idx_list_audio = [ii % max_audio_index for ii in range(index_start * 4 * audio_token_per_frame, index_end * 4 * audio_token_per_frame)] latents = latents_all[:, :, idx_list].clone() - if latents.shape[2] < frames_per_batch and isinstance(self.scheduler, FlowUniPCMultistepScheduler) : - pad_shape = list(latents.shape) - pad_shape[2] = frames_per_batch - latents.shape[2] - pad = torch.zeros(pad_shape, dtype=latents.dtype, device=latents.device) - latents = torch.cat([latents, pad], dim=2) + + # 检查并重置scheduler状态 + if isinstance(self.scheduler, FlowUniPCMultistepScheduler): + self.scheduler.prepare_for_step_with_different_shape(latents.shape) + + if isinstance(self.scheduler, WanStepDistillScheduler) : + self.scheduler.prepare_for_step_with_different_shape(latents.shape) + overlapping_weight_scheme == "LCM" + + if hasattr(self.scheduler, '_step_index'): + self.scheduler._step_index = None + sub_vocal_input_values = vocal_input_values[idx_list_audio] sub_vocal_input_values = self.wav2vec_processor(sub_vocal_input_values, sampling_rate=sr, return_tensors="pt").input_values.to(device) sub_vocal_embeddings = self.wav2vec(sub_vocal_input_values).last_hidden_state + latent_model_input = torch.cat([latents] * 3) if do_classifier_free_guidance else latents + if hasattr(self.scheduler, "scale_model_input"): latent_model_input = self.scheduler.scale_model_input(latent_model_input, t) timestep = t.expand(latent_model_input.shape[0]) target_shape = (self.vae.latent_channels, (num_frames - 1) // self.vae.temporal_compression_ratio + 1, width // self.vae.spacial_compression_ratio, height // self.vae.spacial_compression_ratio) seq_len = math.ceil((target_shape[2] * target_shape[3]) / (self.transformer.config.patch_size[1] * self.transformer.config.patch_size[2]) * target_shape[1]) + if text_guide_scale is not None and audio_guide_scale is not None: sub_vocal_embeddings = torch.cat([torch.zeros_like(sub_vocal_embeddings), sub_vocal_embeddings, sub_vocal_embeddings], dim=0) @@ -807,8 +832,11 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): if do_classifier_free_guidance: noise_pred_uncond, noise_pred_drop_audio, noise_pred_cond = noise_pred.chunk(3) noise_pred = noise_pred_uncond + audio_guide_scale * (noise_pred_drop_audio - noise_pred_uncond) + text_guide_scale * (noise_pred_cond - noise_pred_drop_audio) + latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0] + torch.cuda.empty_cache() + if index_start != 0 and i != 0: overlap_window_weight = torch.zeros(1, 1, overlap_window_length, 1, 1).to(device=latents.device, dtype=latents.dtype) if overlapping_weight_scheme == "uniform": @@ -820,8 +848,14 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): norm_weights = (init_weight - init_weight.min()) / (init_weight.max() - init_weight.min()) for j in range(overlap_window_length): overlap_window_weight[:, :, j] = norm_weights[j] + else: + # 默认线性过渡 + for j in range(overlap_window_length): + overlap_window_weight[:, :, j] = j / (overlap_window_length-1) + overlap_idx_list_start = [ii % latents.shape[2] for ii in range(0, overlap_window_length)] overlap_idx_list_end = [ii % latents_all.shape[2] for ii in range(index_previous_end-overlap_window_length, index_previous_end)] + latents[:, :, overlap_idx_list_start] = latents[:, :, overlap_idx_list_start] * overlap_window_weight + pred_latents[:, :, overlap_idx_list_end] * (1-overlap_window_weight) latents = latents.to(torch.bfloat16) for iii in range(legal_compressed_frames_num): @@ -886,3 +920,5 @@ def apply_scale_to_latent(latent, mean, std_inv): # 反标准化: z = z / (1/std) + mean scaled_latent = latent / std_inv.view(1, -1, 1, 1, 1) + mean.view(1, -1, 1, 1, 1) return scaled_latent + +