mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-04 00:27:05 +08:00
Fix length
This commit is contained in:
parent
74e5afcbf2
commit
b293fa4d5b
@ -1227,7 +1227,7 @@ class WAN21(BaseModel):
|
|||||||
|
|
||||||
tracks = kwargs.get("tracks", None)
|
tracks = kwargs.get("tracks", None)
|
||||||
if tracks is not None:
|
if tracks is not None:
|
||||||
res = patch_motion(tracks.to(device), res[0], 220.0, (4, 16), 2)[None]
|
res = patch_motion(tracks.to(device), res[0], kwargs.get("ati_temperature", None), (4, 16), kwargs.get("ati_topk", None))[None]
|
||||||
|
|
||||||
return res
|
return res
|
||||||
|
|
||||||
|
|||||||
@ -1,3 +1,4 @@
|
|||||||
|
import math
|
||||||
import nodes
|
import nodes
|
||||||
import node_helpers
|
import node_helpers
|
||||||
import torch
|
import torch
|
||||||
@ -490,7 +491,7 @@ def merge_final(vert_attr: torch.Tensor, weight: torch.Tensor, vert_assign: torc
|
|||||||
final_attr = torch.sum(sel_attr * weight.unsqueeze(-1), dim=-2)
|
final_attr = torch.sum(sel_attr * weight.unsqueeze(-1), dim=-2)
|
||||||
return final_attr
|
return final_attr
|
||||||
|
|
||||||
def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], 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()
|
||||||
@ -516,7 +517,9 @@ def process_tracks(tracks_np: np.ndarray, frame_size: Tuple[int, int], quant_mul
|
|||||||
out_0 = out_[:1]
|
out_0 = out_[:1]
|
||||||
|
|
||||||
out_l = out_[1:] # 121 => 120 | 1
|
out_l = out_[1:] # 121 => 120 | 1
|
||||||
out_l = torch.repeat_interleave(out_l, 2, dim=0)[1::3] # 120 => 240 => 80
|
a = 120 // 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
|
||||||
|
|
||||||
final_result = torch.cat([out_0, out_l], dim=0)
|
final_result = torch.cat([out_0, out_l], dim=0)
|
||||||
|
|
||||||
@ -566,59 +569,67 @@ class WanTrackToVideo:
|
|||||||
# Parse tracks from JSON
|
# Parse tracks from JSON
|
||||||
tracks_data = parse_json_tracks(tracks)
|
tracks_data = parse_json_tracks(tracks)
|
||||||
|
|
||||||
if tracks_data:
|
if not tracks_data:
|
||||||
# Convert tracks to tensor format
|
return WanImageToVideo().encode(positive, negative, vae, width, height, length, batch_size, start_image=start_image, clip_vision_output=clip_vision_output)
|
||||||
arrs = []
|
|
||||||
for i, track in enumerate(tracks_data):
|
latent = torch.zeros([batch_size, 16, ((length - 1) // 4) + 1, height // 8, width // 8],
|
||||||
pts = pad_pts(track)
|
device=comfy.model_management.intermediate_device())
|
||||||
arrs.append(pts)
|
# Convert tracks to tensor format
|
||||||
|
arrs = []
|
||||||
|
for i, track in enumerate(tracks_data):
|
||||||
|
pts = pad_pts(track)
|
||||||
|
arrs.append(pts)
|
||||||
|
|
||||||
tracks_np = np.stack(arrs, axis=0)
|
tracks_np = np.stack(arrs, axis=0)
|
||||||
processed_tracks = process_tracks(tracks_np, (width, height)).unsqueeze(0)
|
processed_tracks = process_tracks(tracks_np, (width, height), length - 1).unsqueeze(0)
|
||||||
|
|
||||||
|
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)
|
||||||
|
|
||||||
if start_image is not None:
|
lat_h = height // 8
|
||||||
start_image = comfy.utils.common_upscale(start_image[:length].movedim(-1, 1), width, height, "bilinear", "center").movedim(1, -1)
|
lat_w = width // 8
|
||||||
|
|
||||||
lat_h = height // 8
|
|
||||||
lat_w = width // 8
|
|
||||||
|
|
||||||
msk = torch.ones(1, 81, lat_h, lat_w, device=start_image.device)
|
msk = torch.ones(1, latent.shape[2], 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:]
|
||||||
|
],
|
||||||
|
dim=1)
|
||||||
|
|
||||||
|
# Reshape mask into groups of 4 frames
|
||||||
|
msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w)
|
||||||
|
|
||||||
|
# first batch
|
||||||
|
msk = msk.transpose(1, 2)
|
||||||
|
|
||||||
|
dummy_frames = torch.ones(3, length - 1, height, width) * .5
|
||||||
|
|
||||||
|
start_image = start_image.permute(3,0,1,2) # C, T, H, W
|
||||||
|
res = torch.concat([
|
||||||
|
start_image.to(start_image.device),
|
||||||
|
dummy_frames
|
||||||
],
|
],
|
||||||
dim=1)
|
dim=1).to(start_image.device)
|
||||||
|
|
||||||
# Reshape mask into groups of 4 frames
|
res = res.permute(1,2,3,0)[:, :, :, :3] # T, H, W, C
|
||||||
msk = msk.view(1, msk.shape[1] // 4, 4, lat_h, lat_w)
|
|
||||||
|
y = vae.encode(res)
|
||||||
# first batch
|
|
||||||
msk = msk.transpose(1, 2)
|
# Add motion features to conditioning
|
||||||
|
positive = node_helpers.conditioning_set_values(positive,
|
||||||
zero_frames = torch.ones(3, 81 - 1, height, width) * .5
|
{"tracks": processed_tracks,
|
||||||
|
"concat_mask": msk,
|
||||||
start_image = start_image.permute(3,0,1,2) # C, T, H, W
|
"concat_latent_image": y,
|
||||||
res = torch.concat([
|
"ati_temperature": temperature,
|
||||||
start_image.to(start_image.device),
|
"ati_topk": topk})
|
||||||
zero_frames
|
negative = node_helpers.conditioning_set_values(negative,
|
||||||
],
|
{"tracks": processed_tracks,
|
||||||
dim=1).to(start_image.device)
|
"concat_mask": msk,
|
||||||
|
"concat_latent_image": y,
|
||||||
res = res.permute(1,2,3,0)[:, :, :, :3] # T, H, W, C
|
"ati_temperature": temperature,
|
||||||
|
"ati_topk": topk})
|
||||||
y = vae.encode(res)
|
|
||||||
|
|
||||||
# Add motion features to conditioning
|
|
||||||
positive = node_helpers.conditioning_set_values(positive,
|
|
||||||
{"tracks": processed_tracks,
|
|
||||||
"concat_mask": msk,
|
|
||||||
"concat_latent_image": y})
|
|
||||||
negative = node_helpers.conditioning_set_values(negative,
|
|
||||||
{"tracks": processed_tracks,
|
|
||||||
"concat_mask": msk,
|
|
||||||
"concat_latent_image": y})
|
|
||||||
|
|
||||||
|
|
||||||
# Handle clip vision output if provided
|
# Handle clip vision output if provided
|
||||||
@ -626,8 +637,6 @@ class WanTrackToVideo:
|
|||||||
positive = node_helpers.conditioning_set_values(positive, {"clip_vision_output": clip_vision_output})
|
positive = node_helpers.conditioning_set_values(positive, {"clip_vision_output": clip_vision_output})
|
||||||
negative = node_helpers.conditioning_set_values(negative, {"clip_vision_output": clip_vision_output})
|
negative = node_helpers.conditioning_set_values(negative, {"clip_vision_output": clip_vision_output})
|
||||||
|
|
||||||
latent = torch.zeros([batch_size, 16, ((length - 1) // 4) + 1, height // 8, width // 8],
|
|
||||||
device=comfy.model_management.intermediate_device())
|
|
||||||
out_latent = {}
|
out_latent = {}
|
||||||
out_latent["samples"] = latent
|
out_latent["samples"] = latent
|
||||||
return (positive, negative, out_latent)
|
return (positive, negative, out_latent)
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user