mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-03 01:07:05 +08:00
save
This commit is contained in:
parent
54dbbc428e
commit
e643977d61
@ -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
|
||||||
|
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user