diff --git a/comfy/model_base.py b/comfy/model_base.py index 716d6d34a..4392355ea 100644 --- a/comfy/model_base.py +++ b/comfy/model_base.py @@ -1073,111 +1073,6 @@ class Lumina2(BaseModel): out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn) 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""" - _, 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 - 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): 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) @@ -1222,14 +1117,7 @@ class WAN21(BaseModel): mask = mask.repeat(1, 4, 1, 1, 1) mask = utils.resize_to_batch_size(mask, noise.shape[0]) - res = torch.cat((mask, image), dim=1) - - tracks = kwargs.get("tracks", None) - if tracks is not None: - res = patch_motion(tracks.to(device), res[0], kwargs.get("ati_temperature", None), (4, 16), kwargs.get("ati_topk", None))[None] - - return res - + return torch.cat((mask, image), dim=1) def extra_conds(self, **kwargs): out = super().extra_conds(**kwargs) diff --git a/comfy_extras/nodes_wan.py b/comfy_extras/nodes_wan.py index 2722d838b..38ca17c63 100644 --- a/comfy_extras/nodes_wan.py +++ b/comfy_extras/nodes_wan.py @@ -491,6 +491,142 @@ def pad_pts(tr): pts = pts[:FIXED_LENGTH] return pts.reshape(FIXED_LENGTH, 1, 3) +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_single( + tracks: torch.FloatTensor, # (B, T, N, 4) + vid: torch.FloatTensor, # (C, T, H, W) + temperature: float, + vae_divide: tuple, + topk: int, +): + """Apply motion patching based on tracks""" + _, 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 + out_mask_full = torch.cat([torch.ones_like(out_weight[:1]), out_weight], dim=0) # T, H, W + + return out_mask_full[None].expand(vae_divide[0], -1, -1, -1), out_feature_full + + +def patch_motion( + tracks: torch.FloatTensor, # (B, TB, T, N, 4) + vid: torch.FloatTensor, # (C, T, H, W) + temperature: float = 220.0, + vae_divide: tuple = (4, 16), + topk: int = 2, +): + B = len(tracks) + + # Process each batch separately + out_masks = [] + out_features = [] + + for b in range(B): + mask, feature = _patch_motion_single( + tracks[b], # (T, N, 4) + vid, # (C, T, H, W) + temperature, + vae_divide, + topk + ) + out_masks.append(mask) + out_features.append(feature) + + # Stack results: (B, C, T, H, W) + out_mask_full = torch.stack(out_masks, dim=0) + out_feature_full = torch.stack(out_features, dim=0) + + return out_mask_full, out_feature_full + class WanTrackToVideo: @classmethod def INPUT_TYPES(s): @@ -529,14 +665,19 @@ class WanTrackToVideo: latent = torch.zeros([batch_size, 16, ((length - 1) // 4) + 1, height // 8, width // 8], device=comfy.model_management.intermediate_device()) # Convert tracks to tensor format - arrs = [] - for i, track in enumerate(tracks_data): - pts = pad_pts(track) - arrs.append(pts) + if isinstance(tracks_data[0][0], dict): + tracks_data = [tracks_data] - tracks_np = np.stack(arrs, axis=0) - processed_tracks = process_tracks(tracks_np, (width, height), length - 1).unsqueeze(0) + processed_tracks = [] + for batch in tracks_data: + arrs = [] + for track in batch: + pts = pad_pts(track) + arrs.append(pts) + tracks_np = np.stack(arrs, axis=0) + processed_tracks.append(process_tracks(tracks_np, (width, height), length - 1).unsqueeze(0)) + if start_image is not None: start_image = comfy.utils.common_upscale(start_image[:length].movedim(-1, 1), width, height, "bilinear", "center").movedim(1, -1) @@ -570,8 +711,13 @@ class WanTrackToVideo: res = res.permute(1,2,3,0)[:, :, :, :3] # T, H, W, C y = vae.encode(res) + c = torch.cat((msk, y), dim=1) + res = patch_motion( + processed_tracks, c[0], temperature=temperature, topk=topk, vae_divide=(4, 16) + ) - # Add motion features to conditioning + msk, y = res + msk = -msk + 1.0 # Invert mask to match expected format positive = node_helpers.conditioning_set_values(positive, {"tracks": processed_tracks, "concat_mask": msk,