ppbbww

Pillar Point boats, birds, and waves watcher
git clone git@abtrout.com:ppbbww.git
Log | Files | Refs | README | LICENSE

object_detector.py (1558B)


      1 import logging
      2 import sys
      3 import torch
      4 
      5 from PIL import Image, ImageDraw
      6 from time import perf_counter
      7 from transformers import DetrImageProcessor, DetrForObjectDetection
      8 
      9 
     10 class ObjectDetector:
     11     def __init__(self, thresh=0.95, labels=None):
     12         self.thresh = thresh
     13         self.filter_labels = labels
     14 
     15         t0 = perf_counter()
     16         self.device = "cuda:0" if torch.cuda.is_available() else "cpu"
     17         self.processor = DetrImageProcessor.from_pretrained(
     18             "facebook/detr-resnet-50", revision="no_timm"
     19         )
     20         self.model = DetrForObjectDetection.from_pretrained(
     21             "facebook/detr-resnet-50", revision="no_timm"
     22         ).to(self.device)
     23         logging.info(
     24             f"ObjectDetector initialized for device {self.device}; took {perf_counter() - t0} seconds."
     25         )
     26 
     27     def find(self, img):
     28         res = self.__match_results(img)
     29         for label, score, box in zip(res["labels"], res["scores"], res["boxes"]):
     30             label = self.model.config.id2label[label.item()]
     31             if (not self.filter_labels) or (label in self.filter_labels):
     32                 yield (label, round(score.item(), 2), map(int, box.tolist()))
     33 
     34     def __match_results(self, image):
     35         inputs = self.processor(images=image, return_tensors="pt").to(self.device)
     36         outputs = self.model(**inputs)
     37         target_sizes = torch.tensor([image.size[::-1]]).to(self.device)
     38         return self.processor.post_process_object_detection(
     39             outputs, target_sizes=target_sizes, threshold=self.thresh
     40         )[0]