object-detection / helpers.py
drewinspirit's picture
Upload 7 files
4b45220 verified
Raw
History Blame Contribute Delete
7.5 kB
import numpy as np
import cv2
from PIL import Image, ImageDraw, ImageFont
from copy import deepcopy
import colorsys
class BoundBox:
def __init__(self, xmin, ymin, xmax, ymax, objness=None, classes=None):
self.xmin = xmin
self.ymin = ymin
self.xmax = xmax
self.ymax = ymax
self.objness = objness
self.classes = classes
self.label = -1
self.score = -1
def get_label(self):
if self.label == -1:
self.label = np.argmax(self.classes)
return self.label
def get_score(self):
if self.score == -1:
self.score = self.classes[self.get_label()]
return self.score
def _interval_overlap(interval_a, interval_b):
x1, x2 = interval_a
x3, x4 = interval_b
if x3 < x1:
return 0 if x4 < x1 else min(x2, x4) - x1
else:
return 0 if x2 < x3 else min(x2, x4) - x3
def _sigmoid(x):
return 1. / (1. + np.exp(-x))
def bbox_iou(box1, box2):
intersect_w = _interval_overlap([box1.xmin, box1.xmax], [box2.xmin, box2.xmax])
intersect_h = _interval_overlap([box1.ymin, box1.ymax], [box2.ymin, box2.ymax])
intersect = intersect_w * intersect_h
w1, h1 = box1.xmax - box1.xmin, box1.ymax - box1.ymin
w2, h2 = box2.xmax - box2.xmin, box2.ymax - box2.ymin
union = w1 * h1 + w2 * h2 - intersect
return float(intersect) / union
def preprocess_input(image_pil, net_h, net_w):
image = np.asarray(image_pil)
new_h, new_w, _ = image.shape
if (float(net_w) / new_w) < (float(net_h) / new_h):
new_h = (new_h * net_w) / new_w
new_w = net_w
else:
new_w = (new_w * net_h) / new_h
new_h = net_h
new_w, new_h = int(new_w), int(new_h)
resized = cv2.resize(image / 255., (new_w, new_h))
new_image = np.ones((net_h, net_w, 3)) * 0.5
new_image[int((net_h - new_h) // 2):int((net_h + new_h) // 2),
int((net_w - new_w) // 2):int((net_w + new_w) // 2), :] = resized
return np.expand_dims(new_image, 0)
def decode_netout(netout_, obj_thresh, anchors_, image_h, image_w, net_h, net_w):
netout_all = deepcopy(netout_)
boxes_all = []
for i in range(len(netout_all)):
netout = netout_all[i][0]
anchors = anchors_[i]
grid_h, grid_w = netout.shape[:2]
nb_box = 3
netout = netout.reshape((grid_h, grid_w, nb_box, -1))
nb_class = netout.shape[-1] - 5
boxes = []
netout[..., :2] = _sigmoid(netout[..., :2])
netout[..., 4:] = _sigmoid(netout[..., 4:])
netout[..., 5:] = netout[..., 4][..., np.newaxis] * netout[..., 5:]
netout[..., 5:] *= netout[..., 5:] > obj_thresh
for i in range(grid_h * grid_w):
row = i // grid_w
col = i % grid_w
for b in range(nb_box):
objectness = netout[row][col][b][4]
classes = netout[row][col][b][5:]
if (classes <= obj_thresh).all():
continue
x, y, w, h = netout[row][col][b][:4]
x = (col + x) / grid_w
y = (row + y) / grid_h
w = anchors[b][0] * np.exp(w) / net_w
h = anchors[b][1] * np.exp(h) / net_h
box = BoundBox(x - w / 2, y - h / 2, x + w / 2, y + h / 2, objectness, classes)
boxes.append(box)
boxes_all += boxes
boxes_all = correct_yolo_boxes(boxes_all, image_h, image_w, net_h, net_w)
return boxes_all
def correct_yolo_boxes(boxes_, image_h, image_w, net_h, net_w):
boxes = deepcopy(boxes_)
if (float(net_w) / image_w) < (float(net_h) / image_h):
new_w = net_w
new_h = (image_h * net_w) / image_w
else:
new_h = net_w
new_w = (image_w * net_h) / image_h
for i in range(len(boxes)):
x_offset = (net_w - new_w) / 2. / net_w
x_scale = float(new_w) / net_w
y_offset = (net_h - new_h) / 2. / net_h
y_scale = float(new_h) / net_h
boxes[i].xmin = int((boxes[i].xmin - x_offset) / x_scale * image_w)
boxes[i].xmax = int((boxes[i].xmax - x_offset) / x_scale * image_w)
boxes[i].ymin = int((boxes[i].ymin - y_offset) / y_scale * image_h)
boxes[i].ymax = int((boxes[i].ymax - y_offset) / y_scale * image_h)
return boxes
def do_nms(boxes_, nms_thresh, obj_thresh):
boxes = deepcopy(boxes_)
if len(boxes) > 0:
num_class = len(boxes[0].classes)
else:
return []
for c in range(num_class):
sorted_indices = np.argsort([-box.classes[c] for box in boxes])
for i in range(len(sorted_indices)):
index_i = sorted_indices[i]
if boxes[index_i].classes[c] == 0:
continue
for j in range(i + 1, len(sorted_indices)):
index_j = sorted_indices[j]
if bbox_iou(boxes[index_i], boxes[index_j]) >= nms_thresh:
boxes[index_j].classes[c] = 0
new_boxes = []
for box in boxes:
for i in range(num_class):
if box.classes[i] > obj_thresh:
box.label = i
box.score = box.classes[i]
new_boxes.append(box)
break
return new_boxes
def draw_boxes(image_, boxes, labels):
image = image_.copy()
image_w, image_h = image.size
try:
font = ImageFont.truetype(
font='/usr/share/fonts/truetype/liberation/LiberationMono-Regular.ttf',
size=np.floor(3e-2 * image_h + 0.5).astype('int32')
)
except Exception:
font = ImageFont.load_default()
thickness = (image_w + image_h) // 300
hsv_tuples = [(x / len(labels), 1., 1.) for x in range(len(labels))]
colors = list(map(lambda x: colorsys.hsv_to_rgb(*x), hsv_tuples))
colors = list(map(lambda x: (int(x[0] * 255), int(x[1] * 255), int(x[2] * 255)), colors))
np.random.seed(10101)
np.random.shuffle(colors)
np.random.seed(None)
for i, box in reversed(list(enumerate(boxes))):
c = box.get_label()
predicted_class = labels[c]
score = box.get_score()
top, left, bottom, right = box.ymin, box.xmin, box.ymax, box.xmax
label = '{} {:.2f}'.format(predicted_class, score)
draw = ImageDraw.Draw(image)
label_size = draw.textbbox((0, 0), label, font)
label_size = (label_size[2], label_size[3])
top = max(0, np.floor(top + 0.5).astype('int32'))
left = max(0, np.floor(left + 0.5).astype('int32'))
bottom = min(image_h, np.floor(bottom + 0.5).astype('int32'))
right = min(image_w, np.floor(right + 0.5).astype('int32'))
if top - label_size[1] >= 0:
text_origin = np.array([left, top - label_size[1]])
else:
text_origin = np.array([left, top + 1])
if right > left and bottom > top:
draw.rectangle([left, top, right, bottom], outline=colors[c], width=thickness)
draw.text(text_origin, label, fill=(0, 0, 0), font=font)
del draw
return image
def detect_image(image_pil, model, anchors, labels, obj_thresh=0.4, nms_thresh=0.45, net_h=416, net_w=416):
image_w, image_h = image_pil.size
new_image = preprocess_input(image_pil, net_h, net_w)
yolo_outputs = model.predict(new_image)
boxes = decode_netout(yolo_outputs, obj_thresh, anchors, image_h, image_w, net_h, net_w)
boxes = do_nms(boxes, nms_thresh, obj_thresh)
return draw_boxes(image_pil, boxes, labels)