ExpressionVideo2VideoNode

This commit is contained in:
shadowcz007
2024-08-08 22:58:23 +08:00
parent f78867374a
commit 8ec5337668
2 changed files with 63 additions and 3 deletions
+5 -3
View File
@@ -1,5 +1,5 @@
from .nodes.live_portrait import LivePortraitNode,FaceCropInfo,Retargeting,LivePortraitVideoNode
from .nodes.expression_editor import ExpressionEditor,ExpressionVideoNode
from .nodes.expression_editor import ExpressionEditor,ExpressionVideoNode,ExpressionVideo2VideoNode
NODE_CLASS_MAPPINGS = {
"LivePortraitNode": LivePortraitNode,
@@ -7,7 +7,8 @@ NODE_CLASS_MAPPINGS = {
"FaceCropInfo":FaceCropInfo,
"Retargeting":Retargeting,
"ExpressionEditor_":ExpressionEditor,
"ExpressionVideoNode":ExpressionVideoNode
"ExpressionVideoNode":ExpressionVideoNode,
"ExpressionVideo2VideoNode":ExpressionVideo2VideoNode
}
# dict = { "key":value }
@@ -18,7 +19,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"FaceCropInfo":"Face Crop Info",
"Retargeting":"Retargeting",
"ExpressionEditor_":"Expression Editor",
"ExpressionVideoNode":"Expression Video"
"ExpressionVideoNode":"Expression Video",
"ExpressionVideo2VideoNode":"Expression Video 2 Video"
}
# web ui的节点功能
+58
View File
@@ -540,7 +540,65 @@ class ExpressionVideoNode:
result=torch.cat(result, dim=0)
return (result,)
#
class ExpressionVideo2VideoNode:
def __init__(self):
self.src_image = None
@classmethod
def INPUT_TYPES(s):
return {"required": {
"src_frames": ("IMAGE",), #batch
"from_expression":("STRING", {"forceInput": True,"dynamicPrompts": False}),
"to_expression":("STRING", {"forceInput": True,"dynamicPrompts": False}),
"interpolation_type":( ['linear', 'nearest', 'cubic'],
{"default": "cubic"}),
},
}
RETURN_TYPES = ("IMAGE",)
RETURN_NAMES = ("frames",)
FUNCTION = "run"
# OUTPUT_NODE = True
CATEGORY = "♾️Mixlab/Video/LivePortrait"
INPUT_IS_LIST = False
OUTPUT_IS_LIST = (False,) #x list
def run(self,src_frames,from_expression, to_expression,interpolation_type ):
images = [src_frames[i:i + 1, ...] for i in range(src_frames.shape[0])]
interpolations_num=len(images)
from_expression = json.loads(from_expression)
to_expression = json.loads(to_expression)
from_expression=update_expression_json(from_expression)
to_expression=update_expression_json(to_expression)
exps=interpolate_dicts(from_expression,to_expression,interpolations_num,interpolation_type)
result=[]
for i in range(interpolations_num):
src_image = images[i]
exp=exps[i]
self.psi = g_engine.prepare_source(src_image)
self.src_image = src_image
out_img = expression_run(self.psi,exp)
result.append(out_img)
result=torch.cat(result, dim=0)
return (result,)
class ExpressionEditor:
def __init__(self):