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]