diff --git a/nodes_sampler.py b/nodes_sampler.py index d41a27a..7901ba0 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -830,14 +830,15 @@ class WanVideoSampler: extra_channel_latents = extra_channel_latents[0].to(noise) # FlashVSR - flashvsr_LQ_latent = None + flashvsr_LQ_latent = LQ_images = None flashvsr_LQ_images = image_embeds.get("flashvsr_LQ_images", None) if flashvsr_LQ_images is not None: LQ_images = flashvsr_LQ_images.unsqueeze(0).movedim(-1, 1).to(device, dtype) * 2 - 1 - flashvsr_LQ_latent = transformer.LQ_proj_in(LQ_images) - log.info(f"flashvsr_LQ_latent: {flashvsr_LQ_latent[0].shape}") - noise = noise[:, :-1] - seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1]) + if context_options is None: + flashvsr_LQ_latent = transformer.LQ_proj_in(LQ_images) + log.info(f"flashvsr_LQ_latent: {flashvsr_LQ_latent[0].shape}") + noise = noise[:, :-1] + seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1]) latent = noise @@ -1109,7 +1110,7 @@ class WanVideoSampler: 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, humo_image_cond=None, humo_image_cond_neg=None, humo_audio=None, humo_audio_neg=None, wananim_pose_latents=None, - wananim_face_pixels=None, uni3c_data=None, latent_model_input_ovi=None): + wananim_face_pixels=None, uni3c_data=None, latent_model_input_ovi=None, flashvsr_LQ_latent=None,): nonlocal transformer autocast_enabled = ("fp8" in model["quantization"] and not transformer.patched_linear) @@ -1943,6 +1944,16 @@ class WanVideoSampler: center_indices = torch.clamp(center_indices, min=0, max=wananim_pose_latents.shape[2] - 1) partial_wananim_pose_latents = wananim_pose_latents[:, :, center_indices][:, :, :context_frames-1].to(device, dtype) + partial_flashvsr_LQ_latent = None + if LQ_images is not None: + start = c[0] * 4 + end = c[-1] * 4 + 1 + 4 + center_indices = torch.arange(start, end, 1) + center_indices = torch.clamp(center_indices, min=0, max=LQ_images.shape[2] - 1) + print("FlashVSR LQ image indices:", center_indices) + partial_flashvsr_LQ_images = LQ_images[:, :, center_indices].to(device, dtype) + partial_flashvsr_LQ_latent = transformer.LQ_proj_in(partial_flashvsr_LQ_images) + if len(timestep.shape) != 1: partial_timestep = timestep[:, c] partial_timestep[:, :1] = 0 @@ -1958,7 +1969,8 @@ class WanVideoSampler: 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, humo_image_cond=humo_image_cond, humo_image_cond_neg=humo_image_cond_neg, humo_audio=humo_audio, humo_audio_neg=humo_audio_neg, - wananim_face_pixels=partial_wananim_face_pixels, wananim_pose_latents=partial_wananim_pose_latents, multitalk_audio_embeds=multitalk_audio_embeds) + wananim_face_pixels=partial_wananim_face_pixels, wananim_pose_latents=partial_wananim_pose_latents, multitalk_audio_embeds=multitalk_audio_embeds, + flashvsr_LQ_latent=partial_flashvsr_LQ_latent) if cache_args is not None: self.window_tracker.cache_states[window_id] = new_teacache @@ -2847,7 +2859,7 @@ class WanVideoSampler: 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, 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, - wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, uni3c_data = uni3c_data, latent_model_input_ovi=latent_model_input_ovi + wananim_face_pixels=wananim_face_pixels, wananim_pose_latents=wananim_pose_latents, uni3c_data = uni3c_data, latent_model_input_ovi=latent_model_input_ovi, flashvsr_LQ_latent=flashvsr_LQ_latent, ) if bidirectional_sampling: noise_pred_flipped, _,self.cache_state = predict_with_cfg(