cleanup
This commit is contained in:
+1
-35
@@ -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))
|
||||
|
||||
@@ -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
@@ -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 = {
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user