This commit is contained in:
kijai
2025-05-31 01:58:21 +03:00
parent cd2884d88a
commit 8bc74daad9
4 changed files with 34 additions and 62 deletions
+1 -35
View File
@@ -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))
-14
View File
@@ -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)
+22 -11
View File
@@ -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 = {
+11 -2
View File
@@ -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