feat: infer model type from checkpoint name
This commit is contained in:
@@ -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(() => {
|
||||
|
||||
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user