File size: 2,712 Bytes
7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 50ce3c9 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 ade4ba6 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 7206ed3 752f8f0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 | import math
import numpy as np
import PIL
import torch
############################################
# Predict detections with MegaDetector v5a model
def predict_md(
im,
megadetector_model, # Megadet_Models[mega_model_input]
size=640,
):
# resize image
g = size / max(im.size) # multipl factor to make max size of the image equal to input size
im = im.resize((int(x * g) for x in im.size), PIL.Image.Resampling.LANCZOS) # resize
# device: yolov5's select_device expects a CUDA index ('0') or 'cpu', not 'cuda'
md_device = "0" if torch.cuda.is_available() else "cpu"
# megadetector
MD_model = torch.hub.load(
"ultralytics/yolov5", # repo_or_dir
"custom", # model
megadetector_model, # args for callable model
skip_validation=True, # avoid GitHub API rate limit (403)
device=md_device,
trust_repo=True,
)
## detect objects
# vars(results).keys(): imgs, pred, names, files, times, xyxy, xywh, xyxyn, xywhn, n, t, s
results = MD_model(im)
return results
##########################################
def crop_animal_detections(img_in, yolo_results, likelihood_th):
## Extract animal crops
list_labels_as_str = [i for i in yolo_results.names.values()] # ['animal', 'person', 'vehicle']
list_np_animal_crops = []
list_animal_bboxes = [] # detection rows [x1,y1,x2,y2,conf,label] matching each crop
# image to crop (scale as input for megadetector)
img_in = img_in.resize((yolo_results.ims[0].shape[1], yolo_results.ims[0].shape[0]))
# for every detection in the img
for det_array in yolo_results.xyxy:
# for every detection
for j in range(det_array.shape[0]):
# compute coords around bbox rounded to the nearest integer (for pasting later)
xmin_rd = int(math.floor(det_array[j, 0])) # int() should suffice?
ymin_rd = int(math.floor(det_array[j, 1]))
xmax_rd = int(math.ceil(det_array[j, 2]))
ymax_rd = int(math.ceil(det_array[j, 3]))
pred_llk = det_array[j, 4]
pred_label = det_array[j, 5]
# keep animal crops above threshold
if (pred_label == list_labels_as_str.index("animal")) and (pred_llk >= likelihood_th):
area = (xmin_rd, ymin_rd, xmax_rd, ymax_rd)
# pdb.set_trace()
crop = img_in.crop(area) # Image.fromarray(img_in).crop(area)
crop_np = np.asarray(crop)
# add to list
list_np_animal_crops.append(crop_np)
list_animal_bboxes.append(det_array[j, :].tolist())
return list_np_animal_crops, list_animal_bboxes
|