feat: add sam load and save emb
This commit is contained in:
@@ -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 "
|
||||
}
|
||||
@@ -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 "
|
||||
}
|
||||
Reference in New Issue
Block a user