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:
+83
-12
@@ -31,7 +31,7 @@ def p3d_to_p2d(point_3d, height, width): # point3d n*1024*3
|
||||
|
||||
def get_pose_images(smpl_data, offset):
|
||||
pose_images = []
|
||||
for data in smpl_data:
|
||||
for data in smpl_data:
|
||||
if isinstance(data, np.ndarray):
|
||||
joints3d = data
|
||||
else:
|
||||
@@ -43,28 +43,33 @@ def get_pose_images(smpl_data, offset):
|
||||
return pose_images
|
||||
|
||||
|
||||
def get_control_conditions(poses, h, w):
|
||||
video_transforms = transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True)
|
||||
def get_control_conditions(poses, h, w, stick_width=1.0, point_radius=2, style="original"):
|
||||
control_images = []
|
||||
for idx, pose in enumerate(poses):
|
||||
canvas = np.zeros(shape=(h, w, 3), dtype=np.uint8)
|
||||
try:
|
||||
joints3d = p3d_to_p2d(pose, h, w)
|
||||
canvas = draw_3d_points(
|
||||
canvas,
|
||||
joints3d[0],
|
||||
stickwidth=int(h / 350),
|
||||
)
|
||||
if style == "original":
|
||||
canvas = draw_3d_points(
|
||||
canvas,
|
||||
joints3d[0],
|
||||
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))
|
||||
# Image.fromarray(resized_canvas).save(f'tmp/{idx}_pose.jpg')
|
||||
control_images.append(resized_canvas)
|
||||
except Exception as e:
|
||||
print("wrong:", e)
|
||||
except Exception:
|
||||
control_images.append(Image.fromarray(canvas))
|
||||
control_pixel_values = np.array(control_images)
|
||||
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
|
||||
|
||||
|
||||
@@ -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])
|
||||
|
||||
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
|
||||
|
||||
+87
-30
@@ -1,23 +1,40 @@
|
||||
import os
|
||||
import torch
|
||||
import gc
|
||||
from ..utils import log, dict_to_device
|
||||
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
|
||||
from comfy.utils import load_torch_file
|
||||
import folder_paths
|
||||
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
script_directory = os.path.dirname(os.path.abspath(__file__))
|
||||
device = mm.get_torch_device()
|
||||
offload_device = mm.unet_offload_device()
|
||||
|
||||
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 .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:
|
||||
@classmethod
|
||||
@@ -30,6 +47,9 @@ class DownloadAndLoadNLFModel:
|
||||
],
|
||||
)
|
||||
},
|
||||
"optional": {
|
||||
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NLFMODEL",)
|
||||
@@ -37,8 +57,10 @@ class DownloadAndLoadNLFModel:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, url):
|
||||
|
||||
def loadmodel(self, url, warmup=True):
|
||||
|
||||
check_jit_script_function()
|
||||
|
||||
if not os.path.exists(local_model_path):
|
||||
log.info(f"Downloading NLF model to: {local_model_path}")
|
||||
import requests
|
||||
@@ -52,6 +74,20 @@ class DownloadAndLoadNLFModel:
|
||||
|
||||
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,)
|
||||
|
||||
class LoadNLFModel:
|
||||
@@ -61,6 +97,9 @@ class LoadNLFModel:
|
||||
"required": {
|
||||
"path": ("STRING", {"default": local_model_path}),
|
||||
},
|
||||
"optional": {
|
||||
"warmup": ("BOOLEAN", {"default": True, "tooltip": "Whether to warmup the model after loading"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("NLFMODEL",)
|
||||
@@ -68,8 +107,22 @@ class LoadNLFModel:
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def loadmodel(self, path):
|
||||
model = torch.jit.load(path).eval()
|
||||
def loadmodel(self, path, warmup=True):
|
||||
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,
|
||||
|
||||
@@ -108,7 +161,7 @@ class LoadVQVAE:
|
||||
frame_upsample_rate=[2.0, 2.0],
|
||||
joint_upsample_rate=[1.0, 1.0]
|
||||
)
|
||||
|
||||
|
||||
vqvae = SMPL_VQVAE(motion_encoder, motion_decoder, motion_quant).to(device)
|
||||
vqvae.load_state_dict(vae_sd, strict=True)
|
||||
|
||||
@@ -131,15 +184,6 @@ class MTVCrafterEncodePoses:
|
||||
|
||||
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_std = np.load(os.path.join(script_directory, "data", "std.npy"))
|
||||
|
||||
@@ -153,7 +197,7 @@ class MTVCrafterEncodePoses:
|
||||
|
||||
vqvae.to(device)
|
||||
motion_tokens, vq_loss = vqvae(norm_poses.to(device), return_vq=True)
|
||||
|
||||
|
||||
recon_motion = vqvae(norm_poses.to(device))[0][0].to(dtype=torch.float32).cpu().detach() * global_std + global_mean
|
||||
vqvae.to(offload_device)
|
||||
|
||||
@@ -162,7 +206,7 @@ class MTVCrafterEncodePoses:
|
||||
'global_mean': global_mean,
|
||||
'global_std': global_std
|
||||
}
|
||||
|
||||
|
||||
return poses_dict, recon_motion
|
||||
|
||||
|
||||
@@ -181,10 +225,17 @@ class NLFPredict:
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
|
||||
def predict(self, model, images):
|
||||
|
||||
model.to(device)
|
||||
pred = model.detect_smpl_batched(images.permute(0, 3, 1, 2).to(device))
|
||||
model.to(offload_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))
|
||||
finally:
|
||||
torch._C._jit_set_profiling_executor(jit_profiling_prev_state)
|
||||
|
||||
model = model.to(offload_device)
|
||||
|
||||
pred = dict_to_device(pred, offload_device)
|
||||
|
||||
@@ -197,7 +248,7 @@ class NLFPredict:
|
||||
pose_results[key].append(pred[key])
|
||||
else:
|
||||
pose_results[key].append(None)
|
||||
|
||||
|
||||
return (pose_results,)
|
||||
|
||||
class DrawNLFPoses:
|
||||
@@ -208,21 +259,27 @@ class DrawNLFPoses:
|
||||
"width": ("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_NAMES = ("image",)
|
||||
FUNCTION = "predict"
|
||||
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
|
||||
print(type(poses))
|
||||
|
||||
if isinstance(poses, dict):
|
||||
pose_input = poses['joints3d_nonparam'][0] if 'joints3d_nonparam' in poses else poses
|
||||
else:
|
||||
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,)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user