This commit is contained in:
Eugene Fairley 2025-06-28 22:46:26 -07:00
parent af03f5dbe3
commit 54dbbc428e
2 changed files with 146 additions and 27 deletions

View File

@ -1056,6 +1056,115 @@ class Lumina2(BaseModel):
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
return out return out
def ind_sel(target: torch.Tensor, ind: torch.Tensor, dim: int = 1):
"""Index selection utility function"""
assert (
len(ind.shape) > dim
), "Index must have the target dim, but get dim: %d, ind shape: %s" % (dim, str(ind.shape))
target = target.expand(
*tuple(
[ind.shape[k] if target.shape[k] == 1 else -1 for k in range(dim)]
+ [
-1,
]
* (len(target.shape) - dim)
)
)
ind_pad = ind
if len(target.shape) > dim + 1:
for _ in range(len(target.shape) - (dim + 1)):
ind_pad = ind_pad.unsqueeze(-1)
ind_pad = ind_pad.expand(*(-1,) * (dim + 1), *target.shape[(dim + 1) : :])
return torch.gather(target, dim=dim, index=ind_pad)
def merge_final(vert_attr: torch.Tensor, weight: torch.Tensor, vert_assign: torch.Tensor):
"""Merge vertex attributes with weights"""
target_dim = len(vert_assign.shape) - 1
if len(vert_attr.shape) == 2:
assert vert_attr.shape[0] > vert_assign.max()
new_shape = [1] * target_dim + list(vert_attr.shape)
tensor = vert_attr.reshape(new_shape)
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
else:
assert vert_attr.shape[1] > vert_assign.max()
new_shape = [vert_attr.shape[0]] + [1] * (target_dim - 1) + list(vert_attr.shape[1:])
tensor = vert_attr.reshape(new_shape)
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
final_attr = torch.sum(sel_attr * weight.unsqueeze(-1), dim=-2)
return final_attr
def patch_motion(
tracks: torch.FloatTensor, # (B, T, N, 4)
vid: torch.FloatTensor, # (C, T, H, W)
temperature: float = 220.0,
vae_divide: tuple = (4, 16),
topk: int = 2,
):
"""Apply motion patching based on tracks"""
with torch.no_grad():
print("vid shape:", vid)
_, T, H, W = vid.shape
N = tracks.shape[2]
_, tracks_xy, visible = torch.split(
tracks, [1, 2, 1], dim=-1
) # (B, T, N, 2) | (B, T, N, 1)
tracks_n = tracks_xy / torch.tensor([W / min(H, W), H / min(H, W)], device=tracks_xy.device)
tracks_n = tracks_n.clamp(-1, 1)
visible = visible.clamp(0, 1)
xx = torch.linspace(-W / min(H, W), W / min(H, W), W)
yy = torch.linspace(-H / min(H, W), H / min(H, W), H)
grid = torch.stack(torch.meshgrid(yy, xx, indexing="ij")[::-1], dim=-1).to(
tracks_xy.device
)
tracks_pad = tracks_xy[:, 1:]
visible_pad = visible[:, 1:]
visible_align = visible_pad.view(T - 1, 4, *visible_pad.shape[2:]).sum(1)
tracks_align = (tracks_pad * visible_pad).view(T - 1, 4, *tracks_pad.shape[2:]).sum(
1
) / (visible_align + 1e-5)
dist_ = (
(tracks_align[:, None, None] - grid[None, :, :, None]).pow(2).sum(-1)
) # T, H, W, N
weight = torch.exp(-dist_ * temperature) * visible_align.clamp(0, 1).view(
T - 1, 1, 1, N
)
vert_weight, vert_index = torch.topk(
weight, k=min(topk, weight.shape[-1]), dim=-1
)
grid_mode = "bilinear"
point_feature = torch.nn.functional.grid_sample(
vid[vae_divide[0]:].permute(1, 0, 2, 3)[:1],
tracks_n[:, :1].type(vid.dtype),
mode=grid_mode,
padding_mode="zeros",
align_corners=False,
)
point_feature = point_feature.squeeze(0).squeeze(1).permute(1, 0) # N, C=16
out_feature = merge_final(point_feature, vert_weight, vert_index).permute(3, 0, 1, 2) # T - 1, H, W, C => C, T - 1, H, W
out_weight = vert_weight.sum(-1) # T - 1, H, W
# out feature -> already soft weighted
mix_feature = out_feature + vid[vae_divide[0]:, 1:] * (1 - out_weight.clamp(0, 1))
out_feature_full = torch.cat([vid[vae_divide[0]:, :1], mix_feature], dim=1) # C, T, H, W
print("out_feature_full:", out_feature_full)
out_mask_full = torch.cat([torch.ones_like(out_weight[:1]), out_weight], dim=0) # T, H, W
return torch.cat([out_mask_full[None].expand(vae_divide[0], -1, -1, -1), out_feature_full], dim=0)
class WAN21(BaseModel): class WAN21(BaseModel):
def __init__(self, model_config, model_type=ModelType.FLOW, image_to_video=False, device=None): def __init__(self, model_config, model_type=ModelType.FLOW, image_to_video=False, device=None):
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.wan.model.WanModel) super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.wan.model.WanModel)
@ -1077,10 +1186,10 @@ class WAN21(BaseModel):
image = torch.zeros(shape_image, dtype=noise.dtype, layout=noise.layout, device=noise.device) image = torch.zeros(shape_image, dtype=noise.dtype, layout=noise.layout, device=noise.device)
else: else:
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], 36): for i in range(0, image.shape[1], 16):
# image[:, i: i + 36] = self.process_latent_in(image[:, i: i + 36]) image[:, i: i + 16] = self.process_latent_in(image[:, i: i + 36])
# 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}")
if not self.image_to_video or extra_channels == image.shape[1]: if not self.image_to_video or extra_channels == image.shape[1]:
@ -1090,22 +1199,28 @@ 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}")
return torch.cat((mask, image), dim=1) res = torch.cat((mask, image), dim=1)
tracks = kwargs.get("tracks", None)
if tracks is not None:
res = patch_motion(tracks, res, 220.0, (4, 16), 2)[None]
return res
def extra_conds(self, **kwargs): def extra_conds(self, **kwargs):
out = super().extra_conds(**kwargs) out = super().extra_conds(**kwargs)

