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.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 = { NODE_CLASS_MAPPINGS = {
"LivePortraitNode": LivePortraitNode, "LivePortraitNode": LivePortraitNode,
@@ -7,7 +7,8 @@ NODE_CLASS_MAPPINGS = {
"FaceCropInfo":FaceCropInfo, "FaceCropInfo":FaceCropInfo,
"Retargeting":Retargeting, "Retargeting":Retargeting,
"ExpressionEditor_":ExpressionEditor, "ExpressionEditor_":ExpressionEditor,
"ExpressionVideoNode":ExpressionVideoNode "ExpressionVideoNode":ExpressionVideoNode,
"ExpressionVideo2VideoNode":ExpressionVideo2VideoNode
} }
# dict = { "key":value } # dict = { "key":value }
@@ -18,7 +19,8 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"FaceCropInfo":"Face Crop Info", "FaceCropInfo":"Face Crop Info",
"Retargeting":"Retargeting", "Retargeting":"Retargeting",
"ExpressionEditor_":"Expression Editor", "ExpressionEditor_":"Expression Editor",
"ExpressionVideoNode":"Expression Video" "ExpressionVideoNode":"Expression Video",
"ExpressionVideo2VideoNode":"Expression Video 2 Video"
} }
# web ui的节点功能 # web ui的节点功能
+58
View File
@@ -540,7 +540,65 @@ class ExpressionVideoNode:
result=torch.cat(result, dim=0) result=torch.cat(result, dim=0)
return (result,) 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: class ExpressionEditor:
def __init__(self): def __init__(self):