import os import torch import yaml import folder_paths import comfy.model_management as mm import comfy.utils script_directory = os.path.dirname(os.path.abspath(__file__)) from .liveportrait.config.argument_config import ArgumentConfig from .liveportrait.live_portrait_pipeline import LivePortraitPipeline from .liveportrait.utils.cropper import Cropper from .liveportrait.modules.spade_generator import SPADEDecoder from .liveportrait.modules.warping_network import WarpingNetwork from .liveportrait.modules.motion_extractor import MotionExtractor from .liveportrait.modules.appearance_feature_extractor import AppearanceFeatureExtractor from .liveportrait.modules.stitching_retargeting_network import StitchingRetargetingNetwork class InferenceConfig: def __init__(self, mask_crop = None, flag_use_half_precision=True, flag_lip_zero=True, lip_zero_threshold=0.03, flag_eye_retargeting=False, flag_lip_retargeting=False, flag_stitching=True, flag_relative=True, anchor_frame=0, input_shape=(256, 256), flag_write_result=True, flag_pasteback=True, ref_max_shape=1280, ref_shape_n=2, device_id=0, flag_do_crop=True, flag_do_rot=True): self.flag_use_half_precision = flag_use_half_precision self.flag_lip_zero = flag_lip_zero self.lip_zero_threshold = lip_zero_threshold self.flag_eye_retargeting = flag_eye_retargeting self.flag_lip_retargeting = flag_lip_retargeting self.flag_stitching = flag_stitching self.flag_relative = flag_relative self.anchor_frame = anchor_frame self.input_shape = input_shape self.flag_write_result = flag_write_result self.flag_pasteback = flag_pasteback self.ref_max_shape = ref_max_shape self.ref_shape_n = ref_shape_n self.device_id = device_id self.flag_do_crop = flag_do_crop self.flag_do_rot = flag_do_rot self.mask_crop=mask_crop class CropConfig: def __init__(self, dsize=512, scale=2.3, vx_ratio=0, vy_ratio=-0.125): self.dsize = dsize self.scale = scale self.vx_ratio = vx_ratio self.vy_ratio = vy_ratio class ArgumentConfig: def __init__(self, device_id=0, flag_lip_zero=True, flag_eye_retargeting=False, flag_lip_retargeting=False, flag_stitching=True, flag_relative=True, flag_pasteback=True, flag_do_crop=True, flag_do_rot=True, dsize=512, scale=2.3, vx_ratio=0, vy_ratio=-0.125, ): self.device_id = device_id self.flag_lip_zero = flag_lip_zero self.flag_eye_retargeting = flag_eye_retargeting self.flag_lip_retargeting = flag_lip_retargeting self.flag_stitching = flag_stitching self.flag_relative = flag_relative self.flag_pasteback = flag_pasteback self.flag_do_crop = flag_do_crop self.flag_do_rot = flag_do_rot self.dsize = dsize self.scale = scale self.vx_ratio = vx_ratio self.vy_ratio = vy_ratio class DownloadAndLoadLivePortraitModels: @classmethod def INPUT_TYPES(s): return {"required": { }, "optional": { "precision": ( [ 'fp16', 'fp32', ], { "default": 'fp16' }), } } RETURN_TYPES = ("LIVEPORTRAITPIPE",) RETURN_NAMES = ("live_portrait_pipe",) FUNCTION = "loadmodel" CATEGORY = "LivePortrait" def loadmodel(self, precision='fp16'): device = mm.get_torch_device() mm.soft_empty_cache() pbar = comfy.utils.ProgressBar(3) download_path = os.path.join(folder_paths.models_dir, "liveportrait") model_path = os.path.join(download_path) if not os.path.exists(model_path): print(f"Downloading model to: {model_path}") from huggingface_hub import snapshot_download snapshot_download(repo_id="Kijai/LivePortrait_safetensors", local_dir=download_path, local_dir_use_symlinks=False) model_config_path = os.path.join(script_directory, 'liveportrait', 'config', 'models.yaml') with open(model_config_path, 'r') as file: model_config = yaml.safe_load(file) feature_extractor_path = os.path.join(model_path, 'appearance_feature_extractor.safetensors') motion_extractor_path = os.path.join(model_path, 'motion_extractor.safetensors') warping_module_path = os.path.join(model_path, 'warping_module.safetensors') spade_generator_path = os.path.join(model_path, 'spade_generator.safetensors') stitching_retargeting_path = os.path.join(model_path, 'stitching_retargeting_module.safetensors') # init F model_params = model_config['model_params']['appearance_feature_extractor_params'] self.appearance_feature_extractor = AppearanceFeatureExtractor(**model_params).to(device) self.appearance_feature_extractor.load_state_dict(comfy.utils.load_torch_file(feature_extractor_path)) self.appearance_feature_extractor.eval() print('Load appearance_feature_extractor done.') pbar.update(1) # init M model_params = model_config['model_params']['motion_extractor_params'] self.motion_extractor = MotionExtractor(**model_params).to(device) self.motion_extractor.load_state_dict(comfy.utils.load_torch_file(motion_extractor_path)) self.motion_extractor.eval() print('Load motion_extractor done.') pbar.update(1) # init W model_params = model_config['model_params']['warping_module_params'] self.warping_module = WarpingNetwork(**model_params).to(device) self.warping_module.load_state_dict(comfy.utils.load_torch_file(warping_module_path)) self.warping_module.eval() print('Load warping_module done.') pbar.update(1) # init G model_params = model_config['model_params']['spade_generator_params'] self.spade_generator = SPADEDecoder(**model_params).to(device) self.spade_generator.load_state_dict(comfy.utils.load_torch_file(spade_generator_path)) self.spade_generator.eval() print('Load spade_generator done.') pbar.update(1) def filter_checkpoint_for_model(checkpoint, prefix): """Filter and adjust the checkpoint dictionary for a specific model based on the prefix.""" # Create a new dictionary where keys are adjusted by removing the prefix and the model name filtered_checkpoint = {key.replace(prefix + "_module.", ""): value for key, value in checkpoint.items() if key.startswith(prefix)} return filtered_checkpoint config = model_config['model_params']['stitching_retargeting_module_params'] checkpoint = comfy.utils.load_torch_file(stitching_retargeting_path) stitcher_prefix = 'retarget_shoulder' stitcher_checkpoint = filter_checkpoint_for_model(checkpoint, stitcher_prefix) stitcher = StitchingRetargetingNetwork(**config.get('stitching')) stitcher.load_state_dict(stitcher_checkpoint) stitcher = stitcher.to(device) stitcher.eval() lip_prefix = 'retarget_mouth' lip_checkpoint = filter_checkpoint_for_model(checkpoint, lip_prefix) retargetor_lip = StitchingRetargetingNetwork(**config.get('lip')) retargetor_lip.load_state_dict(lip_checkpoint) retargetor_lip = retargetor_lip.to(device) retargetor_lip.eval() eye_prefix = 'retarget_eye' eye_checkpoint = filter_checkpoint_for_model(checkpoint, eye_prefix) retargetor_eye = StitchingRetargetingNetwork(**config.get('eye')) retargetor_eye.load_state_dict(eye_checkpoint) retargetor_eye = retargetor_eye.to(device) retargetor_eye.eval() print('Load stitching_retargeting_module done.') self.stich_retargeting_module = { 'stitching': stitcher, 'lip': retargetor_lip, 'eye': retargetor_eye } pipeline = LivePortraitPipeline( self.appearance_feature_extractor, self.motion_extractor, self.warping_module, self.spade_generator, self.stich_retargeting_module, InferenceConfig( device_id=device, flag_use_half_precision = True if precision == 'fp16' else False ) ) return (pipeline,) class LivePortraitProcess: @classmethod def INPUT_TYPES(s): return {"required": { "pipeline": ("LIVEPORTRAITPIPE",), "source_image": ("IMAGE",), "driving_images": ("IMAGE",), "dsize": ("INT", {"default": 512, "min": 64, "max": 2048}), "scale": ("FLOAT", {"default": 2.3, "min": 1.0, "max": 4.0, "step": 0.01}), "vx_ratio": ("FLOAT", {"default": 0.0, "min": -1.0, "max": 1.0, "step": 0.01}), "vy_ratio": ("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": { "onnx_device": ( [ 'CPU', 'CUDA', ], { "default": 'CPU' }), } } RETURN_TYPES = ("IMAGE", "IMAGE",) RETURN_NAMES = ("cropped_images", "full_images",) FUNCTION = "process" CATEGORY = "LivePortrait" def process(self, source_image, driving_images, dsize, scale, vx_ratio, vy_ratio, pipeline, lip_zero, eye_retargeting, lip_retargeting, stitching, relative, eyes_retargeting_multiplier, lip_retargeting_multiplier, onnx_device='CUDA'): source_image_np = (source_image * 255).byte().numpy() driving_images_np = (driving_images * 255).byte().numpy() crop_cfg = CropConfig( dsize = dsize, scale = scale, vx_ratio = vx_ratio, vy_ratio = vy_ratio, ) cropper = Cropper(crop_cfg=crop_cfg, provider=onnx_device) pipeline.cropper = cropper pipeline.live_portrait_wrapper.cfg.flag_eye_retargeting = eye_retargeting pipeline.live_portrait_wrapper.cfg.eyes_retargeting_multiplier = eyes_retargeting_multiplier pipeline.live_portrait_wrapper.cfg.flag_lip_retargeting = lip_retargeting pipeline.live_portrait_wrapper.cfg.lip_retargeting_multiplier = lip_retargeting_multiplier pipeline.live_portrait_wrapper.cfg.flag_stitching = stitching pipeline.live_portrait_wrapper.cfg.flag_relative = relative pipeline.live_portrait_wrapper.cfg.flag_lip_zero = lip_zero cropped_out_list = [] full_out_list = [] for img in source_image_np: cropped_frames, full_frame = pipeline.execute(img, driving_images_np) cropped_tensors = [torch.from_numpy(np_array) for np_array in cropped_frames] cropped_tensors_out = torch.stack(cropped_tensors) / 255 cropped_tensors_out = cropped_tensors_out.cpu().float() full_tensors = [torch.from_numpy(np_array) for np_array in full_frame] full_tensors_out = torch.stack(full_tensors) / 255 full_tensors_out = full_tensors_out.cpu().float() cropped_out_list.append(cropped_tensors_out) full_out_list.append(full_tensors_out) cropped_tensors_out = torch.cat(cropped_out_list, dim=0) full_tensors_out = torch.cat(full_out_list, dim=0) return (cropped_tensors_out, full_tensors_out) NODE_CLASS_MAPPINGS = { "DownloadAndLoadLivePortraitModels": DownloadAndLoadLivePortraitModels, "LivePortraitProcess": LivePortraitProcess, } NODE_DISPLAY_NAME_MAPPINGS = { "DownloadAndLoadLivePortraitModels": "(Down)Load LivePortraitModels", "LivePortraitProcess": "LivePortraitProcess", }