92 lines
3.1 KiB
Python
92 lines
3.1 KiB
Python
# -*- coding: utf-8 -*-
|
|
"""Minimal wrapper around Coral EdgeTPU detection/classification."""
|
|
import os
|
|
import sys
|
|
|
|
try:
|
|
from pycoral.utils.edgetpu import make_interpreter
|
|
from pycoral.adapters import common, detect
|
|
except ImportError: # pragma: no cover
|
|
make_interpreter = None
|
|
common = None
|
|
detect = None
|
|
|
|
import numpy as np
|
|
from PIL import Image
|
|
|
|
|
|
class CoralVision(object):
|
|
def __init__(self, model_path, labels_path, score_threshold=0.4):
|
|
self.model_path = model_path
|
|
self.labels_path = labels_path
|
|
self.score_threshold = score_threshold
|
|
self.available = False
|
|
self.labels = {}
|
|
self.interpreter = None
|
|
self._load()
|
|
|
|
def _load_labels(self):
|
|
labels = {}
|
|
if not os.path.exists(self.labels_path):
|
|
return labels
|
|
with open(self.labels_path, 'r') as f:
|
|
for line in f:
|
|
if not line.strip():
|
|
continue
|
|
pair = line.strip().split(None, 1)
|
|
if len(pair) == 2 and pair[0].isdigit():
|
|
labels[int(pair[0])] = pair[1]
|
|
else:
|
|
# fallback for plain list
|
|
labels[len(labels)] = pair[0]
|
|
return labels
|
|
|
|
def _load(self):
|
|
if make_interpreter is None:
|
|
print('[WARN] pycoral not installed; Coral disabled', file=sys.stderr)
|
|
return
|
|
if not os.path.exists(self.model_path):
|
|
print('[WARN] Coral model not found at {} - EdgeTPU disabled'.format(self.model_path), file=sys.stderr)
|
|
return
|
|
try:
|
|
self.interpreter = make_interpreter(self.model_path)
|
|
self.interpreter.allocate_tensors()
|
|
self.labels = self._load_labels()
|
|
self.available = True
|
|
print('[INFO] Coral model loaded: {}'.format(self.model_path))
|
|
except Exception as exc:
|
|
print('[WARN] Failed to init Coral: {}'.format(exc), file=sys.stderr)
|
|
self.available = False
|
|
|
|
def infer(self, img_rgb):
|
|
if not self.available:
|
|
return [], False
|
|
|
|
h, w, _ = img_rgb.shape
|
|
input_size = common.input_size(self.interpreter)
|
|
resized = _resize_keep_aspect(img_rgb, input_size)
|
|
common.set_input(self.interpreter, resized)
|
|
self.interpreter.invoke()
|
|
objs = detect.get_objects(self.interpreter, score_threshold=self.score_threshold)
|
|
|
|
results = []
|
|
for obj in objs:
|
|
bbox = obj.bbox
|
|
scale_x = float(w) / float(input_size[0])
|
|
scale_y = float(h) / float(input_size[1])
|
|
x1 = int(bbox.xmin * scale_x)
|
|
y1 = int(bbox.ymin * scale_y)
|
|
x2 = int(bbox.xmax * scale_x)
|
|
y2 = int(bbox.ymax * scale_y)
|
|
label = self.labels.get(obj.id, str(obj.id))
|
|
results.append({
|
|
"label": label,
|
|
"score": float(obj.score),
|
|
"bbox": [x1, y1, x2, y2],
|
|
})
|
|
return results, True
|
|
|
|
|
|
def _resize_keep_aspect(img_rgb, target_size):
|
|
tw, th = target_size
|
|
return np.array(Image.fromarray(img_rgb).resize((tw, th)))
|