From 9a88b9e40a2108b74b37bce9d2f563a3e37e3d00 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 15 Oct 2025 19:11:41 +0300 Subject: [PATCH] FlashVSR: Add strength setting --- FlashVSR/flashvsr_nodes.py | 4 +++- nodes_sampler.py | 6 ++++-- wanvideo/modules/model.py | 4 ++-- 3 files changed, 9 insertions(+), 5 deletions(-) diff --git a/FlashVSR/flashvsr_nodes.py b/FlashVSR/flashvsr_nodes.py index 81f289d..6930294 100644 --- a/FlashVSR/flashvsr_nodes.py +++ b/FlashVSR/flashvsr_nodes.py @@ -12,6 +12,7 @@ class WanVideoAddFlashVSRInput: return {"required": { "embeds": ("WANVIDIMAGE_EMBEDS",), "images": ("IMAGE", {"tooltip": "Low-res video frames to enhance"}), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 2.0, "step": 0.01, "tooltip": "Strength to apply the FlashVSR latent"}), } } @@ -20,9 +21,10 @@ class WanVideoAddFlashVSRInput: FUNCTION = "add" CATEGORY = "WanVideoWrapper" - def add(self, embeds, images): + def add(self, embeds, images, strength): updated = dict(embeds) updated["flashvsr_LQ_images"] = images + updated["flashvsr_strength"] = strength return (updated,) diff --git a/nodes_sampler.py b/nodes_sampler.py index 7901ba0..99742b0 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -832,12 +832,14 @@ class WanVideoSampler: # FlashVSR flashvsr_LQ_latent = LQ_images = None flashvsr_LQ_images = image_embeds.get("flashvsr_LQ_images", None) + flashvsr_strength = image_embeds.get("flashvsr_strength", 1.0) if flashvsr_LQ_images is not None: LQ_images = flashvsr_LQ_images.unsqueeze(0).movedim(-1, 1).to(device, dtype) * 2 - 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] + if noise.shape[1] != 1: + noise = noise[:, :-1] seq_len = math.ceil((noise.shape[2] * noise.shape[3]) / 4 * noise.shape[1]) latent = noise @@ -1346,6 +1348,7 @@ class WanVideoSampler: "seq_len_ovi": seq_len_ovi, # Audio latent model sequence length for Ovi "ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi "flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling + "flashvsr_strength": flashvsr_strength, # FlashVSR strength } batch_size = 1 @@ -1950,7 +1953,6 @@ class WanVideoSampler: 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) diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 22f2439..525d57d 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -2153,7 +2153,7 @@ class WanModel(torch.nn.Module): wananim_pose_strength=1.0, wananim_face_strength=1.0, lynx_embeds=None, x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None, - flashvsr_LQ_latent=None, + flashvsr_LQ_latent=None, flashvsr_strength=1.0, ): r""" Forward pass through the diffusion model @@ -2810,7 +2810,7 @@ class WanModel(torch.nn.Module): lynx_ref_feature = None # FlashVSR if flashvsr_LQ_latent is not None and b < len(flashvsr_LQ_latent): - x += flashvsr_LQ_latent[b].to(x) + x += flashvsr_LQ_latent[b].to(x) * flashvsr_strength # Prefetch blocks if enabled if self.prefetch_blocks > 0: for prefetch_offset in range(1, self.prefetch_blocks + 1):