From 6168b3a2ac38b5eebed3daf9e52df5742abf6813 Mon Sep 17 00:00:00 2001 From: melMass Date: Mon, 17 Jul 2023 22:52:25 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20=E2=9A=A1=EF=B8=8F=20from=20tensor2np=20?= =?UTF-8?q?always=20returning=20a=20list?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- __init__.py | 21 ++++++++++++--------- nodes/faceenhance.py | 7 +++---- nodes/faceswap.py | 4 ++-- utils.py | 2 +- 4 files changed, 18 insertions(+), 16 deletions(-) diff --git a/__init__.py b/__init__.py index ec213b3..03c9f48 100644 --- a/__init__.py +++ b/__init__.py @@ -14,8 +14,14 @@ NODE_DISPLAY_NAME_MAPPINGS = {} NODE_CLASS_MAPPINGS_DEBUG = {} -def extract_nodes_from_source(source_code): +def extract_nodes_from_source(filename): + source_code = "" + + with open(filename, "r") as file: + source_code = file.read() + nodes = [] + try: parsed = ast.parse(source_code) for node in ast.walk(parsed): @@ -36,6 +42,7 @@ def extract_nodes_from_source(source_code): except SyntaxError: log.error("Failed to parse") pass # File couldn't be parsed + return nodes @@ -64,13 +71,10 @@ def load_nodes(): f"Failed to import module {module_name} because {error_message}" ) # Read __nodes__ variable from the source file - with open(filename, "r") as file: - source_code = file.read() - extracted_nodes = extract_nodes_from_source(source_code) - nodes_failed.extend(extracted_nodes) + nodes_failed.extend(extract_nodes_from_source(filename)) if errors: - log.error( + log.info( f"Some nodes failed to load:\n\t" + "\n\t".join(errors) + "\n\n" @@ -91,11 +95,11 @@ elif web_extensions_root.exists(): try: os.symlink((here / "web"), web_mtb.as_posix()) except Exception: # OSError - log.error( + log.warn( f"Failed to create symlink to {web_mtb}. Please copy the folder manually." ) else: - log.error( + log.warn( f"Comfy root probably not found automatically, please copy the folder {web_mtb} manually in the web/extensions folder of ComfyUI" ) @@ -126,4 +130,3 @@ MANIFEST = { "project": "https://github.com/melMass/comfy_mtb", # The address that the `name` value will link to on Node Class Views "description": "Set of nodes that enhance your animation workflow and provide a range of useful tools including features such as manipulating bounding boxes, perform color corrections, swap faces in images, interpolate frames for smooth animation, export to ProRes format, apply various image operations, work with latent spaces, generate QR codes, and create normal and height maps for textures.", } - diff --git a/nodes/faceenhance.py b/nodes/faceenhance.py index 33296eb..95a55f9 100644 --- a/nodes/faceenhance.py +++ b/nodes/faceenhance.py @@ -111,7 +111,7 @@ class BGUpscaleWrapper: self.upscale_model.cpu() s = torch.clamp(s.movedim(-3, -1), min=0, max=1.0) - return (tensor2np(s),) + return (tensor2np(s)[0],) import sys @@ -150,9 +150,8 @@ class RestoreFace: weight, save_tmp_steps, ) -> torch.Tensor: - pimage = tensor2pil(image) - width, height = pimage.size - + pimage = tensor2np(image)[0] + width, height = pimage.shape[1], pimage.shape[0] source_img = cv2.cvtColor(np.array(pimage), cv2.COLOR_RGB2BGR) sys.stdout = NullWriter() diff --git a/nodes/faceswap.py b/nodes/faceswap.py index 8dcf7ac..0a154ba 100644 --- a/nodes/faceswap.py +++ b/nodes/faceswap.py @@ -100,8 +100,8 @@ class FaceSwap: ): def do_swap(img): model_management.throw_exception_if_processing_interrupted() - img = tensor2pil(img) - ref = tensor2pil(reference) + img = tensor2pil(img)[0] + ref = tensor2pil(reference)[0] face_ids = { int(x) for x in faces_index.strip(",").split(",") if x.isnumeric() } diff --git a/utils.py b/utils.py index 0a93e7f..cafe31d 100644 --- a/utils.py +++ b/utils.py @@ -77,7 +77,7 @@ def np2tensor(img_np: np.ndarray | List[np.ndarray]) -> torch.Tensor: return torch.from_numpy(img_np.astype(np.float32) / 255.0).unsqueeze(0) -def tensor2np(tensor: torch.Tensor) -> Union[np.ndarray, List[np.ndarray]]: +def tensor2np(tensor: torch.Tensor) -> List[np.ndarray]: batch_count = 1 if len(tensor.shape) > 3: batch_count = tensor.size(0)