File size: 3,782 Bytes
e9eb33d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7c098ec
e9eb33d
7c098ec
 
e9eb33d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
import argparse
import glob
import os

import cv2 as cv
import numpy as np

here = os.path.dirname(os.path.abspath(__file__))
sz = 512


def build_anchors():
    scales = [2.0 ** (i / 3.0) for i in range(3)]
    aspects = [(1.0, 1.0), (1.4, 0.7), (0.7, 1.4)]
    base = []
    for s in scales:
        for aw, ah in aspects:
            base.append((32.0 * s * aw, 32.0 * s * ah))
    anchors = []
    for lvl in range(5):
        f = sz // (8 * 2 ** lvl)
        step = 8 * 2 ** lvl
        m = 2 ** lvl
        for y in range(f):
            for x in range(f):
                cx = (x + 0.5) * step
                cy = (y + 0.5) * step
                for bw, bh in base:
                    anchors.append((cx, cy, bw * m, bh * m))
    return np.array(anchors, np.float32)


def main():
    parser = argparse.ArgumentParser(description="EfficientDet-D0 (ONNX) object detection demo")
    parser.add_argument("--model", default=None)
    parser.add_argument("--image", default=os.path.join(here, "example_outputs", "input_image.png"))
    parser.add_argument("--output", default=os.path.join(here, "example_outputs", "output_image.png"))
    parser.add_argument("--conf", type=float, default=0.4)
    args = parser.parse_args()

    model = args.model
    if model is None:
        found = glob.glob(os.path.join(here, "*.onnx"))
        if not found:
            raise SystemExit("no onnx, run convert_to_onnx.py")
        model = found[0]

    img = cv.imread(args.image)
    if img is None:
        raise SystemExit("could not read image: %s" % args.image)

    anchors = build_anchors()
    acx, acy, aw, ah = anchors[:, 0], anchors[:, 1], anchors[:, 2], anchors[:, 3]

    net = cv.dnn.readNetFromONNX(model)
    inp = cv.resize(cv.cvtColor(img, cv.COLOR_BGR2RGB), (sz, sz))
    net.setInput(inp[None].astype(np.uint8))
    res = net.forward(net.getUnconnectedOutLayersNames())
    box = next(a for a in res if a.shape[-1] == 4).reshape(-1, 4)
    cls = next(a for a in res if a.shape[-1] != 4).reshape(box.shape[0], -1)

    ycenter = box[:, 0] * ah + acy
    xcenter = box[:, 1] * aw + acx
    bh = np.exp(box[:, 2]) * ah
    bw = np.exp(box[:, 3]) * aw
    boxes = np.stack([xcenter - bw / 2, ycenter - bh / 2, xcenter + bw / 2, ycenter + bh / 2], 1) / sz

    prob = 1.0 / (1.0 + np.exp(-cls))
    cid = prob.argmax(1)
    scores = prob.max(1)

    keep = scores > args.conf
    boxes = boxes[keep]
    scores = scores[keep]
    cid = cid[keep]
    order = scores.argsort()[::-1]
    pick = []
    while order.size:
        i = order[0]
        pick.append(i)
        xx1 = np.maximum(boxes[i, 0], boxes[order[1:], 0])
        yy1 = np.maximum(boxes[i, 1], boxes[order[1:], 1])
        xx2 = np.minimum(boxes[i, 2], boxes[order[1:], 2])
        yy2 = np.minimum(boxes[i, 3], boxes[order[1:], 3])
        iw = np.maximum(0, xx2 - xx1)
        ih = np.maximum(0, yy2 - yy1)
        inter = iw * ih
        ai = (boxes[i, 2] - boxes[i, 0]) * (boxes[i, 3] - boxes[i, 1])
        aj = (boxes[order[1:], 2] - boxes[order[1:], 0]) * (boxes[order[1:], 3] - boxes[order[1:], 1])
        iou = inter / (ai + aj - inter + 1e-9)
        order = order[1:][iou <= 0.6]

    print("efficientdet-d0", len(pick), "detections")
    h, w = img.shape[:2]
    for i in pick:
        x1, y1, x2, y2 = boxes[i]
        print(int(cid[i]), round(float(scores[i]), 3), round(float(x1), 3), round(float(y1), 3), round(float(x2), 3), round(float(y2), 3))
        cv.rectangle(img, (int(x1 * w), int(y1 * h)), (int(x2 * w), int(y2 * h)), (0, 255, 0), 2)
        cv.putText(img, "%d:%.2f" % (int(cid[i]), scores[i]), (int(x1 * w), int(y1 * h) - 5), cv.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 0), 1)
    cv.imwrite(args.output, img)
    print("wrote", args.output)


if __name__ == "__main__":
    main()