Files
Fannovel16-ComfyUI-MotionDiff/__init__.py
T

329 lines
12 KiB
Python

import torch
import os
import sys
from pathlib import Path
EXTENSION_PATH = Path(__file__).parent
sys.path.insert(0, str(EXTENSION_PATH.resolve()))
import custom_mmpkg.custom_mmcv as mmcv
import numpy as np
import torch
from motiondiff_modules.mogen.models import build_architecture
from custom_mmpkg.custom_mmcv.runner import load_checkpoint
from motiondiff_modules.mogen.utils.plot_utils import (
plot_3d_motion,
t2m_kinematic_chain
)
from comfy.model_management import get_torch_device
from .md_config import get_model_dataset_dict
from .utils import *
#from custom_mmpkg.custom_mmhuman3d.core.conventions.keypoints_mapping import convert_kps
#from custom_mmpkg.custom_mmhuman3d.core.visualization.visualize_keypoints3d import visualize_kp3d
from pathlib import Path
import importlib, traceback
def create_mdm_model(model_config):
cfg = mmcv.Config.fromstring(model_config.config_code, '.py')
mdm = build_architecture(cfg.model)
load_checkpoint(mdm, str(model_config.ckpt_path), map_location='cpu')
mdm.eval().cpu()
return mdm
class MotionDiffModelWrapper(torch.nn.Module): #Anything beside CLIP (mdm.model)
def __init__(self, mdm, dataset) -> None:
super(MotionDiffModelWrapper, self).__init__()
self.loss_recon=mdm.loss_recon
self.diffusion_train=mdm.diffusion_train
self.diffusion_test=mdm.diffusion_test
self.sampler=mdm.sampler
self.dataset = dataset
def forward(self, clip, cond_dict, **kwargs):
motion, motion_mask = kwargs['motion'].float(), kwargs['motion_mask'].float()
sample_idx = kwargs.get('sample_idx', None)
clip_feat = kwargs.get('clip_feat', None)
sampler = kwargs.get('sampler', 'ddpm')
seed = kwargs.get('seed', None)
B, T = motion.shape[:2]
dim_pose = kwargs['motion'].shape[-1]
model_kwargs = cond_dict
model_kwargs['motion_mask'] = motion_mask
model_kwargs['sample_idx'] = sample_idx
inference_kwargs = kwargs.get('inference_kwargs', {})
if sampler == 'ddpm':
output = self.diffusion_test.p_sample_loop(
clip,
(B, T, dim_pose),
clip_denoised=False,
progress=False,
model_kwargs=model_kwargs,
seed=seed,
**inference_kwargs
)
else:
output = self.diffusion_test.ddim_sample_loop(
clip,
(B, T, dim_pose),
clip_denoised=False,
progress=False,
model_kwargs=model_kwargs,
eta=0,
seed=seed,
**inference_kwargs
)
if getattr(clip, "post_process") is not None:
output = clip.post_process(output)
results = kwargs
results['pred_motion'] = output
results = self.split_results(results)
return results
def split_results(self, results):
B = results['motion'].shape[0]
output = []
for i in range(B):
batch_output = dict()
batch_output['motion'] = to_cpu(results['motion'][i])
batch_output['pred_motion'] = to_cpu(results['pred_motion'][i])
batch_output['motion_length'] = to_cpu(results['motion_length'][i])
batch_output['motion_mask'] = to_cpu(results['motion_mask'][i])
if 'pred_motion_length' in results.keys():
batch_output['pred_motion_length'] = to_cpu(results['pred_motion_length'][i])
else:
batch_output['pred_motion_length'] = to_cpu(results['motion_length'][i])
if 'pred_motion_mask' in results:
batch_output['pred_motion_mask'] = to_cpu(results['pred_motion_mask'][i])
else:
batch_output['pred_motion_mask'] = to_cpu(results['motion_mask'][i])
if 'motion_metas' in results.keys():
motion_metas = results['motion_metas'][i]
if 'text' in motion_metas.keys():
batch_output['text'] = motion_metas['text']
if 'token' in motion_metas.keys():
batch_output['token'] = motion_metas['token']
output.append(batch_output)
return output
class MotionDiffCLIPWrapper(torch.nn.Module):
def __init__(self, mdm):
super(MotionDiffCLIPWrapper, self).__init__()
self.model = mdm.model
def forward(self, text, motion_data):
self.model.to(get_torch_device())
B, T = motion_data["motion"].shape[:2]
texts = []
for _ in range(B):
texts.append(text)
out = self.model.get_precompute_condition(device=get_torch_device(), text=texts, **motion_data)
self.model.cpu()
return out
model_dataset_dict = None
class MotionDiffLoader:
@classmethod
def INPUT_TYPES(s):
global model_dataset_dict
model_dataset_dict = get_model_dataset_dict()
return {
"required": {
"model_dataset": (
list(model_dataset_dict.keys()),
{ "default": "-human_ml3d" }
)
},
}
RETURN_TYPES = ("MD_MODEL", "MD_CLIP")
CATEGORY = "MotionDiff"
FUNCTION = "load_mdm"
def load_mdm(self, model_dataset):
global model_dataset_dict
if model_dataset_dict is None:
model_dataset_dict = get_model_dataset_dict() #In case of API users
model_config = model_dataset_dict[model_dataset]()
mdm = create_mdm_model(model_config)
return (MotionDiffModelWrapper(mdm, dataset=model_config.dataset), MotionDiffCLIPWrapper(mdm))
class MotionCLIPTextEncode:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"md_clip": ("MD_CLIP", ),
"motion_data": ("MOTION_DATA", ),
"text": ("STRING", {"default": "a person performs a cartwheel" ,"multiline": True})
},
}
RETURN_TYPES = ("MD_CONDITIONING",)
CATEGORY = "MotionDiff"
FUNCTION = "encode_text"
def encode_text(self, md_clip, motion_data, text):
return (md_clip(text, motion_data), )
class EmptyMotionData:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"frames": ("INT", {"default": 196, "min": 1, "max": 196})
}
}
RETURN_TYPES = ("MOTION_DATA", )
CATEGORY = "MotionDiff"
FUNCTION = "encode_text"
def encode_text(self, frames):
return ({
'motion': torch.zeros(1, frames, 263),
'motion_mask': torch.ones(1, frames),
'motion_length': torch.Tensor([frames]).long(),
}, )
class MotionDiffSimpleSampler:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"sampler_name": (["ddpm", "ddim"], {"default": "ddim"}),
"md_model": ("MD_MODEL", ),
"md_clip": ("MD_CLIP", ),
"md_cond": ("MD_CONDITIONING", ),
"motion_data": ("MOTION_DATA",),
"seed": ("INT", {"default": 123,"min": 0, "max": 0xffffffffffffffff, "step": 1}),
}
}
RETURN_TYPES = ("MOTION_DATA",)
CATEGORY = "MotionDiff"
FUNCTION = "sample"
def sample(self, sampler_name, md_model: MotionDiffModelWrapper, md_clip, md_cond, motion_data, seed):
md_model.to(get_torch_device())
md_clip.to(get_torch_device())
for key in motion_data:
motion_data[key] = to_gpu(motion_data[key])
kwargs = {
**motion_data,
'inference_kwargs': {},
'sampler': sampler_name,
'seed': seed
}
with torch.no_grad():
output = md_model(md_clip.model, cond_dict=md_cond, **kwargs)[0]['pred_motion']
pred_motion = output * md_model.dataset.std + md_model.dataset.mean
pred_motion = pred_motion.cpu().detach()
md_model.cpu(), md_clip.cpu()
for key in motion_data:
motion_data[key] = to_cpu(motion_data[key])
return ({
'motion': pred_motion,
'motion_mask': motion_data['motion_mask'],
'motion_length': motion_data['motion_length'],
}, )
class MotionDataVisualizer:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"motion_data": ("MOTION_DATA", ),
"visualization": (["original", "pseudo-openpose"], {"default": "pseudo-openpose"}),
"distance": ("FLOAT", {"default": 7.0, "min": 0.0, "max": 10.0, "step": 0.1}),
"elevation": ("FLOAT", {"default": 120, "min": 0.0, "max": 300.0, "step": 0.1}),
"rotation": ("FLOAT", {"default": -90, "min": -180, "max": 180, "step": 1}),
"poselinewidth": ("FLOAT", {"default": 4, "min": 0, "max": 50, "step": 0.1}),
},
"optional": {
"opt_title": ("STRING", {"default": '' ,"multiline": False}),
}
}
RETURN_TYPES = ("IMAGE",)
CATEGORY = "MotionDiff"
FUNCTION = "visualize"
def visualize(self, motion_data, visualization, distance, elevation, rotation, poselinewidth, opt_title=None):
if "joints" in motion_data:
joints = motion_data["joints"]
else:
joints = motion_data_to_joints(motion_data["motion"])
pil_frames = plot_3d_motion(
None, t2m_kinematic_chain, joints, distance, elevation, rotation, poselinewidth,
title=opt_title if opt_title is not None else '',
fps=1, save_as_pil_lists=True, visualization=visualization
)
tensor_frames = []
for pil_image in pil_frames:
np_image = np.array(pil_image.convert("RGB")).astype(np.float32) / 255.0
tensor_frames.append(torch.from_numpy(np_image))
return (torch.stack(tensor_frames, dim=0), )
def load_nodes():
shorted_errors = []
full_error_messages = []
node_class_mappings = {}
node_display_name_mappings = {}
for filename in (Path(__file__).parent / "nodes").iterdir():
module_name = filename.stem
if module_name.startswith('.'): continue #Skip hidden files created by the OS (e.g. [.DS_Store](https://en.wikipedia.org/wiki/.DS_Store))
try:
module = importlib.import_module(f".nodes.{module_name}", package=__package__)
node_class_mappings.update(getattr(module, "NODE_CLASS_MAPPINGS"))
if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS"):
node_display_name_mappings.update(getattr(module, "NODE_DISPLAY_NAME_MAPPINGS"))
print(f"Imported {module_name} nodes")
except AttributeError:
pass # wip nodes
except Exception:
error_message = traceback.format_exc()
full_error_messages.append(error_message)
error_message = error_message.splitlines()[-1]
shorted_errors.append(
f"Failed to import module {module_name} because {error_message}"
)
if len(shorted_errors) > 0:
full_err_log = '\n\n'.join(full_error_messages)
print(f"\n\nFull error log from ComfyUI-MotionDiff: \n{full_err_log}\n\n")
print(
f"Some nodes failed to load:\n\t"
+ "\n\t".join(shorted_errors)
+ "\n\n"
+ "Check that you properly installed the dependencies.\n"
+ "If you think this is a bug, please report it on the github page (https://github.com/Fannovel16/ComfyUI-MotionDiff/issues)"
)
return node_class_mappings, node_display_name_mappings
AUX_NODE_MAPPINGS, AUX_DISPLAY_NAME_MAPPINGS = load_nodes()
NODE_CLASS_MAPPINGS = {
"MotionDiffLoader": MotionDiffLoader,
"MotionCLIPTextEncode": MotionCLIPTextEncode,
"MotionDiffSimpleSampler": MotionDiffSimpleSampler,
"EmptyMotionData": EmptyMotionData,
"MotionDataVisualizer": MotionDataVisualizer,
**AUX_NODE_MAPPINGS
}
NODE_DISPLAY_NAME_MAPPINGS = {
"MotionDiffLoader": "MotionDiff Loader",
"MotionCLIPTextEncode": "MotionCLIP Text Encode",
"MotionDiffSimpleSampler": "MotionDiff Simple Sampler",
"EmptyMotionData": "Empty Motion Data",
"MotionDataVisualizer": "Motion Data Visualizer",
**AUX_DISPLAY_NAME_MAPPINGS
}