cleanup experimental context window stuff

Didn't really work anyway
This commit is contained in:
kijai
2025-04-29 17:36:07 +03:00
parent b78505014c
commit 4a41fb0aaf
+11 -49
View File
@@ -2151,9 +2151,7 @@ class WanVideoContextOptions:
"verbose": ("BOOLEAN", {"default": False, "tooltip": "Print debug output"}),
},
"optional": {
"image_cond_start_step": ("INT", {"default": 6, "min": 0, "max": 10000, "step": 1, "tooltip": "!EXPERIMENTAL! Start step of using previous window results as input instead of the init image"}),
"image_cond_window_count": ("INT", {"default": 2, "min": 1, "max": 10000, "step": 1, "tooltip": "!EXPERIMENTAL! Number of image 'prompt windows'"}),
"vae": ("WANVAE",),
"vae": ("WANVAE",),
}
}
@@ -2171,8 +2169,6 @@ class WanVideoContextOptions:
"context_overlap":context_overlap,
"freenoise":freenoise,
"verbose":verbose,
"image_cond_start_step": image_cond_start_step,
"image_cond_window_count": image_cond_window_count,
"vae": vae,
}
@@ -3207,56 +3203,28 @@ class WanVideoSampler:
partial_img_emb = None
partial_control_latents = None
if image_cond is not None:
log.info(f"Image cond shape: {image_cond.shape}")
num_windows= context_options["image_cond_window_count"]
section_size = latent_video_length / num_windows
image_index = min(int(max(c) / section_size), num_windows - 1)
partial_img_emb = image_cond[:, c, :, :]
partial_img_emb = image_cond[:, c]
partial_img_emb[:, 0] = image_cond[:, 0].to(intermediate_device)
if control_latents is not None:
partial_control_latents = control_latents[:, c, :, :]
partial_image_cond = image_cond[:, 0, :, :].to(intermediate_device)
log.info(f"image_index: {image_index}")
if hasattr(self, "previous_noise_pred_context") and image_index > 0: #wip
if idx >= context_options["image_cond_start_step"]:
#strength = 0.5
#partial_image_cond *= strength
mask = torch.ones(4, partial_img_emb.shape[2], partial_img_emb.shape[3], device=partial_img_emb.device, dtype=partial_img_emb.dtype) #torch.Size([20, 10, 104, 60])
if context_vae is not None:
to_decode = self.previous_noise_pred_context[:,-1,:, :].unsqueeze(1).unsqueeze(0).to(context_vae.dtype)
#to_decode = to_decode.permute(0, 1, 3, 2)
#print("to_decode.shape", to_decode.shape)
if isinstance(context_vae, TAEHV):
image = context_vae.decode_video(to_decode.permute(0, 2, 1, 3, 4), parallel=False)
image = context_vae.encode_video(image.repeat(1, 5, 1, 1, 1), parallel=False).permute(0, 2, 1, 3, 4)
else:
image = context_vae.decode(to_decode, device=device, tiled=False)[0]
image = context_vae.encode(image.unsqueeze(0).to(context_vae.dtype), device=device, tiled=False)
#print("decoded image.shape", image.shape) #torch.Size([3, 37, 832, 480])
#print("encoded image.shape", image.shape)
#partial_img_emb[:, 0, :, :] = image[0][:,0,:,:]
#print("partial_img_emb.shape", partial_img_emb.shape)
#print("mask.shape", mask.shape)
#print("self.previous_noise_pred_context.shape", self.previous_noise_pred_context.shape) #torch.Size([16, 10, 104, 60])
partial_img_emb[:, 0, :, :] = torch.cat([image[0][:,0,:,:], mask], dim=0)
else:
partial_img_emb[:, 0, :, :] = partial_image_cond
partial_control_latents = control_latents[:, c]
partial_vace_context = None
if vace_data is not None:
partial_vace_context = vace_data[0]["context"][0][:, c, :, :]
partial_vace_context = vace_data[0]["context"][0][:, c]
if has_ref:
partial_vace_context[:, 0, :, :] = vace_data[0]["context"][0][:, 0, :, :]
partial_vace_context[:, 0] = vace_data[0]["context"][0][:, 0]
partial_vace_context = [partial_vace_context]
partial_audio_proj = None
if fantasytalking_embeds is not None:
partial_audio_proj = audio_proj[:, c]
partial_latent_model_input = latent_model_input[:, c, :, :]
partial_latent_model_input = latent_model_input[:, c]
partial_unianim_data = None
if unianim_data is not None:
partial_dwpose = dwpose_data[:, :, c, :, :]
partial_dwpose = dwpose_data[:, :, c]
partial_dwpose_flat=rearrange(partial_dwpose, 'b c f h w -> b (f h w) c')
partial_unianim_data = {
"dwpose": partial_dwpose_flat,
@@ -3273,18 +3241,12 @@ class WanVideoSampler:
timestep, idx, partial_img_emb, clip_fea, partial_control_latents, partial_vace_context, partial_unianim_data,partial_audio_proj,
current_teacache)
# if callback is not None:
# callback_latent = (noise_pred.to(t.device) * t / 1000).detach().permute(1,0,2,3)
# callback(idx, callback_latent, None, steps)
if teacache_args is not None:
self.window_tracker.teacache_states[window_id] = new_teacache
if image_cond is not None and image_index > 0:
self.previous_noise_pred_context = noise_pred_context
window_mask = create_window_mask(noise_pred_context, c, latent_video_length, context_overlap, looped=is_looped)
noise_pred[:, c, :, :] += noise_pred_context * window_mask
counter[:, c, :, :] += window_mask
noise_pred[:, c] += noise_pred_context * window_mask
counter[:, c] += window_mask
noise_pred /= counter
#normal inference
else: