From fa7134ae386c2f6bb9b623fd17a1cc6f19ef53ea Mon Sep 17 00:00:00 2001 From: smthemex <138738845+smthemex@users.noreply.github.com> Date: Fri, 15 Aug 2025 13:58:17 +0800 Subject: [PATCH] Update wan_inference_long_pipeline.py --- .../wan/pipeline/wan_inference_long_pipeline.py | 16 ++++++++++++---- 1 file changed, 12 insertions(+), 4 deletions(-) diff --git a/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py b/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py index 1501742..3857412 100644 --- a/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py +++ b/StableAvatar/wan/pipeline/wan_inference_long_pipeline.py @@ -578,6 +578,7 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): clip_image_tensor=None, clip_context=None, weight_dtype=torch.bfloat16, + overlapping_weight_scheme="uniform", ) -> Union[WanI2VPipelineTalkingInferenceLongOutput, Tuple]: """ Function invoked when calling the pipeline for generation. @@ -767,7 +768,7 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): sub_vocal_embeddings = torch.cat([torch.zeros_like(sub_vocal_embeddings), sub_vocal_embeddings, sub_vocal_embeddings], dim=0) - with torch.cuda.amp.autocast(dtype=weight_dtype): + with torch.amp.autocast('cuda', dtype=weight_dtype): legal_compressed_frames_num = latents.size()[2] #print(1,latents.shape) #torch.Size([1, 16, 17, 64, 64]) #torch.Size([1, 16, 21, 64, 64]) #print(2,y[:, :, :legal_compressed_frames_num].shape) # 2 torch.Size([3, 20, 17, 64, 64]) torch.Size([3, 20, 21, 64, 64]) @@ -789,8 +790,15 @@ class WanI2VTalkingInferenceLongPipeline(DiffusionPipeline): 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) - for j in range(overlap_window_length): - overlap_window_weight[:, :, j] = j / (overlap_window_length-1) + if overlapping_weight_scheme == "uniform": + for j in range(overlap_window_length): + overlap_window_weight[:, :, j] = j / (overlap_window_length-1) + elif overlapping_weight_scheme == "log": + init_weight = torch.linspace(0, 1, overlap_window_length) + init_weight = torch.log1p(init_weight * (torch.exp(torch.tensor(1.0)) - 1)) + 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] 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) @@ -856,4 +864,4 @@ 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 \ No newline at end of file + return scaled_latent