View File

@ -667,23 +667,27 @@ class WanTrackToVideo:
], ],
dim=1).to(start_image.device) dim=1).to(start_image.device)
res = res.permute(1,2,3,0)[:, :, :, :3] # T, H, W, C
print("start vid:", res.shape, res) print("start vid:", res.shape, res)
res = res.permute(1,2,3,0)[:, :, :, :3] # T, H, W, C
y = vae.encode( y = vae.encode(
res res
)[0] )[0]
print("mask shape:", msk.shape, "y shape:", y.shape) print("mask shape:", msk.shape, "y shape:", y.shape)
y = torch.concat([msk, y]) # y = torch.concat([msk, y])
motion_patched = patch_motion(processed_tracks, y, temperature, (4, 16), topk)[None]
mask, video = motion_patched[:, 0:4], motion_patched[:, 4:] # mask, video = motion_patched[:, 0:4], motion_patched[:, 4:]
# Add motion features to conditioning # Add motion features to conditioning
concat_latent_image = vae.encode(image[:, :, :, :3]) concat_latent_image = vae.encode(image[:, :, :, :3])
positive = node_helpers.conditioning_set_values(positive, positive = node_helpers.conditioning_set_values(positive,
{"concat_mask": mask, {"tracks": processed_tracks,
"concat_latent_image": video}) "concat_mask": msk,
"concat_latent_image": y})
negative = node_helpers.conditioning_set_values(negative, negative = node_helpers.conditioning_set_values(negative,
{"concat_mask": mask, {"tracks": processed_tracks,
"concat_latent_image": video}) "concat_mask": msk,
"concat_latent_image": y})
# Handle clip vision output if provided # Handle clip vision output if provided