feat: infer model type from checkpoint name

This commit is contained in:
Radionic
2023-09-26 18:34:46 +08:00
parent 6e3820b502
commit 722542e92a
4 changed files with 8 additions and 9 deletions
-4
View File
@@ -298,9 +298,6 @@ const ext = {
node.widgets.find((x) => x.name === "embedding_id").value = id;
const ckpt = node.widgets.find((x) => x.name === "ckpt").value;
const model_type = node.widgets.find(
(x) => x.name === "model_type"
).value;
api
.fetchApi("/sam_model", {
@@ -309,7 +306,6 @@ const ext = {
image: connectedImageFileName,
embedding_id: id,
ckpt,
model_type,
}),
})
.then(() => {
+2 -1
View File
@@ -7,6 +7,7 @@ import folder_paths
import json
import numpy as np
import server
import re
# For speeding up ONNX model, see https://github.com/facebookresearch/segment-anything/tree/main/demo#onnx-multithreading-with-sharedarraybuffer
def inject_headers(original_handler):
@@ -65,7 +66,7 @@ async def post_sam_model(request):
if not os.path.exists(emb_filename):
image = load_image(post.get("image"))
ckpt = post.get("ckpt")
model_type = post.get("model_type")
model_type = re.findall(r'vit_[lbh]', ckpt)[0]
ckpt = folder_paths.get_full_path("sams", ckpt)
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)
+3 -2
View File
@@ -2,6 +2,7 @@ import folder_paths
import os
import numpy as np
import torch
import re
from segment_anything import sam_model_registry, SamPredictor
from einops import rearrange, repeat
@@ -24,7 +25,6 @@ class SAM:
return {
"required": {
"image": ("IMAGE",),
"model_type": (["vit_h", "vit_l", "vit_b"],),
"ckpt": (folder_paths.get_filename_list("sams"),),
"embedding_id": (
"STRING",
@@ -40,13 +40,14 @@ class SAM:
RETURN_TYPES = ("SAM_PROMPT",)
FUNCTION = "load_image"
def load_image(self, image, model_type, ckpt, embedding_id, image_prompts_json):
def load_image(self, image, ckpt, embedding_id, image_prompts_json):
import json
global global_predictor
if global_predictor is None:
ckpt = folder_paths.get_full_path("sams", ckpt)
model_type = re.findall(r'vit_[lbh]', ckpt)[0]
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)
global_predictor = predictor
+3 -2
View File
@@ -1,6 +1,7 @@
import folder_paths
import os
import numpy as np
import re
from segment_anything import sam_model_registry, SamPredictor
@@ -19,7 +20,6 @@ class SAM_Prompt_Image:
return {
"required": {
"image": ("IMAGE",),
"model_type": (["vit_h", "vit_l", "vit_b"],),
"ckpt": (folder_paths.get_filename_list("sams"),),
"embedding_id": (
"STRING",
@@ -35,12 +35,13 @@ class SAM_Prompt_Image:
RETURN_TYPES = ("SAM_PROMPT",)
FUNCTION = "load_image"
def load_image(self, image, model_type, ckpt, embedding_id, image_prompts_json):
def load_image(self, image, ckpt, embedding_id, image_prompts_json):
import json
emb_filename = f"{self.output_dir}/{embedding_id}.npy"
if not os.path.exists(emb_filename):
ckpt = folder_paths.get_full_path("sams", ckpt)
model_type = re.findall(r'vit_[lbh]', ckpt)[0]
sam = sam_model_registry[model_type](checkpoint=ckpt)
predictor = SamPredictor(sam)