Spaces:
Runtime error
Runtime error
Update predict.py
Browse files- 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
|