FlashVSR: Add strength setting

This commit is contained in:
kijai
2025-10-15 19:11:41 +03:00
parent 2fd5bb6ffa
commit 9a88b9e40a
3 changed files with 9 additions and 5 deletions
+3 -1
View File
@@ -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,)
+4 -2
View File
@@ -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)
+2 -2
View File
@@ -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):