TheoBH commited on
Commit
4c7e6a7
·
verified ·
1 Parent(s): 7762f58

Update predict.py

Browse files
Files changed (1) hide show
  1. predict.py +5 -4
predict.py CHANGED
@@ -13,17 +13,18 @@ from detectron2.utils.visualizer import ColorMode, Visualizer
13
  from color_palette import ade_palette
14
  from transformers import Mask2FormerImageProcessor, Mask2FormerForUniversalSegmentation
15
 
 
16
  def load_model_and_processor(model_ckpt: str):
17
  device = "cuda" if torch.cuda.is_available() else "cpu"
18
  model = Mask2FormerForUniversalSegmentation.from_pretrained(model_ckpt).to(torch.device(device))
19
  model.eval()
20
  image_preprocessor = Mask2FormerImageProcessor.from_pretrained(model_ckpt)
21
  return model, image_preprocessor
22
-
23
  def load_default_ckpt():
24
  default_ckpt = "facebook/mask2former-swin-tiny-coco-panoptic"
25
  return default_ckpt
26
-
27
  def draw_panoptic_segmentation(predicted_segmentation_map, seg_info, image):
28
  metadata = MetadataCatalog.get("coco_2017_val_panoptic")
29
  for res in seg_info:
@@ -41,7 +42,7 @@ def draw_panoptic_segmentation(predicted_segmentation_map, seg_info, image):
41
  return output_img, labels
42
 
43
 
44
-
45
  def predict_masks(input_img_path: str):
46
 
47
  #load model and image processor
@@ -66,7 +67,7 @@ def predict_masks(input_img_path: str):
66
 
67
 
68
 
69
-
70
  def get_mask_for_label(results, label):
71
  import numpy as np
72
  from PIL import Image
 
13
  from color_palette import ade_palette
14
  from transformers import Mask2FormerImageProcessor, Mask2FormerForUniversalSegmentation
15
 
16
+ @spaces.GPU
17
  def load_model_and_processor(model_ckpt: str):
18
  device = "cuda" if torch.cuda.is_available() else "cpu"
19
  model = Mask2FormerForUniversalSegmentation.from_pretrained(model_ckpt).to(torch.device(device))
20
  model.eval()
21
  image_preprocessor = Mask2FormerImageProcessor.from_pretrained(model_ckpt)
22
  return model, image_preprocessor
23
+ @spaces.GPU
24
  def load_default_ckpt():
25
  default_ckpt = "facebook/mask2former-swin-tiny-coco-panoptic"
26
  return default_ckpt
27
+ @spaces.GPU
28
  def draw_panoptic_segmentation(predicted_segmentation_map, seg_info, image):
29
  metadata = MetadataCatalog.get("coco_2017_val_panoptic")
30
  for res in seg_info:
 
42
  return output_img, labels
43
 
44
 
45
+ @spaces.GPU
46
  def predict_masks(input_img_path: str):
47
 
48
  #load model and image processor
 
67
 
68
 
69
 
70
+ @spaces.GPU
71
  def get_mask_for_label(results, label):
72
  import numpy as np
73
  from PIL import Image