Spaces:
Build error
Build error
Fix ClipSeg error
Browse files
app.py
CHANGED
|
@@ -41,7 +41,7 @@ radio_options = [
|
|
| 41 |
"Upload mask",
|
| 42 |
"Draw mask above",
|
| 43 |
"Infer mask with MaskFormer",
|
| 44 |
-
"Infer mask with CLIPSeg"
|
| 45 |
]
|
| 46 |
presets = {
|
| 47 |
"custom": [None, None, None, None],
|
|
@@ -58,15 +58,15 @@ hparams_presets = {
|
|
| 58 |
# ==========================
|
| 59 |
# CLIPSeg Init
|
| 60 |
# ==========================
|
| 61 |
-
clip_seg_model = CLIPDensePredT(version='ViT-B/16', reduce_dim=64)
|
| 62 |
-
clip_seg_model.eval()
|
| 63 |
-
clip_seg_model.load_state_dict(torch.load('clipseg/weights/rd64-uni.pth'), strict=False)
|
| 64 |
-
clip_seg_model.to(device)
|
| 65 |
-
clip_seg_transform = transforms.Compose([
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
])
|
| 70 |
|
| 71 |
# ==========================
|
| 72 |
# MaskFormer Init
|
|
@@ -125,16 +125,16 @@ def activate_mask_upload(radio):
|
|
| 125 |
is_upload = radio == radio_options[0]
|
| 126 |
return gr.update(interactive=is_upload)
|
| 127 |
|
| 128 |
-
def infer_clip_seg(img, source_prompt, threshold=100):
|
| 129 |
-
|
| 130 |
-
|
| 131 |
-
|
| 132 |
-
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
|
| 137 |
-
|
| 138 |
|
| 139 |
def infer_maskformer(img, source_prompt):
|
| 140 |
category_mapping = {cat["name"]: i for i, cat in enumerate(COCO_CATEGORIES)}
|
|
@@ -158,8 +158,8 @@ def get_mask(radio, image_upload, mask_upload, source_prompt, threshold=100):
|
|
| 158 |
mask = image_upload["mask"]
|
| 159 |
elif radio == radio_options[2]:
|
| 160 |
mask = infer_maskformer(image_upload["image"], source_prompt)
|
| 161 |
-
elif radio == radio_options[3]:
|
| 162 |
-
|
| 163 |
|
| 164 |
if mask is None:
|
| 165 |
raise gr.Error("Missing input mask. Try running Update Mask again.")
|
|
|
|
| 41 |
"Upload mask",
|
| 42 |
"Draw mask above",
|
| 43 |
"Infer mask with MaskFormer",
|
| 44 |
+
# "Infer mask with CLIPSeg"
|
| 45 |
]
|
| 46 |
presets = {
|
| 47 |
"custom": [None, None, None, None],
|
|
|
|
| 58 |
# ==========================
|
| 59 |
# CLIPSeg Init
|
| 60 |
# ==========================
|
| 61 |
+
# clip_seg_model = CLIPDensePredT(version='ViT-B/16', reduce_dim=64)
|
| 62 |
+
# clip_seg_model.eval()
|
| 63 |
+
# clip_seg_model.load_state_dict(torch.load('clipseg/weights/rd64-uni.pth'), strict=False)
|
| 64 |
+
# clip_seg_model.to(device)
|
| 65 |
+
# clip_seg_transform = transforms.Compose([
|
| 66 |
+
# transforms.ToTensor(),
|
| 67 |
+
# transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),
|
| 68 |
+
# transforms.Resize((512, 512)),
|
| 69 |
+
# ])
|
| 70 |
|
| 71 |
# ==========================
|
| 72 |
# MaskFormer Init
|
|
|
|
| 125 |
is_upload = radio == radio_options[0]
|
| 126 |
return gr.update(interactive=is_upload)
|
| 127 |
|
| 128 |
+
# def infer_clip_seg(img, source_prompt, threshold=100):
|
| 129 |
+
# img = clip_seg_transform(img).unsqueeze(0)
|
| 130 |
+
# with torch.no_grad():
|
| 131 |
+
# preds = clip_seg_model(img, [source_prompt])[0]
|
| 132 |
+
# mask_preds = torch.sigmoid(preds[0][0]).detach().cpu().numpy()
|
| 133 |
+
# mask_preds = mask_preds * 255.0
|
| 134 |
+
# mask_preds = np.where(mask_preds > threshold, 255, 0)
|
| 135 |
+
# mask_preds = np.uint8(mask_preds)
|
| 136 |
+
# mask_preds = Image.fromarray(mask_preds, "L")
|
| 137 |
+
# return mask_preds
|
| 138 |
|
| 139 |
def infer_maskformer(img, source_prompt):
|
| 140 |
category_mapping = {cat["name"]: i for i, cat in enumerate(COCO_CATEGORIES)}
|
|
|
|
| 158 |
mask = image_upload["mask"]
|
| 159 |
elif radio == radio_options[2]:
|
| 160 |
mask = infer_maskformer(image_upload["image"], source_prompt)
|
| 161 |
+
# elif radio == radio_options[3]:
|
| 162 |
+
# mask = infer_clip_seg(image_upload["image"], source_prompt)
|
| 163 |
|
| 164 |
if mask is None:
|
| 165 |
raise gr.Error("Missing input mask. Try running Update Mask again.")
|