From 8bc74daad9e58e985db76304050ec74b7b78361f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 31 May 2025 01:58:21 +0300 Subject: [PATCH] cleanup --- ATI/motion.py | 36 +----------------------------------- ATI/motion_patch.py | 14 -------------- ATI/nodes.py | 33 ++++++++++++++++++++++----------- nodes.py | 13 +++++++++++-- 4 files changed, 34 insertions(+), 62 deletions(-) diff --git a/ATI/motion.py b/ATI/motion.py index 3861d18..3d615e6 100644 --- a/ATI/motion.py +++ b/ATI/motion.py @@ -12,49 +12,15 @@ # See the License for the specific language governing permissions and # limitations under the License. -import io from typing import Dict, List, Optional, Tuple, Union import numpy as np import torch - -def get_tracks_inference(tracks, height, width, quant_multi: Optional[int] = 8, **kwargs): - if isinstance(tracks, str): - tracks = torch.load(tracks) - - tracks_np = unzip_to_array(tracks) - print("tracks_np shape: ", tracks_np.shape) - print(tracks_np) - - tracks = process_tracks( - tracks_np, (width, height), quant_multi=1, **kwargs - ) - - return tracks - - -def unzip_to_array( - data: bytes, key: Union[str, List[str]] = "array" -) -> Union[np.ndarray, Dict[str, np.ndarray]]: - bytes_io = io.BytesIO(data) - - if isinstance(key, str): - # Load the NPZ data from the BytesIO object - with np.load(bytes_io) as data: - return data[key] - else: - get = {} - with np.load(bytes_io) as data: - for k in key: - get[k] = data[k] - return get - - def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], quant_multi: int = 8, **kwargs): # tracks: shape [t, h, w, 3] => samples align with 24 fps, model trained with 16 fps. # frame_size: tuple (W, H) - tracks = torch.from_numpy(tracks_np).float()# / quant_multi + tracks = torch.from_numpy(tracks_np).float() if tracks.shape[1] == 121: tracks = torch.permute(tracks, (1, 0, 2, 3)) diff --git a/ATI/motion_patch.py b/ATI/motion_patch.py index a428753..bde62b0 100644 --- a/ATI/motion_patch.py +++ b/ATI/motion_patch.py @@ -78,8 +78,6 @@ def patch_motion( tracks: torch.FloatTensor, # (B, T, N, 4) vid: torch.FloatTensor, # (C, T, H, W) temperature: float = 220.0, - training: bool = True, - tail_dropout: float = 0.2, vae_divide: tuple = (4, 16), topk: int = 2, ): @@ -93,18 +91,6 @@ def patch_motion( tracks_n = tracks_n.clamp(-1, 1) visible = visible.clamp(0, 1) - if tail_dropout > 0 and training: - TT = visible.shape[1] - rrange = torch.arange(TT, device=visible.device, dtype=visible.dtype)[ - None, :, None, None - ] - rand_nn = torch.rand_like(visible[:, :1]) - rand_rr = torch.rand_like(visible[:, :1]) * (TT - 1) - visible = visible * ( - (rand_nn > tail_dropout).type_as(visible) - + (rrange < rand_rr).type_as(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) diff --git a/ATI/nodes.py b/ATI/nodes.py index e57be7b..b19c7bd 100644 --- a/ATI/nodes.py +++ b/ATI/nodes.py @@ -1,9 +1,6 @@ -import os, io import json -import torch -script_directory = os.path.dirname(os.path.abspath(__file__)) -from .motion import get_tracks_inference, process_tracks +from .motion import process_tracks import numpy as np FIXED_LENGTH = 121 def pad_pts(tr): @@ -25,6 +22,10 @@ class WanVideoATITracks: "tracks": ("STRING",), "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"}), + "temperature": ("FLOAT", {"default": 220.0, "min": 0.0, "max": 1000.0, "step": 0.1}), + "topk": ("INT", {"default": 2, "min": 1, "max": 10, "step": 1}), + "start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percent of the steps to apply ATI"}), + "end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percent of the steps to apply ATI"}), }, } @@ -33,15 +34,21 @@ class WanVideoATITracks: FUNCTION = "patchmodel" CATEGORY = "WanVideoWrapper" - def patchmodel(self, model, tracks, width, height): - tracks_data = json.loads(tracks) + def patchmodel(self, model, tracks, width, height, temperature, topk, start_percent, end_percent): + + if len(tracks) < 10: + tracks_data = [] + for coords in tracks: + coords = json.loads(coords.replace("'", '"')) + tracks_data.append(coords) + else: + coords = json.loads(tracks.replace("'", '"')) - if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]: - # It's a single track, wrap it in a list to make it a list of tracks - tracks_data = [tracks_data] + if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]: + # It's a single track, wrap it in a list to make it a list of tracks + tracks_data = [tracks_data] arrs = [] - for track in tracks_data: pts = pad_pts(track) arrs.append(pts) @@ -52,7 +59,11 @@ class WanVideoATITracks: patcher = model.clone() patcher.model_options["transformer_options"]["ati_tracks"] = processed_tracks.unsqueeze(0) - + patcher.model_options["transformer_options"]["ati_temperature"] = temperature + patcher.model_options["transformer_options"]["ati_topk"] = topk + patcher.model_options["transformer_options"]["ati_start_percent"] = start_percent + patcher.model_options["transformer_options"]["ati_end_percent"] = end_percent + return (patcher,) NODE_CLASS_MAPPINGS = { diff --git a/nodes.py b/nodes.py index 7146203..147c9a3 100644 --- a/nodes.py +++ b/nodes.py @@ -2415,6 +2415,7 @@ class WanVideoSampler: fun_ref_image = None image_cond = image_embeds.get("image_embeds", None) + ATI_tracks = None if image_cond is not None: log.info(f"image_cond shape: {image_cond.shape}") @@ -2423,7 +2424,11 @@ class WanVideoSampler: ATI_tracks = transformer_options.get("ati_tracks", None) if ATI_tracks is not None: from .ATI.motion_patch import patch_motion - image_cond = patch_motion(ATI_tracks.to(image_cond.device, image_cond.dtype), image_cond, training=False) + topk = transformer_options.get("ati_topk", 2) + temperature = transformer_options.get("ati_temperature", 220.0) + ati_start_percent = transformer_options.get("ati_start_percentage", 0.0) + ati_end_percent = transformer_options.get("ati_end_percentage", 1.0) + image_cond_ati = patch_motion(ATI_tracks.to(image_cond.device, image_cond.dtype), image_cond, topk=topk, temperature=temperature) log.info(f"ATI tracks shape: {ATI_tracks.shape}") end_image = image_embeds.get("end_image", None) @@ -2906,7 +2911,11 @@ class WanVideoSampler: if not patcher.model.is_patched: log.info("Loading LoRA...") patcher = apply_lora(patcher, device, device, low_mem_load=False) - patcher.model.is_patched = True + patcher.model.is_patched = True + elif ATI_tracks is not None: + if (ati_start_percent <= current_step_percentage <= ati_end_percent) or \ + (ati_end_percent > 0 and idx == 0 and current_step_percentage >= ati_start_percent): + image_cond_input = image_cond_ati.to(z) else: image_cond_input = image_cond.to(z) if image_cond is not None else None