File size: 2,712 Bytes
7206ed3
 
2e1b62d
 
 
7206ed3
 
 
 
2e1b62d
 
 
 
 
 
7206ed3
2e1b62d
 
 
 
 
 
 
 
 
 
 
 
 
 
7206ed3
 
2e1b62d
 
 
 
7206ed3
 
 
2e1b62d
7206ed3
 
2e1b62d
7206ed3
2e1b62d
7206ed3
 
2e1b62d
 
7206ed3
 
 
 
2e1b62d
 
7206ed3
2e1b62d
 
7206ed3
2e1b62d
 
7206ed3
2e1b62d
7206ed3
 
2e1b62d
 
7206ed3
 
 
 
2e1b62d
7206ed3
2e1b62d
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