Squashed commit of the following:
commit 916fc0b1bcfd37b6bd9ece0daeb5b3cbaa53d0a9 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Dec 15 17:30:37 2025 +0200 Update nodes.py commit 63818324f5dbb0b300064bea0402c4cd1bd57b2b Author: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Dec 15 17:30:26 2025 +0200 Refactor RoPE caching commit bb0c55da4d8f8bca4968704e877fd057a90a1eeb Author: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Dec 15 01:59:16 2025 +0200 Update nodes_sampler.py commit a0447d55534857051606ee4201bc7f4e25aa73ae Author: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Dec 15 01:28:09 2025 +0200 Fix non scale wfs commit fa761cc2f2a426faa9c391aeede62cf6f0fd7266 Merge: ea1677b3aae54fAuthor: kijai <40791699+kijai@users.noreply.github.com> Date: Mon Dec 15 01:26:23 2025 +0200 Merge branch 'main' into SCAIL commit ea1677bd4ad42f19e369551590a9d4f17a36fa29 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Dec 14 19:41:43 2025 +0200 Handle torchscript issue better Some other custom nodes globally set torch._C._jit_set_profiling_executor(False) which breaks the NLF model commit e3cfa64bd3712ac153ce84a75215842c884a8ba4 Merge: ad7a0b93611341Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Dec 14 16:49:04 2025 +0200 Merge branch 'main' into SCAIL commit ad7a0b925de61ff705b928cd802e752e46089b42 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Dec 14 16:10:34 2025 +0200 Fix possible uni3c issue commit 74d97fa4bb7c58a0edf8516cc9fad4468da5c57e Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Dec 14 15:58:42 2025 +0200 Match Uni3C temporal dim commit 056d8ad96ffa5a223a8cd88c900a573a8d450e22 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Dec 14 14:47:58 2025 +0200 Add warning for potential other overrides on torch.jit.script commit f6dff002ffdcd880451955db298872ea90a4e3f8 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Dec 14 14:19:33 2025 +0200 Add option to warmup the NLF model on load and fix it's offloading commit a19107501dff23804e7db984d7da304a9955adc9 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sun Dec 14 13:45:20 2025 +0200 Add error to indicate ComfyUI-RMBG currently breaks the NLF model commit e2cfa486e48ead50195884167d9794c7caf0a69f Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 23:29:49 2025 +0200 Cleanup unnecessary code commit 462b61855fb96b0cb18cbccd48593256992808d7 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 18:05:10 2025 +0200 context windows commit e57d4baeebf12c43e851c6c2467d698d7dbb4d03 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 16:55:23 2025 +0200 Start/end percentages and strength commit 3e507ae32256ed3e41cea69d9e26c30b5272968e Merge: 1e5c7cb0fa5383Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 16:09:16 2025 +0200 Merge branch 'main' into SCAIL commit 1e5c7cb2113138bdeae562d266f911c1e3edee91 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 15:45:39 2025 +0200 Update nodes.py commit 98f8e56bcacfc07e12cbb4b26555b2b28d9db92f Merge: 965214678e3e18Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 15:42:44 2025 +0200 Merge branch 'main' into SCAIL commit 9652146763fb27e916a6853a8125efd0a67cd601 Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 02:41:06 2025 +0200 Add imitation of SCAIL pose drawing to the existing NLF node This only draws the pose with same colors, it's not meant as final solution, just for testing. commit 1f86cebdaa97570ed88da0c9986b85c6664d62dc Author: kijai <40791699+kijai@users.noreply.github.com> Date: Sat Dec 13 01:11:56 2025 +0200 test pose inputs commit b348b21dbef0dcb92c0961df85959648e78da6aa Author: kijai <40791699+kijai@users.noreply.github.com> Date: Fri Dec 12 20:10:48 2025 +0200 Init
This commit is contained in:
+78
-7
@@ -43,28 +43,33 @@ def get_pose_images(smpl_data, offset):
|
|||||||
return pose_images
|
return pose_images
|
||||||
|
|
||||||
|
|
||||||
def get_control_conditions(poses, h, w):
|
def get_control_conditions(poses, h, w, stick_width=1.0, point_radius=2, style="original"):
|
||||||
video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
|
|
||||||
control_images = []
|
control_images = []
|
||||||
for idx, pose in enumerate(poses):
|
for idx, pose in enumerate(poses):
|
||||||
canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
|
canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
|
||||||
try:
|
try:
|
||||||
joints3d = p3d_to_p2d(pose, h, w)
|
joints3d = p3d_to_p2d(pose, h, w)
|
||||||
|
if style == "original":
|
||||||
canvas = draw_3d_points(
|
canvas = draw_3d_points(
|
||||||
canvas,
|
canvas,
|
||||||
joints3d[0],
|
joints3d[0],
|
||||||
stickwidth=int(h / 350),
|
stickwidth=int(h / 350 * stick_width),
|
||||||
|
r=point_radius,
|
||||||
|
)
|
||||||
|
elif style == "scail":
|
||||||
|
canvas = draw_3d_points_scail(
|
||||||
|
canvas,
|
||||||
|
joints3d[0],
|
||||||
|
stickwidth=int(h / 350 * stick_width),
|
||||||
|
r=point_radius,
|
||||||
)
|
)
|
||||||
resized_canvas = cv2.resize(canvas, (w, h))
|
resized_canvas = cv2.resize(canvas, (w, h))
|
||||||
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
|
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
|
||||||
control_images.append(resized_canvas)
|
control_images.append(resized_canvas)
|
||||||
except Exception as e:
|
except Exception:
|
||||||
print("wrong:", e)
|
|
||||||
control_images.append(Image.fromarray(canvas))
|
control_images.append(Image.fromarray(canvas))
|
||||||
control_pixel_values = np.array(control_images)
|
control_pixel_values = np.array(control_images)
|
||||||
control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
|
control_pixel_values = torch.from_numpy(control_pixel_values).contiguous() / 255.
|
||||||
print("control_pixel_values.shape", control_pixel_values.shape)
|
|
||||||
#control_pixel_values = video_transforms(control_pixel_values)
|
|
||||||
return control_pixel_values
|
return control_pixel_values
|
||||||
|
|
||||||
|
|
||||||
@@ -140,3 +145,69 @@ def draw_3d_points(canvas, points, stickwidth=2, r=2, draw_line=True):
|
|||||||
cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
|
cv2.fillConvexPoly(canvas, polygon, connection_colors[i%17])
|
||||||
|
|
||||||
return canvas
|
return canvas
|
||||||
|
|
||||||
|
def draw_3d_points_scail(canvas, points, stickwidth=2, r=2, draw_line=True):
|
||||||
|
|
||||||
|
connetions = [
|
||||||
|
[15,12],[12, 16],[16, 18],[18, 20],[20, 22], # 0-4: Left arm chain
|
||||||
|
[12,17],[17,19],[19,21], # 5-7: Right arm chain
|
||||||
|
[21,23], # 8: Right hand
|
||||||
|
[12,1],[1,4],[4,7], # 9-11: Neck to left leg (hip, thigh, shin)
|
||||||
|
[12,2],[2,5],[5,8], # 12-14: Neck to right leg (hip, thigh, shin)
|
||||||
|
]
|
||||||
|
|
||||||
|
# Warm colors for right side, cool colors for left side
|
||||||
|
connection_colors = [
|
||||||
|
[180, 180, 180], # 0: [15,12] - L. clavicle (Bright Cyan)
|
||||||
|
[0, 200, 255], # 1: [12,16] - L. shoulder (Bright Cyan)
|
||||||
|
[0, 120, 255], # 2: [16,18] - L. upper arm (Bright Blue)
|
||||||
|
[0, 60, 255], # 3: [18,20] - L. forearm (Deep Blue)
|
||||||
|
[60, 0, 255], # 4: [20,22] - L. hand (Blue-Purple)
|
||||||
|
[255, 0, 0], # 5: [12,17] - R. clavicle (Bright Red)
|
||||||
|
[255, 100, 0], # 6: [17,19] - R. upper arm (Bright Orange)
|
||||||
|
[255, 180, 0], # 7: [19,21] - R. forearm (Golden Orange)
|
||||||
|
[255, 255, 0], # 8: [21,23] - R. hand (Bright Yellow)
|
||||||
|
[30, 27, 160], # 9: [12,1] - Neck to L. hip (purple-blue)
|
||||||
|
[73, 27, 177], # 10: [1,4] - L. thigh (purple)
|
||||||
|
[145, 27, 194], # 11: [4,7] - L. shin (magenta)
|
||||||
|
[200, 255, 100], # 12: [12,2] - Neck to R. hip (yellow)
|
||||||
|
[54, 201, 52], # 13: [2,5] - R. thigh (green)
|
||||||
|
[30, 176, 85], # 14: [5,8] - R. shin (green)
|
||||||
|
]
|
||||||
|
|
||||||
|
# draw line
|
||||||
|
if draw_line:
|
||||||
|
# Collect all joints that are part of connections
|
||||||
|
joints_in_use = set()
|
||||||
|
for connection in connetions:
|
||||||
|
joints_in_use.add(connection[0])
|
||||||
|
joints_in_use.add(connection[1])
|
||||||
|
|
||||||
|
for i in range(len(connetions)):
|
||||||
|
point1_idx, point2_idx = connetions[i][0:2]
|
||||||
|
point1 = points[point1_idx]
|
||||||
|
point2 = points[point2_idx]
|
||||||
|
x1, y1 = int(point1[0]), int(point1[1])
|
||||||
|
x2, y2 = int(point2[0]), int(point2[1])
|
||||||
|
cv2.line(canvas, (x1, y1), (x2, y2), connection_colors[i], stickwidth)
|
||||||
|
|
||||||
|
# draw points for joints that have connections
|
||||||
|
joints_in_use = set()
|
||||||
|
for connection in connetions:
|
||||||
|
joints_in_use.add(connection[0])
|
||||||
|
joints_in_use.add(connection[1])
|
||||||
|
|
||||||
|
for joint_idx in joints_in_use:
|
||||||
|
if joint_idx >= len(points):
|
||||||
|
continue
|
||||||
|
x, y = points[joint_idx][0:2]
|
||||||
|
x, y = int(x), int(y)
|
||||||
|
# Use the color from the first connection involving this joint
|
||||||
|
joint_color = [180, 180, 180] # default grey
|
||||||
|
for i, connection in enumerate(connetions):
|
||||||
|
if connection[0] == joint_idx or connection[1] == joint_idx:
|
||||||
|
joint_color = connection_colors[i]
|
||||||
|
break
|
||||||
|
cv2.circle(canvas, (x, y), r, joint_color, thickness=-1)
|
||||||
|
|
||||||
|
return canvas
|
||||||
|
|||||||
+78
-21
@@ -1,10 +1,7 @@
|
|||||||
import os
|
import os
|
||||||
import torch
|
import torch
|
||||||
import gc
|
|
||||||
from ..utils import log, dict_to_device
|
from ..utils import log, dict_to_device
|
||||||
import numpy as np
|
import numpy as np
|
||||||
from accelerate import init_empty_weights
|
|
||||||
from accelerate.utils import set_module_tensor_to_device
|
|
||||||
|
|
||||||
import comfy.model_management as mm
|
import comfy.model_management as mm
|
||||||
from comfy.utils import load_torch_file
|
from comfy.utils import load_torch_file
|
||||||
@@ -17,7 +14,27 @@ offload_device = mm.unet_offload_device()
|
|||||||
local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript")
|
local_model_path = os.path.join(folder_paths.models_dir, "nlf", "nlf_l_multi_0.3.2.torchscript")
|
||||||
|
|
||||||
from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
|
from .motion4d import SMPL_VQVAE, VectorQuantizer, Encoder, Decoder
|
||||||
from .mtv import prepare_motion_embeddings
|
|
||||||
|
def check_jit_script_function():
|
||||||
|
if torch.jit.script.__name__ != "script":
|
||||||
|
# Get more details about what modified it
|
||||||
|
module = torch.jit.script.__module__
|
||||||
|
qualname = getattr(torch.jit.script, '__qualname__', 'unknown')
|
||||||
|
code_file = None
|
||||||
|
try:
|
||||||
|
code_file = torch.jit.script.__code__.co_filename
|
||||||
|
code_line = torch.jit.script.__code__.co_firstlineno
|
||||||
|
log.warning(f"torch.jit.script has been modified by another custom node.\n"
|
||||||
|
f" Function name: {torch.jit.script.__name__}\n"
|
||||||
|
f" Module: {module}\n"
|
||||||
|
f" Qualified name: {qualname}\n"
|
||||||
|
f" Defined in: {code_file}:{code_line}\n"
|
||||||
|
f"This may cause issues with the NLF model.")
|
||||||
|
except:
|
||||||
|
log.warning("--------------------------------")
|
||||||
|
log.warning(f"torch.jit.script function is: {torch.jit.script.__name__} from module {module}, "
|
||||||
|
f"this has been modified by another custom node. This may cause issues with the NLF model.")
|
||||||
|
log.warning("--------------------------------")
|
||||||
|
|
||||||
class DownloadAndLoadNLFModel:
|
class DownloadAndLoadNLFModel:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -30,6 +47,9 @@ class DownloadAndLoadNLFModel:
|
|||||||
],
|
],
|
||||||
)
|
)
|
||||||
},
|
},
|
||||||
|
"optional": {
|
||||||
|
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("NLFMODEL",)
|
RETURN_TYPES = ("NLFMODEL",)
|
||||||
@@ -37,7 +57,9 @@ class DownloadAndLoadNLFModel:
|
|||||||
FUNCTION = "loadmodel"
|
FUNCTION = "loadmodel"
|
||||||
CATEGORY = "WanVideoWrapper"
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
def loadmodel(self, url):
|
def loadmodel(self, url, warmup=True):
|
||||||
|
|
||||||
|
check_jit_script_function()
|
||||||
|
|
||||||
if not os.path.exists(local_model_path):
|
if not os.path.exists(local_model_path):
|
||||||
log.info(f"Downloading NLF model to: {local_model_path}")
|
log.info(f"Downloading NLF model to: {local_model_path}")
|
||||||
@@ -52,6 +74,20 @@ class DownloadAndLoadNLFModel:
|
|||||||
|
|
||||||
model = torch.jit.load(local_model_path).eval()
|
model = torch.jit.load(local_model_path).eval()
|
||||||
|
|
||||||
|
if warmup:
|
||||||
|
log.info("Warming up NLF model...")
|
||||||
|
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
|
||||||
|
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
|
||||||
|
try:
|
||||||
|
for _ in range(2):
|
||||||
|
_ = model.detect_smpl_batched(dummy_input)
|
||||||
|
finally:
|
||||||
|
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
|
||||||
|
|
||||||
|
log.info("NLF model warmed up")
|
||||||
|
|
||||||
|
model = model.to(offload_device)
|
||||||
|
|
||||||
return (model,)
|
return (model,)
|
||||||
|
|
||||||
class LoadNLFModel:
|
class LoadNLFModel:
|
||||||
@@ -61,6 +97,9 @@ class LoadNLFModel:
|
|||||||
"required": {
|
"required": {
|
||||||
"path": ("STRING", {"default": local_model_path}),
|
"path": ("STRING", {"default": local_model_path}),
|
||||||
},
|
},
|
||||||
|
"optional": {
|
||||||
|
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
|
||||||
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("NLFMODEL",)
|
RETURN_TYPES = ("NLFMODEL",)
|
||||||
@@ -68,8 +107,22 @@ class LoadNLFModel:
|
|||||||
FUNCTION = "loadmodel"
|
FUNCTION = "loadmodel"
|
||||||
CATEGORY = "WanVideoWrapper"
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
def loadmodel(self, path):
|
def loadmodel(self, path, warmup=True):
|
||||||
model = torch.jit.load(path).eval()
|
check_jit_script_function()
|
||||||
|
model = torch.jit.load(path, map_location="cpu").eval()
|
||||||
|
|
||||||
|
if warmup:
|
||||||
|
log.info("Warming up NLF model...")
|
||||||
|
dummy_input = torch.zeros(1, 3, 256, 256, device=device)
|
||||||
|
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
|
||||||
|
try:
|
||||||
|
for _ in range(2):
|
||||||
|
_ = model.detect_smpl_batched(dummy_input)
|
||||||
|
finally:
|
||||||
|
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
|
||||||
|
log.info("NLF model warmed up")
|
||||||
|
|
||||||
|
model = model.to(offload_device)
|
||||||
|
|
||||||
return model,
|
return model,
|
||||||
|
|
||||||
@@ -131,15 +184,6 @@ class MTVCrafterEncodePoses:
|
|||||||
|
|
||||||
def encode(self, vqvae, poses):
|
def encode(self, vqvae, poses):
|
||||||
|
|
||||||
# import pickle
|
|
||||||
# with open(os.path.join(script_directory, "data", "sampled_data.pkl"), 'rb') as f:
|
|
||||||
# data_list = pickle.load(f)
|
|
||||||
# if not isinstance(data_list, list):
|
|
||||||
# data_list = [data_list]
|
|
||||||
# print(data_list)
|
|
||||||
|
|
||||||
# smpl_poses = data_list[1]['pose']
|
|
||||||
|
|
||||||
global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3)
|
global_mean = np.load(os.path.join(script_directory, "data", "mean.npy")) #global_mean.shape: (24, 3)
|
||||||
global_std = np.load(os.path.join(script_directory, "data", "std.npy"))
|
global_std = np.load(os.path.join(script_directory, "data", "std.npy"))
|
||||||
|
|
||||||
@@ -182,9 +226,16 @@ class NLFPredict:
|
|||||||
|
|
||||||
def predict(self, model, images):
|
def predict(self, model, images):
|
||||||
|
|
||||||
model.to(device)
|
check_jit_script_function()
|
||||||
|
model = model.to(device)
|
||||||
|
|
||||||
|
jit_profiling_prev_state = torch._C._jit_set_profiling_executor(True)
|
||||||
|
try:
|
||||||
pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
|
pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
|
||||||
model.to(offload_device)
|
finally:
|
||||||
|
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
|
||||||
|
|
||||||
|
model = model.to(offload_device)
|
||||||
|
|
||||||
pred = dict_to_device(pred, offload_device)
|
pred = dict_to_device(pred, offload_device)
|
||||||
|
|
||||||
@@ -208,6 +259,11 @@ class DrawNLFPoses:
|
|||||||
"width": ("INT", {"default": 512}),
|
"width": ("INT", {"default": 512}),
|
||||||
"height": ("INT", {"default": 512}),
|
"height": ("INT", {"default": 512}),
|
||||||
},
|
},
|
||||||
|
"optional": {
|
||||||
|
"stick_width": ("FLOAT", {"default": 4.0, "min": 0.0, "max": 1000.0, "step": 0.01, "tooltip": "Stick width multiplier"}),
|
||||||
|
"point_radius": ("INT", {"default": 5, "min": 1, "max": 10, "step": 1, "tooltip": "Point radius for drawing the pose"}),
|
||||||
|
"style": (["original", "scail"], {"default": "original", "tooltip": "style of the pose drawing"}),
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
RETURN_TYPES = ("IMAGE", )
|
RETURN_TYPES = ("IMAGE", )
|
||||||
@@ -215,14 +271,15 @@ class DrawNLFPoses:
|
|||||||
FUNCTION = "predict"
|
FUNCTION = "predict"
|
||||||
CATEGORY = "WanVideoWrapper"
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
def predict(self, poses, width, height):
|
def predict(self, poses, width, height, stick_width=1.0, point_radius=2, style="original"):
|
||||||
from .draw_pose import get_control_conditions
|
from .draw_pose import get_control_conditions
|
||||||
print(type(poses))
|
|
||||||
if isinstance(poses, dict):
|
if isinstance(poses, dict):
|
||||||
pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
|
pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
|
||||||
else:
|
else:
|
||||||
pose_input = poses
|
pose_input = poses
|
||||||
control_conditions = get_control_conditions(pose_input, height, width)
|
|
||||||
|
control_conditions = get_control_conditions(pose_input, height, width, stick_width=stick_width, point_radius=point_radius, style=style)
|
||||||
|
|
||||||
return (control_conditions,)
|
return (control_conditions,)
|
||||||
|
|
||||||
|
|||||||
Binary file not shown.
Binary file not shown.
Binary file not shown.
File diff suppressed because one or more lines are too long
@@ -0,0 +1,96 @@
|
|||||||
|
import torch
|
||||||
|
from ..utils import log
|
||||||
|
import comfy.model_management as mm
|
||||||
|
|
||||||
|
device = mm.get_torch_device()
|
||||||
|
offload_device = mm.unet_offload_device()
|
||||||
|
|
||||||
|
class WanVideoAddSCAILReferenceEmbeds:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||||
|
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||||
|
"ref_image": ("IMAGE",),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
|
||||||
|
},
|
||||||
|
"optional": {
|
||||||
|
"clip_embeds": ("WANVIDIMAGE_CLIPEMBEDS", {"tooltip": "Clip vision encoded image"}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||||
|
RETURN_NAMES = ("image_embeds",)
|
||||||
|
FUNCTION = "add"
|
||||||
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
|
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, clip_embeds=None):
|
||||||
|
updated = dict(embeds)
|
||||||
|
|
||||||
|
vae.to(device)
|
||||||
|
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
|
||||||
|
ref_latent = vae.encode([ref_image_in], device, tiled=False)[0]
|
||||||
|
log.info(f"SCAIL ref_latent shape: {ref_latent.shape}")
|
||||||
|
|
||||||
|
ref_mask = torch.ones_like(ref_latent[:4])
|
||||||
|
ref_latent = torch.cat([ref_latent, ref_mask], dim=0)
|
||||||
|
vae.to(offload_device)
|
||||||
|
|
||||||
|
updated.setdefault("scail_embeds", {})
|
||||||
|
updated["scail_embeds"]["ref_latent_pos"] = ref_latent * strength
|
||||||
|
updated["scail_embeds"]["ref_latent_neg"] = torch.zeros_like(ref_latent)
|
||||||
|
updated["scail_embeds"]["ref_start_percent"] = start_percent
|
||||||
|
updated["scail_embeds"]["ref_end_percent"] = end_percent
|
||||||
|
updated["clip_context"] = clip_embeds.get("clip_embeds", None) if clip_embeds is not None else None
|
||||||
|
|
||||||
|
return (updated,)
|
||||||
|
|
||||||
|
class WanVideoAddSCAILPoseEmbeds:
|
||||||
|
@classmethod
|
||||||
|
def INPUT_TYPES(s):
|
||||||
|
return {"required": {
|
||||||
|
"embeds": ("WANVIDIMAGE_EMBEDS",),
|
||||||
|
"vae": ("WANVAE", {"tooltip": "VAE model"}),
|
||||||
|
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
|
||||||
|
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
|
||||||
|
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
|
||||||
|
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
|
||||||
|
RETURN_NAMES = ("image_embeds",)
|
||||||
|
FUNCTION = "add"
|
||||||
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
|
||||||
|
def add(self, embeds, vae, pose_images, strength, start_percent=0.0, end_percent=1.0):
|
||||||
|
updated = dict(embeds)
|
||||||
|
|
||||||
|
vae.to(device)
|
||||||
|
pose_images_in = (pose_images[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
|
||||||
|
pose_latent = vae.encode([pose_images_in], device, tiled=False)[0]
|
||||||
|
pose_mask = torch.ones_like(pose_latent[:4])
|
||||||
|
pose_latent = torch.cat([pose_latent, pose_mask], dim=0)
|
||||||
|
log.info(f"SCAIL pose_latent shape: {pose_latent.shape}")
|
||||||
|
|
||||||
|
vae.to(offload_device)
|
||||||
|
|
||||||
|
updated.setdefault("scail_embeds", {})
|
||||||
|
updated["scail_embeds"]["pose_latent"] = pose_latent
|
||||||
|
updated["scail_embeds"]["pose_strength"] = strength
|
||||||
|
updated["scail_embeds"]["pose_start_percent"] = start_percent
|
||||||
|
updated["scail_embeds"]["pose_end_percent"] = end_percent
|
||||||
|
|
||||||
|
return (updated,)
|
||||||
|
|
||||||
|
|
||||||
|
NODE_CLASS_MAPPINGS = {
|
||||||
|
"WanVideoAddSCAILPoseEmbeds": WanVideoAddSCAILPoseEmbeds,
|
||||||
|
"WanVideoAddSCAILReferenceEmbeds": WanVideoAddSCAILReferenceEmbeds,
|
||||||
|
}
|
||||||
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
|
"WanVideoAddSCAILReferenceEmbeds": "WanVideo Add SCAIL Reference Embeds",
|
||||||
|
"WanVideoAddSCAILPoseEmbeds": "WanVideo Add SCAIL Pose Embeds",
|
||||||
|
}
|
||||||
@@ -160,6 +160,7 @@ class WanMove_native:
|
|||||||
RETURN_NAMES = ("positive", "tracks")
|
RETURN_NAMES = ("positive", "tracks")
|
||||||
FUNCTION = "patchcond"
|
FUNCTION = "patchcond"
|
||||||
CATEGORY = "WanVideoWrapper"
|
CATEGORY = "WanVideoWrapper"
|
||||||
|
DEPRECATED = True
|
||||||
|
|
||||||
def patchcond(self, positive, track_coords, track_mask=None):
|
def patchcond(self, positive, track_coords, track_mask=None):
|
||||||
|
|
||||||
|
|||||||
@@ -47,6 +47,7 @@ OPTIONAL_MODULES = [
|
|||||||
(".steadydancer.nodes", "SteadyDancer"),
|
(".steadydancer.nodes", "SteadyDancer"),
|
||||||
(".onetoall.nodes", "OneToAll"),
|
(".onetoall.nodes", "OneToAll"),
|
||||||
(".WanMove.nodes", "WanMove"),
|
(".WanMove.nodes", "WanMove"),
|
||||||
|
(".SCAIL.nodes", "SCAIL"),
|
||||||
]
|
]
|
||||||
|
|
||||||
def register_nodes(module_path: str, name: str, optional: bool) -> None:
|
def register_nodes(module_path: str, name: str, optional: bool) -> None:
|
||||||
|
|||||||
@@ -857,7 +857,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
|
|||||||
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
|
||||||
except Exception:
|
except Exception:
|
||||||
vace_block_idx = None
|
vace_block_idx = None
|
||||||
elif name.startswith("blocks.") and "face" not in name:
|
elif name.startswith("blocks.") and "face" not in name and "controlnet_blocks." not in name:
|
||||||
try:
|
try:
|
||||||
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
block_idx = int(name.split("blocks.")[1].split(".")[0])
|
||||||
except Exception:
|
except Exception:
|
||||||
@@ -1223,10 +1223,7 @@ class WanVideoModelLoader:
|
|||||||
model_type = "no_cross_attn" #minimaxremover
|
model_type = "no_cross_attn" #minimaxremover
|
||||||
elif "model_type.Wan2_1-FLF2V-14B-720P" in sd or "img_emb.emb_pos" in sd or "flf2v" in model.lower():
|
elif "model_type.Wan2_1-FLF2V-14B-720P" in sd or "img_emb.emb_pos" in sd or "flf2v" in model.lower():
|
||||||
model_type = "fl2v"
|
model_type = "fl2v"
|
||||||
elif in_channels in [36, 48]:
|
if "blocks.0.cross_attn.k_img.weight" in sd:
|
||||||
if "blocks.0.cross_attn.k_img.weight" not in sd:
|
|
||||||
model_type = "t2v"
|
|
||||||
else:
|
|
||||||
model_type = "i2v"
|
model_type = "i2v"
|
||||||
elif in_channels == 16:
|
elif in_channels == 16:
|
||||||
model_type = "t2v"
|
model_type = "t2v"
|
||||||
@@ -1512,6 +1509,11 @@ class WanVideoModelLoader:
|
|||||||
FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU())
|
FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU())
|
||||||
transformer.condition_embedding_align = PoseRefNetNoBNV3(in_channels_x=16, in_channels_c=16, hidden_dim=128, num_heads=8) # Frame-wise Attention Alignment Unit
|
transformer.condition_embedding_align = PoseRefNetNoBNV3(in_channels_x=16, in_channels_c=16, hidden_dim=128, num_heads=8) # Frame-wise Attention Alignment Unit
|
||||||
|
|
||||||
|
# SCAIL
|
||||||
|
if "patch_embedding_pose.weight" in sd:
|
||||||
|
log.info("SCAIL model detected, patching model...")
|
||||||
|
pose_dim = sd["patch_embedding_pose.weight"].shape[1]
|
||||||
|
transformer.patch_embedding_pose = nn.Conv3d(pose_dim, dim, kernel_size=patch_size, stride=patch_size)
|
||||||
|
|
||||||
if "image_to_cond.conv_in.bias" in sd:
|
if "image_to_cond.conv_in.bias" in sd:
|
||||||
# One-to-all
|
# One-to-all
|
||||||
|
|||||||
@@ -1219,6 +1219,17 @@ class WanVideoSampler:
|
|||||||
latents_to_not_step = prev_latents.shape[1]
|
latents_to_not_step = prev_latents.shape[1]
|
||||||
one_to_all_data["num_latent_frames_to_replace"] = latents_to_not_step
|
one_to_all_data["num_latent_frames_to_replace"] = latents_to_not_step
|
||||||
|
|
||||||
|
# SCAIL
|
||||||
|
scail_embeds = image_embeds.get("scail_embeds", None)
|
||||||
|
scail_data = None
|
||||||
|
if scail_embeds is not None:
|
||||||
|
log.info("Using SCAIL embeddings:")
|
||||||
|
for k, v in scail_embeds.items():
|
||||||
|
log.info(f" {k}: {v.shape if isinstance(v, torch.Tensor) else v}")
|
||||||
|
scail_data = scail_embeds.copy()
|
||||||
|
scail_data = dict_to_device(scail_data, device, dtype)
|
||||||
|
|
||||||
|
|
||||||
# WanMove
|
# WanMove
|
||||||
wanmove_embeds = None
|
wanmove_embeds = None
|
||||||
if image_cond is not None:
|
if image_cond is not None:
|
||||||
@@ -1468,6 +1479,16 @@ class WanVideoSampler:
|
|||||||
if background_latents is not None or foreground_latents is not None:
|
if background_latents is not None or foreground_latents is not None:
|
||||||
z = torch.cat([z, foreground_latents.to(z), background_latents.to(z)], dim=0)
|
z = torch.cat([z, foreground_latents.to(z), background_latents.to(z)], dim=0)
|
||||||
|
|
||||||
|
scail_data_in = None
|
||||||
|
if scail_data is not None:
|
||||||
|
ref_concat_mask = torch.zeros_like(z[:4])
|
||||||
|
z = torch.cat([z, ref_concat_mask])
|
||||||
|
if context_window is not None:
|
||||||
|
scail_data_in = scail_data.copy()
|
||||||
|
scail_data_in["pose_latent"] = scail_data["pose_latent"][:, context_window]
|
||||||
|
else:
|
||||||
|
scail_data_in = scail_data
|
||||||
|
|
||||||
if wanmove_embeds is not None and context_window is not None:
|
if wanmove_embeds is not None and context_window is not None:
|
||||||
image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
|
image_cond_input = replace_feature(image_cond_input.unsqueeze(0), track_pos[:, context_window].unsqueeze(0), wanmove_embeds.get("strength", 1.0))[0]
|
||||||
|
|
||||||
@@ -1530,6 +1551,7 @@ class WanVideoSampler:
|
|||||||
"sdancer_input": sdancer_input, # SteadyDancer input
|
"sdancer_input": sdancer_input, # SteadyDancer input
|
||||||
"one_to_all_input": one_to_all_data, # One-to-All input
|
"one_to_all_input": one_to_all_data, # One-to-All input
|
||||||
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
|
"one_to_all_controlnet_strength": one_to_all_data["controlnet_strength"] if one_to_all_data is not None else 0.0,
|
||||||
|
"scail_input": scail_data_in, # SCAIL input
|
||||||
}
|
}
|
||||||
|
|
||||||
batch_size = 1
|
batch_size = 1
|
||||||
|
|||||||
+113
-39
@@ -2,6 +2,7 @@
|
|||||||
import math
|
import math
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
from einops import repeat, rearrange
|
from einops import repeat, rearrange
|
||||||
from ...enhance_a_video.enhance import get_feta_scores
|
from ...enhance_a_video.enhance import get_feta_scores
|
||||||
import time
|
import time
|
||||||
@@ -734,6 +735,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
|||||||
context(Tensor): Shape [B, L2, C]
|
context(Tensor): Shape [B, L2, C]
|
||||||
"""
|
"""
|
||||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||||
|
|
||||||
# compute query
|
# compute query
|
||||||
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype)
|
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype)
|
||||||
|
|
||||||
@@ -2125,7 +2127,9 @@ class WanModel(torch.nn.Module):
|
|||||||
return x.add(residual_out, alpha=strength)
|
return x.add(residual_out, alpha=strength)
|
||||||
|
|
||||||
|
|
||||||
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
|
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, ref_frame_shape=None, pose_frame_shape=None,
|
||||||
|
steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
|
||||||
|
|
||||||
patch_size = self.patch_size
|
patch_size = self.patch_size
|
||||||
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
|
||||||
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
|
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
|
||||||
@@ -2138,40 +2142,70 @@ class WanModel(torch.nn.Module):
|
|||||||
if steps_w is None:
|
if steps_w is None:
|
||||||
steps_w = w_len
|
steps_w = w_len
|
||||||
|
|
||||||
|
# Main frames position IDs
|
||||||
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
|
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
|
||||||
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
|
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
|
||||||
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
|
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
|
||||||
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
|
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
|
||||||
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
|
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
|
||||||
if attn_cond_shape is not None:
|
|
||||||
F_cond, H_cond, W_cond = attn_cond_shape[2], attn_cond_shape[3], attn_cond_shape[4]
|
segments = [img_ids] # Start with main frames
|
||||||
|
|
||||||
|
# Reference frames position IDs
|
||||||
|
if ref_frame_shape is not None:
|
||||||
|
F_cond, H_cond, W_cond = ref_frame_shape[-3], ref_frame_shape[-2], ref_frame_shape[-1]
|
||||||
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
|
cond_f_len = ((F_cond + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||||
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
|
cond_h_len = ((H_cond + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||||
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
|
cond_w_len = ((W_cond + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||||
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=device, dtype=dtype)
|
cond_img_ids = torch.zeros((cond_f_len, cond_h_len, cond_w_len, 3), device=device, dtype=dtype)
|
||||||
|
|
||||||
#shift
|
|
||||||
shift_f_size = 81 # Default value
|
|
||||||
shift_f = False
|
|
||||||
if shift_f:
|
|
||||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(shift_f_size, shift_f_size + cond_f_len - 1,steps=cond_f_len, device=device, dtype=dtype).reshape(-1, 1, 1)
|
|
||||||
else:
|
|
||||||
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=device, dtype=dtype).reshape(-1, 1, 1)
|
cond_img_ids[:, :, :, 0] = cond_img_ids[:, :, :, 0] + torch.linspace(0, cond_f_len - 1, steps=cond_f_len, device=device, dtype=dtype).reshape(-1, 1, 1)
|
||||||
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=device, dtype=dtype).reshape(1, -1, 1)
|
cond_img_ids[:, :, :, 1] = cond_img_ids[:, :, :, 1] + torch.linspace(h_len, h_len + cond_h_len - 1, steps=cond_h_len, device=device, dtype=dtype).reshape(1, -1, 1)
|
||||||
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=device, dtype=dtype).reshape(1, 1, -1)
|
cond_img_ids[:, :, :, 2] = cond_img_ids[:, :, :, 2] + torch.linspace(w_len, w_len + cond_w_len - 1, steps=cond_w_len, device=device, dtype=dtype).reshape(1, 1, -1)
|
||||||
|
|
||||||
# Combine original and conditional position ids
|
segments.insert(0, cond_img_ids.reshape(1, -1, cond_img_ids.shape[-1])) # Ref frames come first
|
||||||
#img_ids = repeat(img_ids, "t h w c -> b (t h w) c", b=1)
|
|
||||||
#cond_img_ids = repeat(cond_img_ids, "t h w c -> b (t h w) c", b=1)
|
|
||||||
cond_img_ids = cond_img_ids.reshape(1, -1, cond_img_ids.shape[-1])
|
|
||||||
combined_img_ids = torch.cat([img_ids, cond_img_ids], dim=1)
|
|
||||||
|
|
||||||
# Generate RoPE frequencies for the combined positions
|
# Pose frames position IDs
|
||||||
|
if pose_frame_shape is not None:
|
||||||
|
F_pose, H_pose, W_pose = pose_frame_shape[-3], pose_frame_shape[-2], pose_frame_shape[-1]
|
||||||
|
|
||||||
|
downscale = H_pose != h
|
||||||
|
pose_f_len_full = ((F_pose + (self.patch_size[0] // 2)) // self.patch_size[0])
|
||||||
|
pose_h_len_full = (((H_pose * (2 if downscale else 1)) + (self.patch_size[1] // 2)) // self.patch_size[1]) # 2x height
|
||||||
|
pose_w_len_full = (((W_pose * (2 if downscale else 1)) + (self.patch_size[2] // 2)) // self.patch_size[2]) # 2x width
|
||||||
|
|
||||||
|
pose_img_ids = torch.zeros((pose_f_len_full, pose_h_len_full, pose_w_len_full, 3), device=device, dtype=dtype)
|
||||||
|
global_h_offset, global_w_offset = 0, 120 # global spatial offset to separate pose from main frames spatially (SCAIL uses 120 as offset)
|
||||||
|
pose_img_ids[:, :, :, 0] = pose_img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (pose_f_len_full - 1), steps=pose_f_len_full, device=device, dtype=dtype).reshape(-1, 1, 1)
|
||||||
|
pose_img_ids[:, :, :, 1] = pose_img_ids[:, :, :, 1] + torch.linspace(global_h_offset + freq_offset, global_h_offset + pose_h_len_full - 1, steps=pose_h_len_full, device=device, dtype=dtype).reshape(1, -1, 1)
|
||||||
|
pose_img_ids[:, :, :, 2] = pose_img_ids[:, :, :, 2] + torch.linspace(global_w_offset + freq_offset, global_w_offset + pose_w_len_full - 1, steps=pose_w_len_full, device=device, dtype=dtype).reshape(1, 1, -1)
|
||||||
|
|
||||||
|
segments.append(pose_img_ids.reshape(1, -1, pose_img_ids.shape[-1]))
|
||||||
|
|
||||||
|
combined_img_ids = torch.cat(segments, dim=1)
|
||||||
freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2)
|
freqs = self.rope_embedder(combined_img_ids, ntk_alphas).movedim(1, 2)
|
||||||
else:
|
|
||||||
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
|
# Downsample pose frequencies to match actual pose input resolution
|
||||||
|
if pose_frame_shape is not None and downscale:
|
||||||
|
pose_h_len_actual = ((H_pose + (self.patch_size[1] // 2)) // self.patch_size[1])
|
||||||
|
pose_w_len_actual = ((W_pose + (self.patch_size[2] // 2)) // self.patch_size[2])
|
||||||
|
|
||||||
|
pose_start_idx = freqs.shape[1] - pose_f_len_full * pose_h_len_full * pose_w_len_full
|
||||||
|
main_freqs, pose_freqs = freqs[:, :pose_start_idx], freqs[:, pose_start_idx:]
|
||||||
|
|
||||||
|
B, _, heads, dim, _, _ = pose_freqs.shape
|
||||||
|
# Reshape and pool: (B, L, heads, dim, 2, 2) -> pool H,W -> (B, L', heads, dim, 2, 2)
|
||||||
|
pose_freqs = pose_freqs.reshape(B, pose_f_len_full, pose_h_len_full, pose_w_len_full, heads, dim, 2, 2)
|
||||||
|
pose_freqs = pose_freqs.permute(0, 1, 4, 5, 6, 7, 2, 3).reshape(-1, pose_h_len_full, pose_w_len_full)
|
||||||
|
pose_freqs = F.avg_pool2d(pose_freqs, kernel_size=2, stride=2)
|
||||||
|
pose_freqs = pose_freqs.reshape(B, pose_f_len_full, heads, dim, 2, 2, pose_h_len_actual, pose_w_len_actual)
|
||||||
|
pose_freqs = pose_freqs.permute(0, 1, 6, 7, 2, 3, 4, 5).reshape(B, -1, heads, dim, 2, 2)
|
||||||
|
|
||||||
|
freqs = torch.cat([main_freqs, pose_freqs], dim=1)
|
||||||
|
|
||||||
return freqs
|
return freqs
|
||||||
|
|
||||||
|
|
||||||
def forward(
|
def forward(
|
||||||
self, x, t, context, seq_len,
|
self, x, t, context, seq_len,
|
||||||
is_uncond=False,
|
is_uncond=False,
|
||||||
@@ -2212,7 +2246,8 @@ class WanModel(torch.nn.Module):
|
|||||||
num_cond_latents=None,
|
num_cond_latents=None,
|
||||||
add_text_emb=None,
|
add_text_emb=None,
|
||||||
sdancer_input=None, # SteadyDancer
|
sdancer_input=None, # SteadyDancer
|
||||||
one_to_all_input=None, one_to_all_controlnet_strength=0.0 # One-to-All
|
one_to_all_input=None, one_to_all_controlnet_strength=0.0, # One-to-All
|
||||||
|
scail_input=None, # SCAIL pose
|
||||||
):
|
):
|
||||||
r"""
|
r"""
|
||||||
Forward pass through the diffusion model
|
Forward pass through the diffusion model
|
||||||
@@ -2296,6 +2331,7 @@ class WanModel(torch.nn.Module):
|
|||||||
freqs = freqs.to(device)
|
freqs = freqs.to(device)
|
||||||
|
|
||||||
_, F, H, W = x[0].shape
|
_, F, H, W = x[0].shape
|
||||||
|
ref_frame_shape = pose_frame_shape = None
|
||||||
|
|
||||||
sdancer_enabled = False
|
sdancer_enabled = False
|
||||||
if sdancer_input is not None and sdancer_input['start_percent'] <= current_step_percentage <= sdancer_input['end_percent']:
|
if sdancer_input is not None and sdancer_input['start_percent'] <= current_step_percentage <= sdancer_input['end_percent']:
|
||||||
@@ -2348,12 +2384,24 @@ class WanModel(torch.nn.Module):
|
|||||||
token_replace_start = (H // self.patch_size[1]) * (W // self.patch_size[2]) # skip first (ref) frame
|
token_replace_start = (H // self.patch_size[1]) * (W // self.patch_size[2]) # skip first (ref) frame
|
||||||
replace_token_num = num_latent_frames_to_replace * token_replace_start # zero next frames
|
replace_token_num = num_latent_frames_to_replace * token_replace_start # zero next frames
|
||||||
|
|
||||||
|
# SCAIL ref
|
||||||
|
if scail_input is not None:
|
||||||
|
ref_latent = scail_input.get("ref_latent_pos", None) if not is_uncond else scail_input.get("ref_latent_neg", None)
|
||||||
|
if ref_latent is not None and scail_input['ref_start_percent'] <= current_step_percentage <= scail_input['ref_end_percent']:
|
||||||
|
x = [torch.cat([v, u], dim=1) for v, u in zip([ref_latent], x)]
|
||||||
|
seq_len += math.ceil((ref_latent.shape[-1] * ref_latent.shape[-2]) / 4 * ref_latent.shape[-3])
|
||||||
|
F += 1
|
||||||
|
prefix_frames = 1
|
||||||
|
suffix_frames += 1
|
||||||
|
|
||||||
#uni3c controlnet
|
#uni3c controlnet
|
||||||
if uni3c_data is not None:
|
if uni3c_data is not None:
|
||||||
render_latent = uni3c_data["render_latent"].to(self.base_dtype)
|
render_latent = uni3c_data["render_latent"].to(self.base_dtype)
|
||||||
hidden_states = x[0].unsqueeze(0).clone().float()
|
hidden_states = x[0].unsqueeze(0).clone().float()
|
||||||
if hidden_states.shape[1] == 16: #T2V work around
|
if hidden_states.shape[1] == 16: #T2V work around
|
||||||
hidden_states = torch.cat([hidden_states, torch.zeros_like(hidden_states[:, :4])], dim=1)
|
hidden_states = torch.cat([hidden_states, torch.zeros_like(hidden_states[:, :4])], dim=1)
|
||||||
|
if hidden_states.shape[2] != render_latent.shape[2]: # temporal resample
|
||||||
|
render_latent = nn.functional.interpolate(render_latent, size=(hidden_states.shape[2], hidden_states.shape[3], hidden_states.shape[4]), mode='trilinear', align_corners=False)
|
||||||
render_latent = torch.cat([hidden_states[:, :20], render_latent], dim=1)
|
render_latent = torch.cat([hidden_states[:, :20], render_latent], dim=1)
|
||||||
|
|
||||||
# SteadyDancer
|
# SteadyDancer
|
||||||
@@ -2424,6 +2472,17 @@ class WanModel(torch.nn.Module):
|
|||||||
original_grid_sizes = grid_sizes.clone()
|
original_grid_sizes = grid_sizes.clone()
|
||||||
x = [u.flatten(2).transpose(1, 2) for u in x]
|
x = [u.flatten(2).transpose(1, 2) for u in x]
|
||||||
self.original_seq_len = x[0].shape[1]
|
self.original_seq_len = x[0].shape[1]
|
||||||
|
|
||||||
|
# SCAIL pose
|
||||||
|
if scail_input is not None:
|
||||||
|
scail_pose_latents = scail_input.get("pose_latent", None)
|
||||||
|
if scail_pose_latents is not None and scail_input['pose_start_percent'] <= current_step_percentage <= scail_input['pose_end_percent']:
|
||||||
|
scail_x = [self.patch_embedding_pose(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in [scail_pose_latents]]
|
||||||
|
scail_x = [u.flatten(2).transpose(1, 2) * scail_input.get("pose_strength", 1) for u in scail_x]
|
||||||
|
x = [torch.cat([u, v], dim=1) for u, v in zip(x, scail_x)]
|
||||||
|
seq_len += scail_x[0].shape[1]
|
||||||
|
pose_frame_shape = scail_pose_latents.shape
|
||||||
|
|
||||||
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
|
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
|
||||||
assert seq_lens.max() <= seq_len, f"max seq len {seq_lens.max()} exceeds provided seq_len {seq_len}"
|
assert seq_lens.max() <= seq_len, f"max seq len {seq_lens.max()} exceeds provided seq_len {seq_len}"
|
||||||
|
|
||||||
@@ -2435,9 +2494,8 @@ class WanModel(torch.nn.Module):
|
|||||||
add_cond = self.add_conv_in(add_cond.to(self.add_conv_in.weight.dtype)).to(x[0].dtype)
|
add_cond = self.add_conv_in(add_cond.to(self.add_conv_in.weight.dtype)).to(x[0].dtype)
|
||||||
add_cond = add_cond.flatten(2).transpose(1, 2)
|
add_cond = add_cond.flatten(2).transpose(1, 2)
|
||||||
x[0] = x[0] + self.add_proj(add_cond)
|
x[0] = x[0] + self.add_proj(add_cond)
|
||||||
attn_cond_shape = None
|
|
||||||
if attn_cond is not None:
|
if attn_cond is not None:
|
||||||
attn_cond_shape = attn_cond.shape
|
ref_frame_shape = attn_cond.shape
|
||||||
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
grid_sizes = torch.stack([torch.tensor([u[0] + 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||||
attn_cond = self.attn_conv_in(attn_cond.to(self.attn_conv_in.weight.dtype)).to(x[0].dtype)
|
attn_cond = self.attn_conv_in(attn_cond.to(self.attn_conv_in.weight.dtype)).to(x[0].dtype)
|
||||||
attn_cond = attn_cond.flatten(2).transpose(1, 2)
|
attn_cond = attn_cond.flatten(2).transpose(1, 2)
|
||||||
@@ -2490,33 +2548,49 @@ class WanModel(torch.nn.Module):
|
|||||||
x_ip = ip_image_patch.flatten(2).transpose(1, 2) # [B, N, D]
|
x_ip = ip_image_patch.flatten(2).transpose(1, 2) # [B, N, D]
|
||||||
freq_offset = standin_input["freq_offset"]
|
freq_offset = standin_input["freq_offset"]
|
||||||
|
|
||||||
|
# region rope freqs
|
||||||
if freqs is None and "comfy" in self.rope_func: #comfy rope
|
if freqs is None and "comfy" in self.rope_func: #comfy rope
|
||||||
current_shape = (F, H, W)
|
# Create cache key from all relevant parameters
|
||||||
|
cache_key = (
|
||||||
has_cond = attn_cond is not None
|
F, H, W,
|
||||||
|
attn_cond is not None,
|
||||||
|
tuple(ref_frame_shape) if ref_frame_shape is not None else None,
|
||||||
|
tuple(pose_frame_shape) if pose_frame_shape is not None else None,
|
||||||
|
self.rope_embedder.k,
|
||||||
|
tuple(ntk_alphas),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check cache using key comparison
|
||||||
if (self.cached_freqs is not None and
|
if (self.cached_freqs is not None and
|
||||||
self.cached_shape == current_shape and
|
hasattr(self, 'cached_key') and
|
||||||
self.cached_cond == has_cond and
|
self.cached_key == cache_key):
|
||||||
self.cached_rope_k == self.rope_embedder.k and
|
|
||||||
self.cached_ntk_alphas == ntk_alphas
|
|
||||||
):
|
|
||||||
freqs = self.cached_freqs
|
freqs = self.cached_freqs
|
||||||
else:
|
else:
|
||||||
freqs = self.rope_encode_comfy(F, H, W, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
|
log.info("Generating new RoPE frequencies")
|
||||||
|
freqs = self.rope_encode_comfy(
|
||||||
|
F, H, W,
|
||||||
|
freq_offset=freq_offset,
|
||||||
|
ntk_alphas=ntk_alphas,
|
||||||
|
ref_frame_shape=ref_frame_shape,
|
||||||
|
pose_frame_shape=pose_frame_shape,
|
||||||
|
device=x.device,
|
||||||
|
dtype=x.dtype
|
||||||
|
)
|
||||||
|
|
||||||
if s2v_ref_latent is not None:
|
if s2v_ref_latent is not None:
|
||||||
freqs_ref = self.rope_encode_comfy(
|
freqs_ref = self.rope_encode_comfy(
|
||||||
s2v_ref_latent.shape[2],
|
s2v_ref_latent.shape[2],
|
||||||
s2v_ref_latent.shape[3],
|
s2v_ref_latent.shape[3],
|
||||||
s2v_ref_latent.shape[4],
|
s2v_ref_latent.shape[4],
|
||||||
t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
|
t_start=max(30, F + 9),
|
||||||
|
device=x.device,
|
||||||
|
dtype=x.dtype
|
||||||
|
)
|
||||||
freqs = torch.cat([freqs, freqs_ref], dim=1)
|
freqs = torch.cat([freqs, freqs_ref], dim=1)
|
||||||
|
|
||||||
|
# Store cache with key
|
||||||
self.cached_freqs = freqs
|
self.cached_freqs = freqs
|
||||||
self.cached_shape = current_shape
|
self.cached_key = cache_key
|
||||||
self.cached_cond = has_cond
|
|
||||||
self.cached_rope_k = self.rope_embedder.k
|
|
||||||
self.cached_ntk_alphas = ntk_alphas
|
|
||||||
|
|
||||||
# Stand-In RoPE frequencies
|
# Stand-In RoPE frequencies
|
||||||
if x_ip is not None:
|
if x_ip is not None:
|
||||||
@@ -3102,16 +3176,16 @@ class WanModel(torch.nn.Module):
|
|||||||
if self.ref_conv is not None and fun_ref is not None:
|
if self.ref_conv is not None and fun_ref is not None:
|
||||||
fun_ref_length = fun_ref.size(1)
|
fun_ref_length = fun_ref.size(1)
|
||||||
x = x[:, fun_ref_length:]
|
x = x[:, fun_ref_length:]
|
||||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
#grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||||
|
|
||||||
if end_ref_latent is not None:
|
if end_ref_latent is not None:
|
||||||
end_ref_latent_length = end_ref_latent.size(1)
|
end_ref_latent_length = end_ref_latent.size(1)
|
||||||
x = x[:, :-end_ref_latent_length]
|
x = x[:, :-end_ref_latent_length]
|
||||||
grid_sizes = torch.stack([torch.tensor([u[0] - end_ref_latent_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
#grid_sizes = torch.stack([torch.tensor([u[0] - end_ref_latent_frames, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||||
|
|
||||||
if attn_cond is not None:
|
#if attn_cond is not None:
|
||||||
x = x[:, :self.original_seq_len]
|
# x = x[:, :self.original_seq_len]
|
||||||
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
#grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
|
||||||
|
|
||||||
|
|
||||||
x = x[:, :self.original_seq_len]
|
x = x[:, :self.original_seq_len]
|
||||||
|
|||||||
Reference in New Issue
Block a user