Update wan_inference_long_pipeline.py

This commit is contained in:
smthemex
2025-08-15 13:58:17 +08:00
committed by GitHub
parent f3dc0cc55b
commit fa7134ae38
@@ -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
return scaled_latent