From 2b7fd659d7d23fd6dd3c2d14eff32d1a51056061 Mon Sep 17 00:00:00 2001 From: Tengyang Chen Date: Fri, 18 Aug 2023 17:06:44 +0800 Subject: [PATCH] small fix for sam node --- WAS_Node_Suite.py | 26 +++++++++++++++++--------- 1 file changed, 17 insertions(+), 9 deletions(-) diff --git a/WAS_Node_Suite.py b/WAS_Node_Suite.py index 830ec89..2b6d7e8 100644 --- a/WAS_Node_Suite.py +++ b/WAS_Node_Suite.py @@ -10802,7 +10802,7 @@ class WAS_SAM_Model_Loader: def INPUT_TYPES(self): return { "required": { - "model_size": (["ViT-H (91M)", "ViT-L (308M)", "ViT-B (636M)"], ), + "model_size": (["ViT-H", "ViT-L", "ViT-B"], ), } } @@ -10815,15 +10815,15 @@ class WAS_SAM_Model_Loader: conf = getSuiteConfig() model_filename_mapping = { - "ViT-H (91M)": "sam_vit_h_4b8939.pth", - "ViT-L (308M)": "sam_vit_l_0b3195.pth", - "ViT-B (636M)": "sam_vit_b_01ec64.pth", + "ViT-H": "sam_vit_h_4b8939.pth", + "ViT-L": "sam_vit_l_0b3195.pth", + "ViT-B": "sam_vit_b_01ec64.pth", } model_url_mapping = { - "ViT-H (91M)": conf['sam_model_vith_url'] if conf.__contains__('sam_model_vith_url') else r"https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth", - "ViT-L (308M)": conf['sam_model_vitl_url'] if conf.__contains__('sam_model_vitl_url') else r"https://dl.fbaipublicfiles.com/segment_anything/sam_vit_l_0b3195.pth", - "ViT-B (636M)": conf['sam_model_vitb_url'] if conf.__contains__('sam_model_vitb_url') else r"https://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pth", + "ViT-H": conf['sam_model_vith_url'] if conf.__contains__('sam_model_vith_url') else r"https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth", + "ViT-L": conf['sam_model_vitl_url'] if conf.__contains__('sam_model_vitl_url') else r"https://dl.fbaipublicfiles.com/segment_anything/sam_vit_l_0b3195.pth", + "ViT-B": conf['sam_model_vitb_url'] if conf.__contains__('sam_model_vitb_url') else r"https://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pth", } model_url = model_url_mapping[model_size] @@ -10850,8 +10850,16 @@ class WAS_SAM_Model_Loader: r = requests.get(model_url, allow_redirects=True) open(sam_file, 'wb').write(r.content) - from segment_anything import build_sam - sam_model = build_sam(checkpoint=sam_file) + from segment_anything import build_sam_vit_h, build_sam_vit_l, build_sam_vit_b + + if model_size == 'ViT-H': + sam_model = build_sam_vit_h(sam_file) + elif model_size == 'ViT-L': + sam_model = build_sam_vit_l(sam_file) + elif model_size == 'ViT-B': + sam_model = build_sam_vit_b(sam_file) + else: + raise ValueError(f"SAM model does not match the model_size: '{model_size}'.") return (sam_model, )