Basic context window support for FlashVSR
This commit is contained in:
+20
-8
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user