From d000fbc645b6cb844de3b87f98292b326f31654f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 9 Dec 2025 19:55:56 +0200 Subject: [PATCH] Support WanMove https://github.com/ali-vilab/Wan-Move --- WanMove/example_tracks.npy | Bin 0 -> 776 bytes WanMove/example_visibility.npy | Bin 0 -> 209 bytes WanMove/nodes.py | 134 ++++++++++++++ WanMove/trajectory.py | 324 +++++++++++++++++++++++++++++++++ nodes_sampler.py | 14 +- 5 files changed, 470 insertions(+), 2 deletions(-) create mode 100644 WanMove/example_tracks.npy create mode 100644 WanMove/example_visibility.npy create mode 100644 WanMove/nodes.py create mode 100644 WanMove/trajectory.py diff --git a/WanMove/example_tracks.npy b/WanMove/example_tracks.npy new file mode 100644 index 0000000000000000000000000000000000000000..f9759ebc83b31d5d5b89244a5facc4badfe2a393 GIT binary patch literal 776 zcmb7&{ZGsR0LJw)WNWCI7UnQ6xm!Y(_pQ&ckhiAo)VJ=szSosHx~oLpeG_?4D9u|K zDiyUNA}_mX>5b8v_q?8@mh?kpBIi%=?0NRt^T}DgA!x12VMy|jRC$zW(i-LSW%7l( zIdX+euG?%(G-?bHn~hQ8Kfg>9XAXqK$54I7KsF zvvO`YKJ719;1I!=%IC;jB2cb<%BaT?1Xj1==&ojZKnt-=YEHL4CSbCLAA296FVG-4 zHSC8|581<=~|ZPAJBP5%SKx*#(pjJjrSQlpygD$6~`A*tP$?v>?x9^ zzKgF(q^q}v!VVFA%Po4piZql}B(bM96neS>nvtvuBnHCgN9|Vq;$- zyE9KPbC0fF9dj_uriK+=ZXiAXqE>9z4 zg+$ND6sqP)T(3wX*-JwBZJ@5yLhkiVR4leItXd@Oi5b_4kt8IVQ4J5HbCQ{)dm%J? Gn)wUE%s<@# literal 0 HcmV?d00001 diff --git a/WanMove/example_visibility.npy b/WanMove/example_visibility.npy new file mode 100644 index 0000000000000000000000000000000000000000..70e1137fc7abb1b19f188a96b840a17ea1b75f33 GIT binary patch literal 209 zcmbR27wQ`j$;eQ~P_3SlTAW;@Zl$1JlVqr_qoAIaUsO_*m=~X4l#&V(cT3DEP6dh= dXCxM+0{I$-Itms*Y^bTDP^&-|;9{gU0031T94G() literal 0 HcmV?d00001 diff --git a/WanMove/nodes.py b/WanMove/nodes.py new file mode 100644 index 0000000..985273c --- /dev/null +++ b/WanMove/nodes.py @@ -0,0 +1,134 @@ +import json +import torch +import torchvision.transforms.functional as TF +from ..utils import log +from .trajectory import create_pos_feature_map, draw_tracks_on_video +import os +from comfy import model_management as mm +device = mm.get_torch_device() +script_directory = os.path.dirname(os.path.abspath(__file__)) + +VAE_STRIDE = (4, 8, 8) # t, h, w + +class WanVideoWanDrawWanMoveTracks: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "tracks": ("WANMOVETRACKS",), + "width": ("INT", {"default": 832, "min": 64, "max": 2048, "step": 8, "tooltip": "Width of the image to encode"}), + "height": ("INT", {"default": 480, "min": 64, "max": 29048, "step": 8, "tooltip": "Height of the image to encode"}), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("image",) + FUNCTION = "execute" + CATEGORY = "WanVideoWrapper" + + def execute(self, tracks, width, height): + track = tracks["tracks"].unsqueeze(0) + track_visibility = tracks["track_visibility"].unsqueeze(0) + + video = torch.zeros((track.shape[1], height, width, 3), device=device) + track_video = draw_tracks_on_video(video, track, track_visibility) + track_video = torch.stack([TF.to_tensor(frame) for frame in track_video], dim=0).movedim(1, -1) + + return (track_video.float().cpu(), ) + + +class WanVideoAddWanMoveTracks: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image_embeds": ("WANVIDIMAGE_EMBEDS",), + "track_coords": ("STRING", {"forceInput": True, "tooltip": "JSON string or list of JSON strings representing the tracks"}), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}), + }, + "optional": { + "track_mask": ("MASK",), + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "WANMOVETRACKS") + RETURN_NAMES = ("image_embeds", "tracks") + FUNCTION = "add" + CATEGORY = "WanVideoWrapper" + + def add(self, image_embeds, track_coords, strength, track_mask=None): + updated = dict(image_embeds) + + target_shape = image_embeds.get("target_shape") + if target_shape is not None: + height = target_shape[2] * VAE_STRIDE[1] + width = target_shape[3] * VAE_STRIDE[2] + else: + height = image_embeds["lat_h"] * VAE_STRIDE[1] + width = image_embeds["lat_w"] * VAE_STRIDE[2] + num_frames = image_embeds["num_frames"] + + tracks_data = parse_json_tracks(track_coords) + track_list = [ + [[track[frame]['x'], track[frame]['y']] for track in tracks_data] + for frame in range(len(tracks_data[0])) + ] + track = torch.tensor(track_list, dtype=torch.float32, device=device) # shape: (frames, num_tracks, 2) + track = track[:num_frames] + + num_tracks = track.shape[-2] + if track_mask is None: + track_visibility = torch.ones((num_frames, num_tracks), dtype=torch.bool, device=device) + else: + track_visibility = (track_mask > 0).any(dim=(1, 2)).unsqueeze(-1) + + feature_map, track_pos = create_pos_feature_map(track, track_visibility, VAE_STRIDE, height, width, 16, track_num=1, device=device) + + updated.setdefault("wanmove_embeds", {}) + updated["wanmove_embeds"]["track_pos"] = track_pos * strength + + tracks_dict = { + "tracks": track, + "track_visibility": track_visibility, + } + + return (updated, tracks_dict,) + + +def parse_json_tracks(tracks): + tracks_data = [] + try: + # If tracks is a string, try to parse it as JSON + if isinstance(tracks, str): + parsed = json.loads(tracks.replace("'", '"')) + tracks_data.extend(parsed) + else: + # If tracks is a list of strings, parse each one + for track_str in tracks: + parsed = json.loads(track_str.replace("'", '"')) + tracks_data.append(parsed) + + # Check if we have a single track (dict with x,y) or a list of tracks + if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]: + # Single track detected, wrap it in a list + tracks_data = [tracks_data] + elif tracks_data and isinstance(tracks_data[0], list) and tracks_data[0] and isinstance(tracks_data[0][0], dict) and 'x' in tracks_data[0][0]: + # Already a list of tracks, nothing to do + pass + else: + # Unexpected format + log.warning(f"Warning: Unexpected track format: {type(tracks_data[0])}") + + except json.JSONDecodeError as e: + log.warning(f"Error parsing tracks JSON: {e}") + tracks_data = [] + + return tracks_data + + +NODE_CLASS_MAPPINGS = { + "WanVideoAddWanMoveTracks": WanVideoAddWanMoveTracks, + "WanVideoWanDrawWanMoveTracks": WanVideoWanDrawWanMoveTracks, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "WanVideoAddWanMoveTracks": "WanVideo Add WanMove Tracks", + "WanVideoWanDrawWanMoveTracks": "WanVideo Draw WanMove Tracks", + } diff --git a/WanMove/trajectory.py b/WanMove/trajectory.py new file mode 100644 index 0000000..2fd656c --- /dev/null +++ b/WanMove/trajectory.py @@ -0,0 +1,324 @@ +# https://github.com/ali-vilab/Wan-Move/blob/main/wan/modules/trajectory.py + +import numpy as np +import torch +from PIL import Image, ImageDraw + +SKIP_ZERO = False + +def get_pos_emb( + pos_k: torch.Tensor, + pos_emb_dim: int, + theta_func: callable = lambda i, d: torch.pow(10000, torch.mul(2, torch.div(i.to(torch.float32), d))), + device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"), + dtype: torch.dtype = torch.float32, +) -> torch.Tensor: + """ + Generate batch position embeddings. + + Args: + pos_k (torch.Tensor): A 1D tensor containing positions for which to generate embeddings. + pos_emb_dim (int): The dimension of position embeddings. + theta_func (callable): Function to compute thetas based on position and embedding dimensions. + device (torch.device): Device to store the position embeddings. + dtype (torch.dtype): Desired data type for computations. + + Returns: + torch.Tensor: The position embeddings with shape (batch_size, pos_emb_dim). + """ + assert pos_emb_dim % 2 == 0, "The dimension of position embeddings must be even." + pos_k = pos_k.to(device, dtype) + if SKIP_ZERO: + pos_k = pos_k + 1 + batch_size = pos_k.size(0) + + denominator = torch.arange(0, pos_emb_dim // 2, device=device, dtype=dtype) + # Expand denominator to match the shape needed for broadcasting + denominator_expanded = denominator.view(1, -1).expand(batch_size, -1) + + thetas = theta_func(denominator_expanded, pos_emb_dim) + + # Ensure pos_k is in the correct shape for broadcasting + pos_k_expanded = pos_k.view(-1, 1).to(dtype) + sin_thetas = torch.sin(torch.div(pos_k_expanded, thetas)) + cos_thetas = torch.cos(torch.div(pos_k_expanded, thetas)) + + # Concatenate sine and cosine embeddings along the last dimension + pos_emb = torch.cat([sin_thetas, cos_thetas], dim=-1) + + return pos_emb + +def create_pos_feature_map( + pred_tracks: torch.Tensor, # [T, N, 2] + pred_visibility: torch.Tensor, # [T, N] + downsample_ratios: list[int], + height: int, + width: int, + pos_emb_dim: int, + track_num: int = -1, + t_down_strategy: str = "sample", + device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"), + dtype: torch.dtype = torch.float32, +): + """ + Create a feature map from the predicted tracks. + + Args: + - pred_tracks: torch.Tensor, the predicted tracks, [T, N, 2] + - pred_visibility: torch.Tensor, the predicted visibility, [T, N] + - downsample_ratios: list[int], the ratios for downsampling time, height, and width + - height: int, the height of the feature map + - width: int, the width of the feature map + - pos_emb_dim: int, the dimension of the position embeddings + - track_num: int, the number of tracks to use + - t_down_strategy: str, the strategy for downsampling time dimension + - device: torch.device, the device + - dtype: torch.dtype, the data type + + Returns: + - feature_map: torch.Tensor, the feature map, [T', H', W', pos_emb_dim] + - track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width + """ + + assert t_down_strategy in ["sample", "average"], "Invalid strategy for downsampling time dimension." + + t, n, _ = pred_tracks.shape + t_down, h_down, w_down = downsample_ratios + feature_map = torch.zeros((t-1) // t_down + 1, height // h_down, width // w_down, pos_emb_dim, device=device, dtype=dtype) + track_pos = - torch.ones(n, (t-1) // t_down + 1, 2, dtype=torch.long) + + if track_num == -1: + track_num = n + + tracks_idx = torch.randperm(n)[:track_num] + tracks = pred_tracks[:, tracks_idx] + visibility = pred_visibility[:, tracks_idx] + tracks_embs = get_pos_emb(torch.randperm(n)[:track_num], pos_emb_dim, device=device, dtype=dtype) + + for t_idx in range(0, t, t_down): + if t_down_strategy == "sample" or t_idx == 0: + cur_tracks = tracks[t_idx] # [N, 2] + cur_visibility = visibility[t_idx] # [N] + else: + cur_tracks = tracks[t_idx:t_idx+t_down].mean(dim=0) + cur_visibility = torch.any(visibility[t_idx:t_idx+t_down], dim=0) + + for i in range(track_num): + if not cur_visibility[i] or cur_tracks[i][0] < 0 or cur_tracks[i][1] < 0 or cur_tracks[i][0] >= width or cur_tracks[i][1] >= height: + continue + x, y = cur_tracks[i] + x, y = int(x // w_down), int(y // h_down) + feature_map[t_idx // t_down, y, x] += tracks_embs[i] + track_pos[i, t_idx // t_down, 0], track_pos[i, t_idx // t_down, 1] = y, x + + return feature_map, track_pos + + +def replace_feature( + vae_feature: torch.Tensor, # [B, C', T', H', W'] + track_pos: torch.Tensor, # [B, N, T', 2] +) -> torch.Tensor: + b, _, t, h, w = vae_feature.shape + assert b == track_pos.shape[0], "Batch size mismatch." + n = track_pos.shape[1] + + # Shuffle the trajectory order + track_pos = track_pos[:, torch.randperm(n), :, :] + + # Extract coordinates at time steps ≥ 1 and generate a valid mask + current_pos = track_pos[:, :, 1:, :] # [B, N, T-1, 2] + mask = (current_pos[..., 0] >= 0) & (current_pos[..., 1] >= 0) # [B, N, T-1] + + # Get all valid indices + valid_indices = mask.nonzero(as_tuple=False) # [num_valid, 3] + num_valid = valid_indices.shape[0] + + if num_valid == 0: + return vae_feature + + # Decompose valid indices into each dimension + batch_idx = valid_indices[:, 0] + track_idx = valid_indices[:, 1] + t_rel = valid_indices[:, 2] + t_target = t_rel + 1 # Convert to original time step indices + + # Extract target position coordinates + h_target = current_pos[batch_idx, track_idx, t_rel, 0].long() # Ensure integer indices + w_target = current_pos[batch_idx, track_idx, t_rel, 1].long() + + # Extract source position coordinates (t=0) + h_source = track_pos[batch_idx, track_idx, 0, 0].long() + w_source = track_pos[batch_idx, track_idx, 0, 1].long() + + # Get source features and assign to target positions + src_features = vae_feature[batch_idx, :, 0, h_source, w_source] + vae_feature[batch_idx, :, t_target, h_target, w_target] = src_features + + return vae_feature + +def get_video_track_video( + model, + video_tensor: torch.Tensor, # [T, C, H, W] + downsample_ratios: list[int], + pos_emb_dim: int, + grid_size: int = 32, + track_num: int = -1, + t_down_strategy: str = "sample", + device: torch.device = torch.device("cuda" if torch.cuda.is_available() else "cpu"), + dtype: torch.dtype = torch.float32, +) -> tuple[torch.Tensor, torch.Tensor]: + """ + Get the track video from the video tensor. + + Args: + - model: torch.nn.Module, the model for tracking, CoTracker + - video_tensor: torch.Tensor, the video tensor, [T, C, H, W] + - downsample_ratios: list[int], the ratios for downsampling time, height, and width + - height: int, the height of the feature map + - width: int, the width of the feature map + - pos_emb_dim: int, the dimension of the position embeddings + - grid_size: int, the size of the grid + - track_num: int, the number of tracks to use + - t_down_strategy: str, the strategy for downsampling time dimension + - device: torch.device, the device + - dtype: torch.dtype, the data type + + Returns: + - track_video: torch.Tensor, the track video, [pos_emb_dim, T', H', W'] + - track_pos: torch.Tensor, the position embeddings, [N, T', 2], 2 = height, width + - pred_tracks: the predicted point trajectories + - pred_visibility: visibility of the predicted point trajectories + """ + + t, c, height, width = video_tensor.shape + with ( + torch.autocast(device_type=device.type, dtype=dtype), + torch.no_grad(), + ): + pred_tracks, pred_visibility = model( + video_tensor.unsqueeze(0), + grid_size=grid_size, + backward_tracking=False, + ) + + track_video, track_pos = create_pos_feature_map( + pred_tracks[0], pred_visibility[0], downsample_ratios, height, width, pos_emb_dim, track_num, t_down_strategy, device, dtype + ) + + return track_video.permute(3, 0, 1, 2), track_pos, pred_tracks, pred_visibility + +# --------------------------- +# Visualize functions +# -------------------------- + +def draw_overall_gradient_polyline_on_image(image, line_width, points, start_color): + """ + - image (Image): target image to draw on. + - line_width (int): initial line width. + - points (list of tuples): list of points forming the polyline, each point is (x, y). + - start_color (tuple): starting color of the line (R, G, B). + + Return: + - Image: original image with the gradient polyline drawn. + """ + + def get_distance(p1, p2): + return ((p2[0] - p1[0]) ** 2 + (p2[1] - p1[1]) ** 2) ** 0.5 + + # Create a new image with the same size as the original + new_image = Image.new('RGBA', image.size) + draw = ImageDraw.Draw(new_image, 'RGBA') + points = points[::-1] + + # Compute total length + total_length = sum(get_distance(points[i], points[i+1]) for i in range(len(points)-1)) + + # Accumulated length + accumulated_length = 0 + + # Draw the gradient polyline + for start_point, end_point in zip(points[:-1], points[1:]): + segment_length = get_distance(start_point, end_point) + steps = int(segment_length) + + for i in range(steps): + # Current accumulated length + current_length = accumulated_length + (i / steps) * segment_length + + # Alpha from fully opaque to fully transparent + alpha = int(255 * (1 - current_length / total_length)) + color = (*start_color, alpha) + + # Interpolated coordinates + x = int(start_point[0] + (end_point[0] - start_point[0]) * i / steps) + y = int(start_point[1] + (end_point[1] - start_point[1]) * i / steps) + + # Dynamic line width, decreasing from initial width to 1 + dynamic_line_width = int(line_width * (1 - (current_length / total_length))) + dynamic_line_width = max(dynamic_line_width, 1) # minimum width is 1 to avoid 0 + + draw.line([(x, y), (x + 1, y)], fill=color, width=dynamic_line_width) + + accumulated_length += segment_length + + return new_image + +def add_weighted(rgb, track): + rgb = np.array(rgb) # [H, W, C] "RGB" + track = np.array(track) # [H, W, C] "RGBA" + + # Compute weights from the alpha channel + alpha = track[:, :, 3] / 255.0 + + # Expand alpha to 3 channels to match RGB + alpha = np.stack([alpha] * 3, axis=-1) + + # Blend the two images + blend_img = track[:, :, :3] * alpha + rgb * (1 - alpha) + + return Image.fromarray(blend_img.astype(np.uint8)) + +def draw_tracks_on_video(video, tracks, visibility=None, track_frame=24): + color_map = [ + (102, 153, 255), + (0, 255, 255), + (255, 255, 0), + (255, 102, 204), + (0, 255, 0) + ] + circle_size = 12 + line_width = 16 + + video = video.byte().cpu().numpy() # (81, 480, 832, 3) + tracks = tracks[0].long().detach().cpu().numpy() + if visibility is not None: + visibility = visibility[0].detach().cpu().numpy() + # print(video.shape, tracks.shape) + + output_frames = [] + # Process the video + for t in range(video.shape[0]): + # Extract current frame + frame = video[t] + frame = Image.fromarray(frame).convert("RGB") + + # Draw tracks + for n in range(tracks.shape[1]): + if visibility is not None and visibility[t, n] == 0: + continue + + # Track coordinate at current frame + track_coord = tracks[t, n] + tracks_coord = tracks[max(t-track_frame, 0):t+1, n] + + # Draw a circle + draw = ImageDraw.Draw(frame) + draw.ellipse((track_coord[0] - circle_size, track_coord[1] - circle_size, track_coord[0] + circle_size, track_coord[1] + circle_size), fill=color_map[n % len(color_map)]) + # Draw the polyline + track_image = draw_overall_gradient_polyline_on_image(frame, line_width, tracks_coord, color_map[n % len(color_map)]) + frame = add_weighted(frame, track_image) + + # Save current frame + output_frames.append(frame.convert("RGB")) + + return output_frames diff --git a/nodes_sampler.py b/nodes_sampler.py index 5439ac3..8931380 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -291,7 +291,7 @@ class WanVideoSampler: else: cfg = [cfg] * (steps + 1) - control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = mocha_embeds = None + control_latents = control_camera_latents = clip_fea = clip_fea_neg = end_image = recammaster = camera_embed = unianim_data = mocha_embeds = image_cond_neg =None vace_data = vace_context = vace_scale = None fun_or_fl2v_model = has_ref = drop_last = False phantom_latents = fun_ref_image = ATI_tracks = None @@ -301,6 +301,15 @@ class WanVideoSampler: #I2V image_cond = image_embeds.get("image_embeds", None) if image_cond is not None: + # WanMove + wanmove_embeds = image_embeds.get("wanmove_embeds", None) + if wanmove_embeds is not None: + from .WanMove.trajectory import replace_feature + track_pos = wanmove_embeds["track_pos"] + if any(not math.isclose(c, 1.0) for c in cfg): + image_cond_neg = torch.cat([image_embeds["mask"], image_cond]) + image_cond = replace_feature(image_cond.unsqueeze(0), track_pos.unsqueeze(0))[0] + if transformer.in_dim == 16: raise ValueError("T2V (text to video) model detected, encoded images only work with I2V (Image to video) models") elif transformer.in_dim not in [48, 32]: # fun 2.1 models don't use the mask @@ -1118,7 +1127,7 @@ class WanVideoSampler: lynx_embeds=lynx_embeds ) log.info(f"Extracted {len(lynx_ref_buffer)} cond ref buffers") - if not math.isclose(cfg[0], 1.0): + if any(not math.isclose(c, 1.0) for c in cfg): log.info("Extracting Lynx ref uncond buffer...") if transformer.in_dim == 36: lynx_ref_input_uncond = torch.cat([lynx_ref_latent_uncond, empty_image_cond], dim=0) @@ -1536,6 +1545,7 @@ class WanVideoSampler: base_params['is_uncond'] = True base_params['clip_fea'] = clip_fea_neg if clip_fea_neg is not None else clip_fea base_params["add_text_emb"] = qwenvl_embeds_neg.to(device) if qwenvl_embeds_neg is not None else None # QwenVL embeddings for Bindweave + base_params['y'] = image_cond_neg if image_cond_neg is not None else base_params['y'] if wananim_face_pixels is not None: base_params['wananim_face_pixel_values'] = torch.zeros_like(wananim_face_pixels).to(device, torch.float32) - 1 if humo_audio_input_neg is not None: