Fix up temporal masking

This commit is contained in:
kijai
2025-03-18 09:09:41 +02:00
parent c4ed468967
commit 557fe5c875
2 changed files with 16 additions and 24 deletions
+6 -3
View File
@@ -84,12 +84,12 @@ def get_previewer(device, latent_format):
taew_sd = comfy.utils.load_torch_file(taehv_path)
taesd = TAEHV(taew_sd).to(device)
previewer = TAESDPreviewerImpl(taesd)
previewer = WrappedPreviewer(previewer)
previewer = WrappedPreviewer(previewer, rate=16)
if previewer is None:
if latent_format.latent_rgb_factors is not None:
previewer = Latent2RGBPreviewer(latent_format.latent_rgb_factors, latent_format.latent_rgb_factors_bias)
previewer = WrappedPreviewer(previewer)
previewer = WrappedPreviewer(previewer, rate=4)
return previewer
def prepare_callback(model, steps, x0_output_dict=None):
@@ -182,7 +182,10 @@ class WrappedPreviewer(LatentPreviewer):
#NOTE: send sync already uses call_soon_threadsafe
serv.send_sync(server.BinaryEventTypes.PREVIEW_IMAGE,
message.getvalue(), serv.client_id)
ind = (ind + 1) % ((leng-1) * 4 - 1)
if self.rate == 16:
ind = (ind + 1) % ((leng-1) * 4 - 1)
else:
ind = (ind + 1) % leng
def decode_latent_to_preview(self, x0):
if hasattr(self, 'taesd'):
x0 = x0.unsqueeze(0)
+10 -21
View File
@@ -1694,6 +1694,7 @@ class WanVideoSampler:
)
mask = masks[idx]
mask = mask.to(latent)
print(mask.shape, image_latent.shape, latent.shape)
latent = image_latent * mask + latent * (1-mask)
# end diff diff
@@ -2105,33 +2106,21 @@ class WanVideoEncode:
vae.to(offload_device)
mm.soft_empty_cache()
print("encoded latents shape",latents.shape)
log.info(f"encoded latents shape {latents.shape}")
if mask is not None: #B, H, W
B, H, W = mask.shape
target_frames = latents.shape[2]
target_h, target_w = latents.shape[3:]
# Temporal: pad/truncate
if B > target_frames:
mask = mask[:target_frames]
elif B < target_frames:
padding = torch.zeros((target_frames - B, H, W), device=mask.device)
mask = torch.cat([mask, padding], dim=0)
# Spatial: resize each frame
mask = torch.nn.functional.interpolate(
mask.unsqueeze(1), # Add channel dim for interpolate
size=(target_h, target_w),
mode='bilinear'
).squeeze(1) # Remove channel dim
mask.unsqueeze(0).unsqueeze(0), # Add batch and channel dims [1,1,T,H,W]
size=(latents.shape[2], target_h, target_w),
mode='trilinear',
align_corners=False
).squeeze(0) # Remove batch dim, keep channel dim
# Add batch & channel dims for final output
mask = mask.unsqueeze(0).unsqueeze(0)
mask = mask.repeat(1, latents.shape[1], 1, 1, 1)
print("mask shape",mask.shape)
mask = mask.unsqueeze(0).repeat(1, latents.shape[1], 1, 1, 1)
return ({"samples": latents, "mask": mask},)
class WanVideoLatentPreview: