feat: auto download sam model

This commit is contained in:
Radionic
2023-09-26 17:45:32 +08:00
parent d935ffe4fb
commit 91000bc953
7 changed files with 45 additions and 22 deletions
+37 -15
View File
@@ -6,6 +6,7 @@
"""
import os
import sys
sys.path.append(os.path.join(os.path.dirname(__file__)))
import routes
@@ -13,7 +14,9 @@ import inspect
import sys
import importlib
import subprocess
from folder_paths import add_model_folder_path
import requests
from folder_paths import add_model_folder_path, get_filename_list, get_folder_paths
from tqdm import tqdm
ag_path = os.path.join(os.path.dirname(__file__))
@@ -32,7 +35,7 @@ ag_path = os.path.join(os.path.dirname(__file__))
def get_python_files(path):
return [f[:-3] for f in os.listdir(path) if f.endswith('.py')]
return [f[:-3] for f in os.listdir(path) if f.endswith(".py")]
def append_to_sys_path(path):
@@ -40,16 +43,36 @@ def append_to_sys_path(path):
sys.path.append(path)
def create_sam_model_dir():
model_dir = os.path.join(ag_path, "../../models/sam")
def download_sam_model():
model_dir = get_folder_paths("sams")[0]
if not os.path.isdir(model_dir):
os.makedirs(model_dir)
add_model_folder_path('sam', model_dir)
add_model_folder_path("sams", model_dir)
files = get_filename_list("sams")
if len(files) == 0:
print("Downloading sam model...")
url = "https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth"
response = requests.get(url, stream=True)
response.raise_for_status()
file_size = int(response.headers.get("Content-Length", 0))
chunk_size = 1024
num_bars = int(file_size / chunk_size)
with open(f"{model_dir}/sam_vit_h_4b8939.pth", "wb") as f:
for chunk in tqdm(
response.iter_content(chunk_size=chunk_size),
total=num_bars,
unit="KB",
desc=url.split("/")[-1],
):
f.write(chunk)
create_sam_model_dir()
download_sam_model()
paths = ['blender', 'sam', 'common']
paths = ["blender", "sam", "common"]
files = []
for path in paths:
@@ -61,6 +84,7 @@ NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
import blender_node
base_class = blender_node.ObjectOps
# Import all the modules and append their mappings
@@ -71,22 +95,20 @@ for file in files:
if inspect.isclass(obj):
if issubclass(obj, base_class) and obj != base_class:
NODE_CLASS_MAPPINGS.update(obj.NODE_CLASS_MAPPINGS())
NODE_DISPLAY_NAME_MAPPINGS.update(
obj.NODE_DISPLAY_NAME_MAPPINGS())
NODE_DISPLAY_NAME_MAPPINGS.update(obj.NODE_DISPLAY_NAME_MAPPINGS())
if (hasattr(module, 'BLENDER_NODES')):
if hasattr(module, "BLENDER_NODES"):
for node in module.BLENDER_NODES:
# print(node)
# NODE_CLASS_MAPPINGS.update({node: module.BLENDER_NODES[node]})
# NODE_DISPLAY_NAME_MAPPINGS.update({node: node})
NODE_CLASS_MAPPINGS.update(node.NODE_CLASS_MAPPINGS())
NODE_DISPLAY_NAME_MAPPINGS.update(
node.NODE_DISPLAY_NAME_MAPPINGS())
NODE_DISPLAY_NAME_MAPPINGS.update(node.NODE_DISPLAY_NAME_MAPPINGS())
if hasattr(module, 'NODE_CLASS_MAPPINGS'):
if hasattr(module, "NODE_CLASS_MAPPINGS"):
NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS)
if hasattr(module, 'NODE_DISPLAY_NAME_MAPPINGS'):
if hasattr(module, "NODE_DISPLAY_NAME_MAPPINGS"):
NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS)
WEB_DIRECTORY = "js"
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS']
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
View File
+1
View File
@@ -5,4 +5,5 @@ opencv-contrib-python
einops
bpy
segment-anything
tqdm
# -e git+https://github.com/facebookresearch/segment-anything.git#egg=segment_anything
+1 -1
View File
@@ -66,7 +66,7 @@ async def post_sam_model(request):
image = load_image(post.get("image"))
ckpt = post.get("ckpt")
model_type = post.get("model_type")
ckpt = folder_paths.get_full_path("sam", ckpt)
ckpt = folder_paths.get_full_path("sams", ckpt)
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)
+2 -2
View File
@@ -25,7 +25,7 @@ class SAM:
"required": {
"image": ("IMAGE",),
"model_type": (["vit_h", "vit_l", "vit_b"],),
"ckpt": (folder_paths.get_filename_list("sam"),),
"ckpt": (folder_paths.get_filename_list("sams"),),
"embedding_id": (
"STRING",
{"multiline": False, "default": "embedding"},
@@ -46,7 +46,7 @@ class SAM:
global global_predictor
if global_predictor is None:
ckpt = folder_paths.get_full_path("sam", ckpt)
ckpt = folder_paths.get_full_path("sams", ckpt)
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)
global_predictor = predictor
+2 -2
View File
@@ -15,7 +15,7 @@ class SAM_Remote_Emb:
return {
"required": {
"model_type": (["vit_h", "vit_l", "vit_b"],),
"ckpt": (folder_paths.get_filename_list("sam"),),
"ckpt": (folder_paths.get_filename_list("sams"),),
},
}
@@ -27,7 +27,7 @@ class SAM_Remote_Emb:
CATEGORY = "image"
def segment(self, model_type, ckpt):
ckpt = folder_paths.get_full_path("sam", ckpt)
ckpt = folder_paths.get_full_path("sams", ckpt)
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)
+2 -2
View File
@@ -20,7 +20,7 @@ class SAM_Prompt_Image:
"required": {
"image": ("IMAGE",),
"model_type": (["vit_h", "vit_l", "vit_b"],),
"ckpt": (folder_paths.get_filename_list("sam"),),
"ckpt": (folder_paths.get_filename_list("sams"),),
"embedding_id": (
"STRING",
{"multiline": False, "default": "embedding"},
@@ -40,7 +40,7 @@ class SAM_Prompt_Image:
emb_filename = f"{self.output_dir}/{embedding_id}.npy"
if not os.path.exists(emb_filename):
ckpt = folder_paths.get_full_path("sam", ckpt)
ckpt = folder_paths.get_full_path("sams", ckpt)
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)