Fix up temporal masking
This commit is contained in:
+6
-3
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user