From 91000bc9538279bf5f4b4aec5adb97debb9f5a2e Mon Sep 17 00:00:00 2001 From: Radionic Date: Tue, 26 Sep 2023 17:45:23 +0800 Subject: [PATCH] feat: auto download sam model --- __init__.py | 52 ++++++++++++++++++++++++++++++++------------- blender/__init__.py | 0 requirements.txt | 1 + routes.py | 2 +- sam/sam.py | 4 ++-- sam/sam_loader.py | 4 ++-- sam/sam_prompt.py | 4 ++-- 7 files changed, 45 insertions(+), 22 deletions(-) delete mode 100644 blender/__init__.py diff --git a/__init__.py b/__init__.py index f7ece75..56c20c6 100644 --- a/__init__.py +++ b/__init__.py @@ -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"] diff --git a/blender/__init__.py b/blender/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/requirements.txt b/requirements.txt index 39acaf4..0ea6151 100644 --- a/requirements.txt +++ b/requirements.txt @@ -5,4 +5,5 @@ opencv-contrib-python einops bpy segment-anything +tqdm # -e git+https://github.com/facebookresearch/segment-anything.git#egg=segment_anything \ No newline at end of file diff --git a/routes.py b/routes.py index fa85b75..2d8ed43 100644 --- a/routes.py +++ b/routes.py @@ -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) diff --git a/sam/sam.py b/sam/sam.py index 6c499fd..81ad70f 100644 --- a/sam/sam.py +++ b/sam/sam.py @@ -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 diff --git a/sam/sam_loader.py b/sam/sam_loader.py index bd7801b..3a4eef7 100644 --- a/sam/sam_loader.py +++ b/sam/sam_loader.py @@ -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) diff --git a/sam/sam_prompt.py b/sam/sam_prompt.py index 1a42f28..7dd99db 100644 --- a/sam/sam_prompt.py +++ b/sam/sam_prompt.py @@ -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)