Files
smthemex-ComfyUI_InteractAv…/InteractAvatar_node.py
T
2026-02-06 09:34:36 +08:00

190 lines
9.1 KiB
Python

# !/usr/bin/env python
# -*- coding: UTF-8 -*-
import numpy as np
import torch
import os
import folder_paths
from typing_extensions import override
from comfy_api.latest import ComfyExtension, io
from .test_wanx_tia2mv_obj_back import generate_video,perdata,load_model,get_z_scale
from .model_loader_utils import trans2path,tensor2pillist,clear_comfyui_cache,get_wav2vec_repo
import nodes
MAX_SEED = np.iinfo(np.int32).max
device = torch.device(
"cuda:0") if torch.cuda.is_available() else torch.device(
"mps") if torch.backends.mps.is_available() else torch.device("cpu")
node_cr_path = os.path.dirname(os.path.abspath(__file__))
weigths_gguf_current_path = os.path.join(folder_paths.models_dir, "gguf")
if not os.path.exists(weigths_gguf_current_path):
os.makedirs(weigths_gguf_current_path)
folder_paths.add_model_folder_path("gguf", weigths_gguf_current_path) # gguf dir
class InteractAvatar_SM_Model(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="InteractAvatar_SM_Model",
display_name="InteractAvatar_SM_Model",
category="InteractAvatar",
inputs=[
io.Combo.Input("dit",options= ["none"] +folder_paths.get_filename_list("diffusion_models") ),
io.Combo.Input("gguf",options= ["none"] +folder_paths.get_filename_list("gguf") ),
io.Combo.Input("lora",options= ["none"] + folder_paths.get_filename_list("loras") ),
io.Combo.Input("back_append_frame",options= [1,2]),
io.Boolean.Input("offload",default=True),
],
outputs=[
io.Model.Output(display_name="model"),
],
)
@classmethod
def execute(cls, dit,gguf,lora,back_append_frame,offload) -> io.NodeOutput:
dit_path=folder_paths.get_full_path("diffusion_models", dit) if dit != "none" else None
gguf_path=folder_paths.get_full_path("gguf", gguf) if gguf != "none" else None
lora_path=folder_paths.get_full_path("loras", lora) if lora != "none" else None
assert dit_path != None , "Please select a model"
clear_comfyui_cache()
model = load_model(
checkpoint_path=dit_path if dit_path != None else gguf_path if gguf_path != None else None,
dit_path=dit_path if dit_path != None else gguf_path if gguf_path != None else None,
lora_path=lora_path,
back_append_frame=int(back_append_frame),
offload=offload
)
return io.NodeOutput(model)
class InteractAvatar_SM_Predata(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="InteractAvatar_SM_Predata",
display_name="InteractAvatar_SM_Predata",
category="InteractAvatar",
inputs=[
io.Clip.Input("clip"),
io.Vae.Input("vae"),
io.Image.Input("images"), # image or video
io.Image.Input("pose_images"), # image or video
io.Int.Input("short_side", default=512, min=256, max=nodes.MAX_RESOLUTION,step=32,display_mode=io.NumberDisplay.number),
io.String.Input("prompt",multiline=True, default="两只手打招呼 伸出大拇指点赞"),
io.String.Input("negative_prompt",multiline=True, default="bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"),
io.String.Input("structured_prompt",multiline=True, default=" (Raise both hands slowly) (Wave both hands side to side for greeting),\n (Make a fist) (Extend the thumb upwards)"),
io.Int.Input("num_frames", default=81, min=8, max=10000,step=1,display_mode=io.NumberDisplay.number),
io.Combo.Input("mode",options=['a2mv','ap2v','mv','a2v','p2v'], default="a2mv"),
io.Combo.Input("back_append_frame",options= [1,2]),
io.Boolean.Input("all_text",default=True),
io.String.Input("wav2vec_repo",multiline=False, default=""),
io.String.Input("object_name",multiline=False, default=""),
io.Audio.Input("audio",optional=True),
io.Image.Input("object_images",optional=True),
io.Mask.Input("object_mask",optional=True),
],
outputs=[
io.Conditioning.Output(display_name="data_dict"),
],
)
@classmethod
def execute(cls, clip,vae,images,pose_images,short_side,prompt,negative_prompt,structured_prompt,num_frames,mode,back_append_frame,all_text,wav2vec_repo,object_name,audio=None,object_images=None,object_mask=None) -> io.NodeOutput:
if back_append_frame==2:
prompt=[ii for ii in prompt.splitlines() if ii]
structured_prompt=[ii for ii in structured_prompt.splitlines() if ii]
assert len(prompt)==len(structured_prompt), "Please make sure the number of prompts and structured prompts are the same,人物提示词和动作提示词的行数要一致"
img=tensor2pillist(images)
dw_img=tensor2pillist(pose_images)
audio_path=trans2path(audio)
data_dict = perdata(
clip,
vae,
img,
dw_img,
object_images,
object_mask,
audio_path,
mode,
prompt=prompt,
negative_prompt=negative_prompt,
structured_prompt=structured_prompt,
frame_num=num_frames,
short_side=short_side,
back_append_frame=back_append_frame,
wav2vec_dir=get_wav2vec_repo(wav2vec_repo),
device=device,
all_text=all_text,
)
return io.NodeOutput(data_dict)
class InteractAvatar_SM_Sampler(io.ComfyNode):
@classmethod
def define_schema(cls):
return io.Schema(
node_id="InteractAvatar_SM_Sampler",
display_name="InteractAvatar_SM_Sampler",
category="InteractAvatar",
inputs=[
io.Model.Input("model"),
io.Conditioning.Input("data_dict"),
io.Int.Input("steps", default=20, min=1, max=10000,display_mode=io.NumberDisplay.number),
io.Int.Input("seed", default=0, min=0, max=MAX_SEED),
io.Float.Input("sample_shift", default=5.0, min=0.0, max=10.0,step=0.1,display_mode=io.NumberDisplay.number),
io.Float.Input("text_guide_scale", default=5.0, min=0.1, max=10.0,step=0.1,display_mode=io.NumberDisplay.number),
io.Float.Input("audio_guide_scale", default=7.5, min=0.1, max=10.0,step=0.1,display_mode=io.NumberDisplay.number),
io.Int.Input("bad_thres", default=800, min=0, max=100000,step=1,display_mode=io.NumberDisplay.number),
io.Int.Input("motion_frame",default=25, min=8, max=10000,step=1,display_mode=io.NumberDisplay.number),
io.Boolean.Input("bad_cfg",default=True),
io.Boolean.Input("three_cfg",default=False),
],
outputs=[
io.Latent.Output(display_name="Latent"),
],
)
@classmethod
def execute(cls, model,data_dict,steps,seed,sample_shift,text_guide_scale,audio_guide_scale,bad_thres,motion_frame,bad_cfg,three_cfg) -> io.NodeOutput:
clear_comfyui_cache()
param_dict = dict(
shift=sample_shift,
sampling_steps=steps,
text_guide_scale=text_guide_scale,
audio_guide_scale=audio_guide_scale,
seed=seed,
bad_cfg=bad_cfg,
three_cfg=three_cfg,
bad_thres=bad_thres,
offload_model=True,
motion_frame=motion_frame,
)
data_dict.update(param_dict)
video = generate_video(model,data_dict)
clear_comfyui_cache() # clear cache
if data_dict["back_append_frame"]==2: # don't need decoder image # B H W C
# import time
# pre_fix=time.strftime("%Y%m%d-%H%M%S")
# save_files=os.path.join(folder_paths.get_output_directory(), f"output_{pre_fix}.pt")
# torch.save(video, save_files) #本地解码
video=model.vae.encode(video) # 统一输出接口
else:
video=get_z_scale(video) #torch.Size([1, 16, 21, 90, 136])
output={"samples":video}
return io.NodeOutput(output)
class InteractAvatar_SM_Extension(ComfyExtension):
@override
async def get_node_list(self) -> list[type[io.ComfyNode]]:
return [
InteractAvatar_SM_Model,
InteractAvatar_SM_Predata,
InteractAvatar_SM_Sampler,
]
async def comfy_entrypoint() -> InteractAvatar_SM_Extension: # ComfyUI calls this to load your extension and its nodes.
return InteractAvatar_SM_Extension()