FlashVSR: Add strength setting
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user