mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-02 15:47:07 +08:00
Ruff fix
This commit is contained in:
parent
acd91d30a7
commit
932002dcf2
@ -8,7 +8,7 @@ import comfy.latent_formats
|
|||||||
import comfy.clip_vision
|
import comfy.clip_vision
|
||||||
import json
|
import json
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from typing import List, Optional, Tuple, Union
|
from typing import Tuple
|
||||||
|
|
||||||
class WanImageToVideo:
|
class WanImageToVideo:
|
||||||
@classmethod
|
@classmethod
|
||||||
@ -399,7 +399,7 @@ def parse_json_tracks(tracks):
|
|||||||
for track_str in tracks:
|
for track_str in tracks:
|
||||||
parsed = json.loads(track_str.replace("'", '"'))
|
parsed = json.loads(track_str.replace("'", '"'))
|
||||||
tracks_data.append(parsed)
|
tracks_data.append(parsed)
|
||||||
|
|
||||||
# Check if we have a single track (dict with x,y) or a list of tracks
|
# 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]:
|
if tracks_data and isinstance(tracks_data[0], dict) and 'x' in tracks_data[0]:
|
||||||
# Single track detected, wrap it in a list
|
# Single track detected, wrap it in a list
|
||||||
@ -409,10 +409,9 @@ def parse_json_tracks(tracks):
|
|||||||
pass
|
pass
|
||||||
else:
|
else:
|
||||||
# Unexpected format
|
# Unexpected format
|
||||||
print(f"Warning: Unexpected track format: {type(tracks_data[0])}")
|
pass
|
||||||
|
|
||||||
except json.JSONDecodeError as e:
|
except json.JSONDecodeError:
|
||||||
print(f"Error parsing tracks JSON: {e}")
|
|
||||||
tracks_data = []
|
tracks_data = []
|
||||||
return tracks_data
|
return tracks_data
|
||||||
|
|
||||||
@ -421,10 +420,10 @@ def tracks_to_tensor(tracks_data, length, width, height, batch_size=1):
|
|||||||
if not tracks_data:
|
if not tracks_data:
|
||||||
# Return empty tracks if no data
|
# Return empty tracks if no data
|
||||||
return torch.zeros((batch_size, length, 1, 4))
|
return torch.zeros((batch_size, length, 1, 4))
|
||||||
|
|
||||||
num_tracks = len(tracks_data)
|
num_tracks = len(tracks_data)
|
||||||
tracks_tensor = torch.zeros((batch_size, length, num_tracks, 4))
|
tracks_tensor = torch.zeros((batch_size, length, num_tracks, 4))
|
||||||
|
|
||||||
for batch_idx in range(batch_size):
|
for batch_idx in range(batch_size):
|
||||||
for track_idx, track in enumerate(tracks_data):
|
for track_idx, track in enumerate(tracks_data):
|
||||||
for frame_idx in range(min(length, len(track))):
|
for frame_idx in range(min(length, len(track))):
|
||||||
@ -436,31 +435,31 @@ def tracks_to_tensor(tracks_data, length, width, height, batch_size=1):
|
|||||||
x_norm = (x / width) * 2 - 1
|
x_norm = (x / width) * 2 - 1
|
||||||
y_norm = (y / height) * 2 - 1
|
y_norm = (y / height) * 2 - 1
|
||||||
visible = point.get('visible', 1)
|
visible = point.get('visible', 1)
|
||||||
|
|
||||||
tracks_tensor[batch_idx, frame_idx, track_idx] = torch.tensor([
|
tracks_tensor[batch_idx, frame_idx, track_idx] = torch.tensor([
|
||||||
track_idx, # track_id
|
track_idx, # track_id
|
||||||
x_norm, # x coordinate
|
x_norm, # x coordinate
|
||||||
y_norm, # y coordinate
|
y_norm, # y coordinate
|
||||||
visible # visibility
|
visible # visibility
|
||||||
])
|
])
|
||||||
|
|
||||||
return tracks_tensor
|
return tracks_tensor
|
||||||
|
|
||||||
def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], num_frames, quant_multi: int = 8, **kwargs):
|
def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], num_frames, quant_multi: int = 8, **kwargs):
|
||||||
# tracks: shape [t, h, w, 3] => samples align with 24 fps, model trained with 16 fps.
|
# tracks: shape [t, h, w, 3] => samples align with 24 fps, model trained with 16 fps.
|
||||||
# frame_size: tuple (W, H)
|
# frame_size: tuple (W, H)
|
||||||
tracks = torch.from_numpy(tracks_np).float()
|
tracks = torch.from_numpy(tracks_np).float()
|
||||||
|
|
||||||
if tracks.shape[1] == 121:
|
if tracks.shape[1] == 121:
|
||||||
tracks = torch.permute(tracks, (1, 0, 2, 3))
|
tracks = torch.permute(tracks, (1, 0, 2, 3))
|
||||||
|
|
||||||
tracks, visibles = tracks[..., :2], tracks[..., 2:3]
|
tracks, visibles = tracks[..., :2], tracks[..., 2:3]
|
||||||
|
|
||||||
short_edge = min(*frame_size)
|
short_edge = min(*frame_size)
|
||||||
|
|
||||||
frame_center = torch.tensor([*frame_size]).type_as(tracks) / 2
|
frame_center = torch.tensor([*frame_size]).type_as(tracks) / 2
|
||||||
tracks = tracks - frame_center
|
tracks = tracks - frame_center
|
||||||
|
|
||||||
tracks = tracks / short_edge * 2
|
tracks = tracks / short_edge * 2
|
||||||
|
|
||||||
visibles = visibles * 2 - 1
|
visibles = visibles * 2 - 1
|
||||||
@ -468,16 +467,16 @@ def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], num_frame
|
|||||||
trange = torch.linspace(-1, 1, tracks.shape[0]).view(-1, 1, 1, 1).expand(*visibles.shape)
|
trange = torch.linspace(-1, 1, tracks.shape[0]).view(-1, 1, 1, 1).expand(*visibles.shape)
|
||||||
|
|
||||||
out_ = torch.cat([trange, tracks, visibles], dim=-1).view(121, -1, 4)
|
out_ = torch.cat([trange, tracks, visibles], dim=-1).view(121, -1, 4)
|
||||||
|
|
||||||
out_0 = out_[:1]
|
out_0 = out_[:1]
|
||||||
|
|
||||||
out_l = out_[1:] # 121 => 120 | 1
|
out_l = out_[1:] # 121 => 120 | 1
|
||||||
a = 120 // math.gcd(120, num_frames)
|
a = 120 // math.gcd(120, num_frames)
|
||||||
b = num_frames // math.gcd(120, num_frames)
|
b = num_frames // math.gcd(120, num_frames)
|
||||||
out_l = torch.repeat_interleave(out_l, b, dim=0)[1::a] # 120 => 120 * b => 120 * b / a == F
|
out_l = torch.repeat_interleave(out_l, b, dim=0)[1::a] # 120 => 120 * b => 120 * b / a == F
|
||||||
|
|
||||||
final_result = torch.cat([out_0, out_l], dim=0)
|
final_result = torch.cat([out_0, out_l], dim=0)
|
||||||
|
|
||||||
return final_result
|
return final_result
|
||||||
|
|
||||||
FIXED_LENGTH = 121
|
FIXED_LENGTH = 121
|
||||||
@ -518,16 +517,16 @@ class WanTrackToVideo:
|
|||||||
|
|
||||||
CATEGORY = "conditioning/video_models"
|
CATEGORY = "conditioning/video_models"
|
||||||
|
|
||||||
def encode(self, positive, negative, vae, tracks, width, height, length, batch_size,
|
def encode(self, positive, negative, vae, tracks, width, height, length, batch_size,
|
||||||
temperature, topk, start_image=None, clip_vision_output=None):
|
temperature, topk, start_image=None, clip_vision_output=None):
|
||||||
|
|
||||||
# Parse tracks from JSON
|
# Parse tracks from JSON
|
||||||
tracks_data = parse_json_tracks(tracks)
|
tracks_data = parse_json_tracks(tracks)
|
||||||
|
|
||||||
if not tracks_data:
|
if not tracks_data:
|
||||||
return WanImageToVideo().encode(positive, negative, vae, width, height, length, batch_size, start_image=start_image, clip_vision_output=clip_vision_output)
|
return WanImageToVideo().encode(positive, negative, vae, width, height, length, batch_size, start_image=start_image, clip_vision_output=clip_vision_output)
|
||||||
|
|
||||||
latent = torch.zeros([batch_size, 16, ((length - 1) // 4) + 1, height // 8, width // 8],
|
latent = torch.zeros([batch_size, 16, ((length - 1) // 4) + 1, height // 8, width // 8],
|
||||||
device=comfy.model_management.intermediate_device())
|
device=comfy.model_management.intermediate_device())
|
||||||
# Convert tracks to tensor format
|
# Convert tracks to tensor format
|
||||||
arrs = []
|
arrs = []
|
||||||
@ -537,16 +536,16 @@ class WanTrackToVideo:
|
|||||||
|
|
||||||
tracks_np = np.stack(arrs, axis=0)
|
tracks_np = np.stack(arrs, axis=0)
|
||||||
processed_tracks = process_tracks(tracks_np, (width, height), length - 1).unsqueeze(0)
|
processed_tracks = process_tracks(tracks_np, (width, height), length - 1).unsqueeze(0)
|
||||||
|
|
||||||
if start_image is not None:
|
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)
|
start_image = comfy.utils.common_upscale(start_image[:length].movedim(-1, 1), width, height, "bilinear", "center").movedim(1, -1)
|
||||||
|
|
||||||
lat_h = height // 8
|
lat_h = height // 8
|
||||||
lat_w = width // 8
|
lat_w = width // 8
|
||||||
|
|
||||||
msk = torch.ones(1, length, lat_h, lat_w, device=start_image.device)
|
msk = torch.ones(1, length, lat_h, lat_w, device=start_image.device)
|
||||||
msk[:, 1:] = 0
|
msk[:, 1:] = 0
|
||||||
|
|
||||||
# repeat first frame 4 times
|
# repeat first frame 4 times
|
||||||
msk = torch.concat([
|
msk = torch.concat([
|
||||||
torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]
|
torch.repeat_interleave(msk[:, 0:1], repeats=4, dim=1), msk[:, 1:]
|
||||||
@ -560,7 +559,7 @@ class WanTrackToVideo:
|
|||||||
msk = msk.transpose(1, 2)
|
msk = msk.transpose(1, 2)
|
||||||
|
|
||||||
dummy_frames = torch.ones(3, length - 1, height, width) * .5
|
dummy_frames = torch.ones(3, length - 1, height, width) * .5
|
||||||
|
|
||||||
start_image = start_image.permute(3,0,1,2) # C, T, H, W
|
start_image = start_image.permute(3,0,1,2) # C, T, H, W
|
||||||
res = torch.concat([
|
res = torch.concat([
|
||||||
start_image.to(start_image.device),
|
start_image.to(start_image.device),
|
||||||
@ -569,9 +568,9 @@ class WanTrackToVideo:
|
|||||||
dim=1).to(start_image.device)
|
dim=1).to(start_image.device)
|
||||||
|
|
||||||
res = res.permute(1,2,3,0)[:, :, :, :3] # T, H, W, C
|
res = res.permute(1,2,3,0)[:, :, :, :3] # T, H, W, C
|
||||||
|
|
||||||
y = vae.encode(res)
|
y = vae.encode(res)
|
||||||
|
|
||||||
# Add motion features to conditioning
|
# Add motion features to conditioning
|
||||||
positive = node_helpers.conditioning_set_values(positive,
|
positive = node_helpers.conditioning_set_values(positive,
|
||||||
{"tracks": processed_tracks,
|
{"tracks": processed_tracks,
|
||||||
@ -579,13 +578,13 @@ class WanTrackToVideo:
|
|||||||
"concat_latent_image": y,
|
"concat_latent_image": y,
|
||||||
"ati_temperature": temperature,
|
"ati_temperature": temperature,
|
||||||
"ati_topk": topk})
|
"ati_topk": topk})
|
||||||
negative = node_helpers.conditioning_set_values(negative,
|
negative = node_helpers.conditioning_set_values(negative,
|
||||||
{"tracks": processed_tracks,
|
{"tracks": processed_tracks,
|
||||||
"concat_mask": msk,
|
"concat_mask": msk,
|
||||||
"concat_latent_image": y,
|
"concat_latent_image": y,
|
||||||
"ati_temperature": temperature,
|
"ati_temperature": temperature,
|
||||||
"ati_topk": topk})
|
"ati_topk": topk})
|
||||||
|
|
||||||
|
|
||||||
# Handle clip vision output if provided
|
# Handle clip vision output if provided
|
||||||
if clip_vision_output is not None:
|
if clip_vision_output is not None:
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user