g-luo commited on
Commit
428226a
·
1 Parent(s): 3173d50

Fix ClipSeg error

Browse files
Files changed (1) hide show
  1. app.py +22 -22
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
- 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,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
- 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,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
- 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.")
 
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.")