fix issues from merge

This commit is contained in:
Mel Massadian
2024-07-08 19:10:53 +02:00
parent 9c7db3c59a
commit eb5fddf4de
+54 -41
View File
@@ -104,17 +104,17 @@ class ArgumentConfig:
class DownloadAndLoadLivePortraitModels: class DownloadAndLoadLivePortraitModels:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { return {
}, "required": {},
"optional": { "optional": {
"precision": ( "precision": (
[ [
'fp16', "fp16",
'fp32', "fp32",
], { ],
"default": 'fp16' {"default": "fp16"},
}), ),
} },
} }
RETURN_TYPES = ("LIVEPORTRAITPIPE",) RETURN_TYPES = ("LIVEPORTRAITPIPE",)
@@ -122,7 +122,7 @@ class DownloadAndLoadLivePortraitModels:
FUNCTION = "loadmodel" FUNCTION = "loadmodel"
CATEGORY = "LivePortrait" CATEGORY = "LivePortrait"
def loadmodel(self, precision='fp16'): def loadmodel(self, precision="fp16"):
device = mm.get_torch_device() device = mm.get_torch_device()
mm.soft_empty_cache() mm.soft_empty_cache()
@@ -246,9 +246,9 @@ class DownloadAndLoadLivePortraitModels:
self.spade_generator, self.spade_generator,
self.stich_retargeting_module, self.stich_retargeting_module,
InferenceConfig( InferenceConfig(
device_id=device, device_id=device,
flag_use_half_precision = True if precision == 'fp16' else False flag_use_half_precision=True if precision == "fp16" else False,
) ),
) )
return (pipeline,) return (pipeline,)
@@ -258,38 +258,51 @@ class DownloadAndLoadLivePortraitModels:
class LivePortraitProcess: class LivePortraitProcess:
@classmethod @classmethod
def INPUT_TYPES(s): def INPUT_TYPES(s):
return {"required": { return {
"required": {
"pipeline": ("LIVEPORTRAITPIPE",), "pipeline": ("LIVEPORTRAITPIPE",),
"source_image": ("IMAGE",), "source_image": ("IMAGE",),
"driving_images": ("IMAGE",), "driving_images": ("IMAGE",),
"dsize": ("INT", {"default": 512, "min": 64, "max": 2048}), "dsize": ("INT", {"default": 512, "min": 64, "max": 2048}),
"scale": ("FLOAT", {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01}), "scale": (
"vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}), "FLOAT",
"vy_ratio": ("FLOAT", {"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.01}), {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01},
"lip_zero": ("BOOLEAN", {"default": True}), ),
"eye_retargeting": ("BOOLEAN", {"default": False}), "vx_ratio": (
"eyes_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}), "FLOAT",
"lip_retargeting": ("BOOLEAN", {"default": False}), {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01},
"lip_retargeting_multiplier": ("FLOAT", {"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001}), ),
"stitching": ("BOOLEAN", {"default": True}), "vy_ratio": (
"relative": ("BOOLEAN", {"default": True}), "FLOAT",
{"default": -0.125, "min": -1.0, "max": 1.0, "step": 0.01},
),
"lip_zero": ("BOOLEAN", {"default": True}),
"eye_retargeting": ("BOOLEAN", {"default": False}),
"eyes_retargeting_multiplier": (
"FLOAT",
{"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001},
),
"lip_retargeting": ("BOOLEAN", {"default": False}),
"lip_retargeting_multiplier": (
"FLOAT",
{"default": 1.0, "min": 0.01, "max": 10.0, "step": 0.001},
),
"stitching": ("BOOLEAN", {"default": True}),
"relative": ("BOOLEAN", {"default": True}),
}, },
"optional": { "optional": {
"mismatch_method": ( "mismatch_method": (
["repeat", "cycle", "mirror", "nearest"], ["repeat", "cycle", "mirror", "nearest"],
{"default": "repeat"}, {"default": "repeat"},
), ),
"onnx_device": ( "onnx_device": (
[ [
'CPU', "CPU",
'CUDA', "CUDA",
], { ],
"default": 'CPU' {"default": "CPU"},
}), ),
} },
} }
RETURN_TYPES = ( RETURN_TYPES = (
@@ -320,7 +333,8 @@ class LivePortraitProcess:
eyes_retargeting_multiplier: float, eyes_retargeting_multiplier: float,
lip_retargeting_multiplier: float, lip_retargeting_multiplier: float,
mismatch_method: str = "repeat", mismatch_method: str = "repeat",
, onnx_device='CUDA'): onnx_device="CUDA",
):
source_np = (source_image * 255).byte().numpy() source_np = (source_image * 255).byte().numpy()
driving_images_np = (driving_images * 255).byte().numpy() driving_images_np = (driving_images * 255).byte().numpy()
@@ -370,4 +384,3 @@ NODE_CLASS_MAPPINGS = {
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels", "DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels",
} }