From dd8833d19555a3883a1670b0468b60ed1b2e5bdc Mon Sep 17 00:00:00 2001 From: BennyKok Date: Wed, 6 Sep 2023 16:47:12 +0800 Subject: [PATCH] feat: add sam load and save emb --- sam/sam_load_emb.py | 47 ++++++++++++++++++++++++++++++++++ sam/sam_save_emb.py | 61 +++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 108 insertions(+) create mode 100644 sam/sam_load_emb.py create mode 100644 sam/sam_save_emb.py diff --git a/sam/sam_load_emb.py b/sam/sam_load_emb.py new file mode 100644 index 0000000..b089379 --- /dev/null +++ b/sam/sam_load_emb.py @@ -0,0 +1,47 @@ +import folder_paths +import torch +from segment_anything import SamAutomaticMaskGenerator, sam_model_registry +from einops import rearrange, repeat + +class SAM_Load_Embedding: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "filename": ("STRING", { + "multiline": False, + "default": "embeddings" + }), + } + } + + RETURN_TYPES = ("EMBEDDINGS",) + RETURN_NAMES = ("EMBEDDINGS",) + + OUTPUT_NODE = True + + FUNCTION = "process" + + CATEGORY = "image" + + def process(self, filename): + import json + import numpy as np + + data = {} + with open(filename, 'r') as f: + data = json.load(f) + + # Convert list to numpy ndarray + data['image_embedding'] = np.array(data['image_embedding']) + + return (data, ) + + +NODE_CLASS_MAPPINGS = { + "SAM_Load_Embedding": SAM_Load_Embedding +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "SAM_Load_Embedding": "SAM_Load_Embedding " +} diff --git a/sam/sam_save_emb.py b/sam/sam_save_emb.py new file mode 100644 index 0000000..56f2142 --- /dev/null +++ b/sam/sam_save_emb.py @@ -0,0 +1,61 @@ +import folder_paths +import torch +from segment_anything import SamAutomaticMaskGenerator, sam_model_registry +from einops import rearrange, repeat + +class SAM_Save_Embedding: + def __init__(self): + self.output_dir = folder_paths.get_output_directory() + self.type = "output" + self.prefix_append = "" + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "embeddings": ("EMBEDDINGS",), + "filename": ("STRING", { + "multiline": False, + "default": "embeddings" + }), + "write_mode": (["Overwrite", "Increment"],), + } + } + + RETURN_TYPES = () + RETURN_NAMES = () + + OUTPUT_NODE = True + + FUNCTION = "process" + + CATEGORY = "image" + + def process(self, embeddings, filename, write_mode): + import json + + filepath = self.output_dir + "/" + filename + '.json' + + if write_mode == "Increment": + count = 0 + # while file exists, increment count + while os.path.exists(self.output_dir + "/" + filename + '_' + str(count) + '.json'): + count += 1 + + filepath = self.output_dir + "/" + filename + '_' + str(count) + '.json' + + # print(embeddings) + embeddings['image_embedding'] = embeddings['image_embedding'].tolist() + with open(filepath, 'w') as f: + json.dump(embeddings, f) + + return { "ui" : { "file": { filepath.replace(f"{self.output_dir}/", "") } } } + + +NODE_CLASS_MAPPINGS = { + "SAM_Save_Embedding": SAM_Save_Embedding +} + +NODE_DISPLAY_NAME_MAPPINGS = { + "SAM_Save_Embedding": "SAM_Save_Embedding " +}