vace ref image fixes

This commit is contained in:
kijai
2025-04-02 00:44:47 +03:00
parent 80bcac3803
commit 3322d868ab
+25 -5
View File
@@ -1656,7 +1656,24 @@ class WanVideoVACEEncode:
input_masks = input_masks.unsqueeze(-1).unsqueeze(0).permute(0, 4, 1, 2, 3).repeat(1, 3, 1, 1, 1) # B, C, T, H, W
if ref_images is not None:
ref_images = common_upscale(ref_images.clone().movedim(-1, 1), width, height, "lanczos", "disabled").movedim(1, -1)
# Create padded image
B, H, W, C = ref_images.shape
current_aspect = W / H
target_aspect = width / height
if current_aspect > target_aspect:
# Image is wider than target, pad height
new_h = int(W / target_aspect)
pad_h = (new_h - H) // 2
padded = torch.ones(ref_images.shape[0], new_h, W, ref_images.shape[3], device=ref_images.device, dtype=ref_images.dtype)
padded[:, pad_h:pad_h+H, :, :] = ref_images
else:
# Image is taller than target, pad width
new_w = int(H * target_aspect)
pad_w = (new_w - W) // 2
padded = torch.ones(ref_images.shape[0], H, new_w, ref_images.shape[3], device=ref_images.device, dtype=ref_images.dtype)
padded[:, :, pad_w:pad_w+W, :] = ref_images
ref_images = common_upscale(padded.movedim(-1, 1), width, height, "lanczos", "center").movedim(1, -1)
ref_images = ref_images.to(self.vae.dtype).to(self.device).unsqueeze(0).permute(0, 4, 1, 2, 3).unsqueeze(0)
ref_images = ref_images * 2 - 1
print("ref_images shape", ref_images.shape)
@@ -1943,7 +1960,7 @@ class WanVideoSampler:
control_latents, clip_fea, clip_fea_neg, end_image = None, None, None, None
vace_context, vace_scale = None, None
fun_model = False
fun_model, has_ref = False, False
image_cond = image_embeds.get("image_embeds", None)
@@ -2669,7 +2686,7 @@ class WanVideoSampler:
pass
return ({
"samples": x0.unsqueeze(0).cpu(), "looped": is_looped, "end_image": end_image if not fun_model else None
"samples": x0.unsqueeze(0).cpu(), "looped": is_looped, "end_image": end_image if not fun_model else None, "has_ref": has_ref
}, )
class WindowTracker:
@@ -2735,6 +2752,8 @@ class WanVideoDecode:
mm.soft_empty_cache()
latents = samples["samples"]
end_image = samples.get("end_image", None)
has_ref = samples.get("has_ref", False)
is_looped = samples.get("looped", False)
vae.to(device)
@@ -2742,8 +2761,9 @@ class WanVideoDecode:
mm.soft_empty_cache()
is_looped = samples.get("looped", False)
warmup_latent_count = 3
if has_ref:
latents = latents[:, :, 1:]
#if is_looped:
# latents = torch.cat([latents[:, :, :warmup_latent_count],latents], dim=2)