This commit is contained in:
Eugene Fairley 2025-06-28 22:53:55 -07:00
parent 54dbbc428e
commit e643977d61

View File

@ -1188,7 +1188,7 @@ class WAN21(BaseModel):
image = image.to(device) image = image.to(device)
image = utils.common_upscale(image.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center") image = utils.common_upscale(image.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center")
for i in range(0, image.shape[1], 16): for i in range(0, image.shape[1], 16):
image[:, i: i + 16] = self.process_latent_in(image[:, i: i + 36]) image[:, i: i + 16] = self.process_latent_in(image[:, i: i + 16])
image = utils.resize_to_batch_size(image, noise.shape[0]) image = utils.resize_to_batch_size(image, noise.shape[0])
print(f"image shape: {image.shape}") print(f"image shape: {image.shape}")
@ -1199,25 +1199,25 @@ class WAN21(BaseModel):
image = image[:, :(extra_channels - 4)] image = image[:, :(extra_channels - 4)]
mask = kwargs.get("concat_mask", kwargs.get("denoise_mask", None)).to(device) if "concat_mask" in kwargs or "denoise_mask" in kwargs else None mask = kwargs.get("concat_mask", kwargs.get("denoise_mask", None)).to(device) if "concat_mask" in kwargs or "denoise_mask" in kwargs else None
if mask is None: # if mask is None:
mask = torch.zeros_like(noise)[:, :4] # mask = torch.zeros_like(noise)[:, :4]
else: # else:
if mask.shape[1] != 4: # if mask.shape[1] != 4:
mask = torch.mean(mask, dim=1, keepdim=True) # mask = torch.mean(mask, dim=1, keepdim=True)
mask = 1.0 - mask # mask = 1.0 - mask
mask = utils.common_upscale(mask.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center") # mask = utils.common_upscale(mask.to(device), noise.shape[-1], noise.shape[-2], "bilinear", "center")
if mask.shape[-3] < noise.shape[-3]: # if mask.shape[-3] < noise.shape[-3]:
mask = torch.nn.functional.pad(mask, (0, 0, 0, 0, 0, noise.shape[-3] - mask.shape[-3]), mode='constant', value=0) # mask = torch.nn.functional.pad(mask, (0, 0, 0, 0, 0, noise.shape[-3] - mask.shape[-3]), mode='constant', value=0)
if mask.shape[1] == 1: # if mask.shape[1] == 1:
mask = mask.repeat(1, 4, 1, 1, 1) # mask = mask.repeat(1, 4, 1, 1, 1)
print(f"Mask shape: {mask.shape}, noise shape: {noise.shape}") # print(f"Mask shape: {mask.shape}, noise shape: {noise.shape}")
mask = utils.resize_to_batch_size(mask, noise.shape[0]) # mask = utils.resize_to_batch_size(mask, noise.shape[0])
print(f"image shape: {image.shape}, mask shape: {mask.shape}") print(f"image shape: {image.shape}, mask shape: {mask.shape}")
res = torch.cat((mask, image), dim=1) res = torch.cat((mask, image), dim=1)
tracks = kwargs.get("tracks", None) tracks = kwargs.get("tracks", None)
if tracks is not None: if tracks is not None:
res = patch_motion(tracks, res, 220.0, (4, 16), 2)[None] res = patch_motion(tracks.to(device), res, 220.0, (4, 16), 2)[None]
return res return res