mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-02 07:07:03 +08:00
Add batch trajectory logic
This commit is contained in:
parent
ef382b319d
commit
877ebd9240
@ -1073,111 +1073,6 @@ class Lumina2(BaseModel):
|
|||||||
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
|
out['c_crossattn'] = comfy.conds.CONDRegular(cross_attn)
|
||||||
return out
|
return out
|
||||||
|
|
||||||
def ind_sel(target: torch.Tensor, ind: torch.Tensor, dim: int = 1):
|
|
||||||
"""Index selection utility function"""
|
|
||||||
assert (
|
|
||||||
len(ind.shape) > dim
|
|
||||||
), "Index must have the target dim, but get dim: %d, ind shape: %s" % (dim, str(ind.shape))
|
|
||||||
|
|
||||||
target = target.expand(
|
|
||||||
*tuple(
|
|
||||||
[ind.shape[k] if target.shape[k] == 1 else -1 for k in range(dim)]
|
|
||||||
+ [
|
|
||||||
-1,
|
|
||||||
]
|
|
||||||
* (len(target.shape) - dim)
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
ind_pad = ind
|
|
||||||
|
|
||||||
if len(target.shape) > dim + 1:
|
|
||||||
for _ in range(len(target.shape) - (dim + 1)):
|
|
||||||
ind_pad = ind_pad.unsqueeze(-1)
|
|
||||||
ind_pad = ind_pad.expand(*(-1,) * (dim + 1), *target.shape[(dim + 1) : :])
|
|
||||||
|
|
||||||
return torch.gather(target, dim=dim, index=ind_pad)
|
|
||||||
|
|
||||||
|
|
||||||
def merge_final(vert_attr: torch.Tensor, weight: torch.Tensor, vert_assign: torch.Tensor):
|
|
||||||
"""Merge vertex attributes with weights"""
|
|
||||||
target_dim = len(vert_assign.shape) - 1
|
|
||||||
if len(vert_attr.shape) == 2:
|
|
||||||
assert vert_attr.shape[0] > vert_assign.max()
|
|
||||||
new_shape = [1] * target_dim + list(vert_attr.shape)
|
|
||||||
tensor = vert_attr.reshape(new_shape)
|
|
||||||
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
|
|
||||||
else:
|
|
||||||
assert vert_attr.shape[1] > vert_assign.max()
|
|
||||||
new_shape = [vert_attr.shape[0]] + [1] * (target_dim - 1) + list(vert_attr.shape[1:])
|
|
||||||
tensor = vert_attr.reshape(new_shape)
|
|
||||||
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
|
|
||||||
|
|
||||||
final_attr = torch.sum(sel_attr * weight.unsqueeze(-1), dim=-2)
|
|
||||||
return final_attr
|
|
||||||
|
|
||||||
|
|
||||||
def patch_motion(
|
|
||||||
tracks: torch.FloatTensor, # (B, T, N, 4)
|
|
||||||
vid: torch.FloatTensor, # (C, T, H, W)
|
|
||||||
temperature: float = 220.0,
|
|
||||||
vae_divide: tuple = (4, 16),
|
|
||||||
topk: int = 2,
|
|
||||||
):
|
|
||||||
"""Apply motion patching based on tracks"""
|
|
||||||
_, T, H, W = vid.shape
|
|
||||||
N = tracks.shape[2]
|
|
||||||
_, tracks_xy, visible = torch.split(
|
|
||||||
tracks, [1, 2, 1], dim=-1
|
|
||||||
) # (B, T, N, 2) | (B, T, N, 1)
|
|
||||||
tracks_n = tracks_xy / torch.tensor([W / min(H, W), H / min(H, W)], device=tracks_xy.device)
|
|
||||||
tracks_n = tracks_n.clamp(-1, 1)
|
|
||||||
visible = 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)
|
|
||||||
|
|
||||||
grid = torch.stack(torch.meshgrid(yy, xx, indexing="ij")[::-1], dim=-1).to(
|
|
||||||
tracks_xy.device
|
|
||||||
)
|
|
||||||
|
|
||||||
tracks_pad = tracks_xy[:, 1:]
|
|
||||||
visible_pad = visible[:, 1:]
|
|
||||||
|
|
||||||
visible_align = visible_pad.view(T - 1, 4, *visible_pad.shape[2:]).sum(1)
|
|
||||||
tracks_align = (tracks_pad * visible_pad).view(T - 1, 4, *tracks_pad.shape[2:]).sum(
|
|
||||||
1
|
|
||||||
) / (visible_align + 1e-5)
|
|
||||||
dist_ = (
|
|
||||||
(tracks_align[:, None, None] - grid[None, :, :, None]).pow(2).sum(-1)
|
|
||||||
) # T, H, W, N
|
|
||||||
weight = torch.exp(-dist_ * temperature) * visible_align.clamp(0, 1).view(
|
|
||||||
T - 1, 1, 1, N
|
|
||||||
)
|
|
||||||
vert_weight, vert_index = torch.topk(
|
|
||||||
weight, k=min(topk, weight.shape[-1]), dim=-1
|
|
||||||
)
|
|
||||||
|
|
||||||
grid_mode = "bilinear"
|
|
||||||
point_feature = torch.nn.functional.grid_sample(
|
|
||||||
vid[vae_divide[0]:].permute(1, 0, 2, 3)[:1],
|
|
||||||
tracks_n[:, :1].type(vid.dtype),
|
|
||||||
mode=grid_mode,
|
|
||||||
padding_mode="zeros",
|
|
||||||
align_corners=False,
|
|
||||||
)
|
|
||||||
point_feature = point_feature.squeeze(0).squeeze(1).permute(1, 0) # N, C=16
|
|
||||||
|
|
||||||
out_feature = merge_final(point_feature, vert_weight, vert_index).permute(3, 0, 1, 2) # T - 1, H, W, C => C, T - 1, H, W
|
|
||||||
out_weight = vert_weight.sum(-1) # T - 1, H, W
|
|
||||||
|
|
||||||
# out feature -> already soft weighted
|
|
||||||
mix_feature = out_feature + vid[vae_divide[0]:, 1:] * (1 - out_weight.clamp(0, 1))
|
|
||||||
|
|
||||||
out_feature_full = torch.cat([vid[vae_divide[0]:, :1], mix_feature], dim=1) # C, T, H, W
|
|
||||||
out_mask_full = torch.cat([torch.ones_like(out_weight[:1]), out_weight], dim=0) # T, H, W
|
|
||||||
return torch.cat([out_mask_full[None].expand(vae_divide[0], -1, -1, -1), out_feature_full], dim=0)
|
|
||||||
|
|
||||||
class WAN21(BaseModel):
|
class WAN21(BaseModel):
|
||||||
def __init__(self, model_config, model_type=ModelType.FLOW, image_to_video=False, device=None):
|
def __init__(self, model_config, model_type=ModelType.FLOW, image_to_video=False, device=None):
|
||||||
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.wan.model.WanModel)
|
super().__init__(model_config, model_type, device=device, unet_model=comfy.ldm.wan.model.WanModel)
|
||||||
@ -1222,14 +1117,7 @@ class WAN21(BaseModel):
|
|||||||
mask = mask.repeat(1, 4, 1, 1, 1)
|
mask = mask.repeat(1, 4, 1, 1, 1)
|
||||||
mask = utils.resize_to_batch_size(mask, noise.shape[0])
|
mask = utils.resize_to_batch_size(mask, noise.shape[0])
|
||||||
|
|
||||||
res = torch.cat((mask, image), dim=1)
|
return torch.cat((mask, image), dim=1)
|
||||||
|
|
||||||
tracks = kwargs.get("tracks", None)
|
|
||||||
if tracks is not None:
|
|
||||||
res = patch_motion(tracks.to(device), res[0], kwargs.get("ati_temperature", None), (4, 16), kwargs.get("ati_topk", None))[None]
|
|
||||||
|
|
||||||
return res
|
|
||||||
|
|
||||||
|
|
||||||
def extra_conds(self, **kwargs):
|
def extra_conds(self, **kwargs):
|
||||||
out = super().extra_conds(**kwargs)
|
out = super().extra_conds(**kwargs)
|
||||||
|
|||||||
@ -491,6 +491,142 @@ def pad_pts(tr):
|
|||||||
pts = pts[:FIXED_LENGTH]
|
pts = pts[:FIXED_LENGTH]
|
||||||
return pts.reshape(FIXED_LENGTH, 1, 3)
|
return pts.reshape(FIXED_LENGTH, 1, 3)
|
||||||
|
|
||||||
|
def ind_sel(target: torch.Tensor, ind: torch.Tensor, dim: int = 1):
|
||||||
|
"""Index selection utility function"""
|
||||||
|
assert (
|
||||||
|
len(ind.shape) > dim
|
||||||
|
), "Index must have the target dim, but get dim: %d, ind shape: %s" % (dim, str(ind.shape))
|
||||||
|
|
||||||
|
target = target.expand(
|
||||||
|
*tuple(
|
||||||
|
[ind.shape[k] if target.shape[k] == 1 else -1 for k in range(dim)]
|
||||||
|
+ [
|
||||||
|
-1,
|
||||||
|
]
|
||||||
|
* (len(target.shape) - dim)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
|
ind_pad = ind
|
||||||
|
|
||||||
|
if len(target.shape) > dim + 1:
|
||||||
|
for _ in range(len(target.shape) - (dim + 1)):
|
||||||
|
ind_pad = ind_pad.unsqueeze(-1)
|
||||||
|
ind_pad = ind_pad.expand(*(-1,) * (dim + 1), *target.shape[(dim + 1) : :])
|
||||||
|
|
||||||
|
return torch.gather(target, dim=dim, index=ind_pad)
|
||||||
|
|
||||||
|
def merge_final(vert_attr: torch.Tensor, weight: torch.Tensor, vert_assign: torch.Tensor):
|
||||||
|
"""Merge vertex attributes with weights"""
|
||||||
|
target_dim = len(vert_assign.shape) - 1
|
||||||
|
if len(vert_attr.shape) == 2:
|
||||||
|
assert vert_attr.shape[0] > vert_assign.max()
|
||||||
|
new_shape = [1] * target_dim + list(vert_attr.shape)
|
||||||
|
tensor = vert_attr.reshape(new_shape)
|
||||||
|
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
|
||||||
|
else:
|
||||||
|
assert vert_attr.shape[1] > vert_assign.max()
|
||||||
|
new_shape = [vert_attr.shape[0]] + [1] * (target_dim - 1) + list(vert_attr.shape[1:])
|
||||||
|
tensor = vert_attr.reshape(new_shape)
|
||||||
|
sel_attr = ind_sel(tensor, vert_assign.type(torch.long), dim=target_dim)
|
||||||
|
|
||||||
|
final_attr = torch.sum(sel_attr * weight.unsqueeze(-1), dim=-2)
|
||||||
|
return final_attr
|
||||||
|
|
||||||
|
|
||||||
|
def _patch_motion_single(
|
||||||
|
tracks: torch.FloatTensor, # (B, T, N, 4)
|
||||||
|
vid: torch.FloatTensor, # (C, T, H, W)
|
||||||
|
temperature: float,
|
||||||
|
vae_divide: tuple,
|
||||||
|
topk: int,
|
||||||
|
):
|
||||||
|
"""Apply motion patching based on tracks"""
|
||||||
|
_, T, H, W = vid.shape
|
||||||
|
N = tracks.shape[2]
|
||||||
|
_, tracks_xy, visible = torch.split(
|
||||||
|
tracks, [1, 2, 1], dim=-1
|
||||||
|
) # (B, T, N, 2) | (B, T, N, 1)
|
||||||
|
tracks_n = tracks_xy / torch.tensor([W / min(H, W), H / min(H, W)], device=tracks_xy.device)
|
||||||
|
tracks_n = tracks_n.clamp(-1, 1)
|
||||||
|
visible = 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)
|
||||||
|
|
||||||
|
grid = torch.stack(torch.meshgrid(yy, xx, indexing="ij")[::-1], dim=-1).to(
|
||||||
|
tracks_xy.device
|
||||||
|
)
|
||||||
|
|
||||||
|
tracks_pad = tracks_xy[:, 1:]
|
||||||
|
visible_pad = visible[:, 1:]
|
||||||
|
|
||||||
|
visible_align = visible_pad.view(T - 1, 4, *visible_pad.shape[2:]).sum(1)
|
||||||
|
tracks_align = (tracks_pad * visible_pad).view(T - 1, 4, *tracks_pad.shape[2:]).sum(
|
||||||
|
1
|
||||||
|
) / (visible_align + 1e-5)
|
||||||
|
dist_ = (
|
||||||
|
(tracks_align[:, None, None] - grid[None, :, :, None]).pow(2).sum(-1)
|
||||||
|
) # T, H, W, N
|
||||||
|
weight = torch.exp(-dist_ * temperature) * visible_align.clamp(0, 1).view(
|
||||||
|
T - 1, 1, 1, N
|
||||||
|
)
|
||||||
|
vert_weight, vert_index = torch.topk(
|
||||||
|
weight, k=min(topk, weight.shape[-1]), dim=-1
|
||||||
|
)
|
||||||
|
|
||||||
|
grid_mode = "bilinear"
|
||||||
|
point_feature = torch.nn.functional.grid_sample(
|
||||||
|
vid[vae_divide[0]:].permute(1, 0, 2, 3)[:1],
|
||||||
|
tracks_n[:, :1].type(vid.dtype),
|
||||||
|
mode=grid_mode,
|
||||||
|
padding_mode="zeros",
|
||||||
|
align_corners=False,
|
||||||
|
)
|
||||||
|
point_feature = point_feature.squeeze(0).squeeze(1).permute(1, 0) # N, C=16
|
||||||
|
|
||||||
|
out_feature = merge_final(point_feature, vert_weight, vert_index).permute(3, 0, 1, 2) # T - 1, H, W, C => C, T - 1, H, W
|
||||||
|
out_weight = vert_weight.sum(-1) # T - 1, H, W
|
||||||
|
|
||||||
|
# out feature -> already soft weighted
|
||||||
|
mix_feature = out_feature + vid[vae_divide[0]:, 1:] * (1 - out_weight.clamp(0, 1))
|
||||||
|
|
||||||
|
out_feature_full = torch.cat([vid[vae_divide[0]:, :1], mix_feature], dim=1) # C, T, H, W
|
||||||
|
out_mask_full = torch.cat([torch.ones_like(out_weight[:1]), out_weight], dim=0) # T, H, W
|
||||||
|
|
||||||
|
return out_mask_full[None].expand(vae_divide[0], -1, -1, -1), out_feature_full
|
||||||
|
|
||||||
|
|
||||||
|
def patch_motion(
|
||||||
|
tracks: torch.FloatTensor, # (B, TB, T, N, 4)
|
||||||
|
vid: torch.FloatTensor, # (C, T, H, W)
|
||||||
|
temperature: float = 220.0,
|
||||||
|
vae_divide: tuple = (4, 16),
|
||||||
|
topk: int = 2,
|
||||||
|
):
|
||||||
|
B = len(tracks)
|
||||||
|
|
||||||
|
# Process each batch separately
|
||||||
|
out_masks = []
|
||||||
|
out_features = []
|
||||||
|
|
||||||
|
for b in range(B):
|
||||||
|
mask, feature = _patch_motion_single(
|
||||||
|
tracks[b], # (T, N, 4)
|
||||||
|
vid, # (C, T, H, W)
|
||||||
|
temperature,
|
||||||
|
vae_divide,
|
||||||
|
topk
|
||||||
|
)
|
||||||
|
out_masks.append(mask)
|
||||||
|
out_features.append(feature)
|
||||||
|
|
||||||
|
# Stack results: (B, C, T, H, W)
|
||||||
|
out_mask_full = torch.stack(out_masks, dim=0)
|
||||||
|
out_feature_full = torch.stack(out_features, dim=0)
|
||||||
|
|
||||||
|
return out_mask_full, out_feature_full
|
||||||
|
|
||||||
class WanTrackToVideo:
|
class WanTrackToVideo:
|
||||||
@classmethod
|
@classmethod
|
||||||
def INPUT_TYPES(s):
|
def INPUT_TYPES(s):
|
||||||
@ -529,14 +665,19 @@ class WanTrackToVideo:
|
|||||||
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 = []
|
if isinstance(tracks_data[0][0], dict):
|
||||||
for i, track in enumerate(tracks_data):
|
tracks_data = [tracks_data]
|
||||||
pts = pad_pts(track)
|
|
||||||
arrs.append(pts)
|
|
||||||
|
|
||||||
tracks_np = np.stack(arrs, axis=0)
|
processed_tracks = []
|
||||||
processed_tracks = process_tracks(tracks_np, (width, height), length - 1).unsqueeze(0)
|
for batch in tracks_data:
|
||||||
|
arrs = []
|
||||||
|
for track in batch:
|
||||||
|
pts = pad_pts(track)
|
||||||
|
arrs.append(pts)
|
||||||
|
|
||||||
|
tracks_np = np.stack(arrs, axis=0)
|
||||||
|
processed_tracks.append(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)
|
||||||
|
|
||||||
@ -570,8 +711,13 @@ class WanTrackToVideo:
|
|||||||
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)
|
||||||
|
c = torch.cat((msk, y), dim=1)
|
||||||
|
res = patch_motion(
|
||||||
|
processed_tracks, c[0], temperature=temperature, topk=topk, vae_divide=(4, 16)
|
||||||
|
)
|
||||||
|
|
||||||
# Add motion features to conditioning
|
msk, y = res
|
||||||
|
msk = -msk + 1.0 # Invert mask to match expected format
|
||||||
positive = node_helpers.conditioning_set_values(positive,
|
positive = node_helpers.conditioning_set_values(positive,
|
||||||
{"tracks": processed_tracks,
|
{"tracks": processed_tracks,
|
||||||
"concat_mask": msk,
|
"concat_mask": msk,
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user