YOLOv8-Seg / ax_infer.py
Fangming Guo
Upload 47 files
f3370c1 verified
Raw
History Blame Contribute Delete
11.4 kB
#!/usr/bin/env python3
import os
import cv2
import numpy as np
from time import time
import argparse
import logging
import axengine as ort
logging.basicConfig(
level=logging.DEBUG,
format='[%(name)s] [%(asctime)s.%(msecs)03d] [%(levelname)s] %(message)s',
datefmt='%H:%M:%S'
)
logger = logging.getLogger("YOLOv8-Seg")
def infer_hw_layout(shape):
"""Infer input height, width and layout from model input shape."""
shape = list(shape)
if len(shape) == 4 and shape[-1] == 3:
h = int(shape[1] or 640)
w = int(shape[2] or 640)
return h, w, "NHWC"
if len(shape) == 4 and shape[1] == 3:
h = int(shape[2] or 640)
w = int(shape[3] or 640)
return h, w, "NCHW"
return 640, 640, "NCHW"
def softmax(x, axis=-1):
"""Compute softmax along axis."""
e_x = np.exp(x - np.max(x, axis=axis, keepdims=True))
return e_x / np.sum(e_x, axis=axis, keepdims=True)
def dfl_decode(box_pred, reg_max=16):
"""Decode DFL (Distribution Focal Loss) box predictions to ltrb distances."""
N = box_pred.shape[0]
box_pred = box_pred.reshape(N, 4, reg_max)
box_pred = softmax(box_pred, axis=-1)
proj = np.arange(reg_max, dtype=np.float32)
return np.sum(box_pred * proj, axis=-1) # (N, 4)
def process_mask(protos, mask_coeffs, boxes, model_shape, orig_shape, scale):
"""
Process masks from prototypes and coefficients.
Args:
protos: (nm, proto_h, proto_w) mask prototypes
mask_coeffs: (N, nm) mask coefficients
boxes: (N, 4) boxes in model coordinates (xyxy)
model_shape: (h, w) model input shape
orig_shape: (h, w) original image shape
scale: preprocessing scale factor
Returns:
masks: (N, orig_h, orig_w) binary masks
"""
nm, proto_h, proto_w = protos.shape
N = len(mask_coeffs)
if N == 0:
return np.zeros((0, orig_shape[0], orig_shape[1]), dtype=np.uint8)
masks = (mask_coeffs @ protos.reshape(nm, -1)).reshape(N, proto_h, proto_w)
masks = 1 / (1 + np.exp(-masks))
width_ratio = proto_w / model_shape[1]
height_ratio = proto_h / model_shape[0]
boxes_proto = boxes.copy()
boxes_proto[:, [0, 2]] *= width_ratio
boxes_proto[:, [1, 3]] *= height_ratio
for i, (x1, y1, x2, y2) in enumerate(boxes_proto.astype(int)):
x1, y1 = max(0, x1), max(0, y1)
x2, y2 = min(proto_w, x2), min(proto_h, y2)
masks[i, :y1, :] = 0
masks[i, y2:, :] = 0
masks[i, :, :x1] = 0
masks[i, :, x2:] = 0
masks_upsampled = np.zeros((N, model_shape[0], model_shape[1]), dtype=np.float32)
for i in range(N):
masks_upsampled[i] = cv2.resize(masks[i], (model_shape[1], model_shape[0]),
interpolation=cv2.INTER_LINEAR)
new_h = int(orig_shape[0] * scale)
new_w = int(orig_shape[1] * scale)
masks_cropped = masks_upsampled[:, :new_h, :new_w]
masks_final = np.zeros((N, orig_shape[0], orig_shape[1]), dtype=np.uint8)
for i in range(N):
mask_resized = cv2.resize(masks_cropped[i], (orig_shape[1], orig_shape[0]),
interpolation=cv2.INTER_LINEAR)
masks_final[i] = (mask_resized > 0.5).astype(np.uint8)
return masks_final
def main():
ap = argparse.ArgumentParser(description='YOLOv8-Seg Inference (AXERARuntime)')
ap.add_argument('--model-path', type=str, default='yolov8n-seg_640x640.axmodel')
ap.add_argument('--test-img', type=str, default='bus.jpg')
ap.add_argument('--img-save-path', type=str, default='result_yolov8_seg.jpg')
ap.add_argument('--score-thres', type=float, default=0.25)
ap.add_argument('--nms-thres', type=float, default=0.7)
ap.add_argument('--providers', type=str, default='AxEngineExecutionProvider')
opt = ap.parse_args()
if not os.path.exists(opt.model_path):
logger.error(f"Model not found: {opt.model_path}")
return
t0 = time()
providers = [p.strip() for p in opt.providers.split(",") if p.strip()] or None
sess = ort.InferenceSession(opt.model_path, providers=providers)
logger.debug(f"\033[1;31mLoad model time = {(time() - t0) * 1000:.2f} ms\033[0m")
inp = sess.get_inputs()[0]
input_name = inp.name
m_h, m_w, layout = infer_hw_layout(inp.shape)
img = cv2.imread(opt.test_img)
if img is None:
logger.error(f"Image not found or unreadable: {opt.test_img}")
return
# Preprocess
t0 = time()
orig_h, orig_w = img.shape[:2]
scale = min(m_h / orig_h, m_w / orig_w)
new_w, new_h = int(orig_w * scale), int(orig_h * scale)
resized = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_LINEAR)
padded = cv2.copyMakeBorder(
resized, 0, m_h - new_h, 0, m_w - new_w,
cv2.BORDER_CONSTANT, value=(127, 127, 127)
)
rgb = cv2.cvtColor(padded, cv2.COLOR_BGR2RGB)
input_tensor = rgb[None, ...].astype(np.uint8) if layout == "NHWC" else np.transpose(rgb, (2, 0, 1))[None, ...].astype(np.uint8)
logger.debug(f"\033[1;31mPre-process time = {(time() - t0) * 1000:.2f} ms\033[0m")
# Inference
t0 = time()
ort_outputs = sess.run(None, {input_name: input_tensor})
out_metas = sess.get_outputs()
logger.debug(f"\033[1;31mForward time = {(time() - t0) * 1000:.2f} ms\033[0m")
# Post-process
t0 = time()
strides = (8, 16, 32)
conf_raw = -np.log(1 / opt.score_thres - 1)
detections = []
all_mask_coeffs = []
output_items = []
for meta, data in zip(out_metas, ort_outputs):
shape = list(meta.shape)
if any(s is None or isinstance(s, str) for s in shape):
shape = list(data.shape)
output_items.append((data, shape))
# Process each scale (9 detection outputs + 1 proto = 10 total)
for scale_idx, stride in enumerate(strides):
box_idx = scale_idx * 3
cls_idx = scale_idx * 3 + 1
mask_idx = scale_idx * 3 + 2
if box_idx >= len(output_items) or cls_idx >= len(output_items) or mask_idx >= len(output_items):
continue
box_data, box_shape = output_items[box_idx]
cls_data, cls_shape = output_items[cls_idx]
mask_data, mask_shape = output_items[mask_idx]
H, W = box_shape[1], box_shape[2]
box_channels = box_shape[-1]
# Determine if DFL is used (YOLOv8: 4*reg_max = 64)
reg_max = None
if box_channels > 4 and box_channels % 4 == 0:
reg_max = box_channels // 4
box_data = box_data[0].reshape(-1, box_channels)
cls_data = cls_data[0].reshape(-1, cls_shape[-1])
mask_data = mask_data[0].reshape(-1, mask_shape[-1])
cls_scores = np.max(cls_data, axis=1)
cls_ids = np.argmax(cls_data, axis=1)
valid_mask = cls_scores >= conf_raw
if not np.any(valid_mask):
continue
v_box = box_data[valid_mask]
v_score = 1 / (1 + np.exp(-cls_scores[valid_mask]))
v_id = cls_ids[valid_mask]
v_mask = mask_data[valid_mask]
gy, gx = np.indices((H, W))
anchors = np.stack((gx.ravel(), gy.ravel()), axis=-1).astype(np.float32) + 0.5
anchors = anchors[valid_mask]
# Decode boxes (DFL -> ltrb -> xyxy)
if reg_max is not None:
v_box = dfl_decode(v_box, reg_max)
lt = v_box[:, :2]
rb = v_box[:, 2:]
x1y1 = anchors - lt
x2y2 = anchors + rb
boxes = np.hstack([x1y1, x2y2]) * stride
for i in range(len(boxes)):
detections.append([*boxes[i], v_score[i], v_id[i]])
all_mask_coeffs.append(v_mask[i])
logger.debug(f"\033[1;31mPost-process time = {(time() - t0) * 1000:.2f} ms\033[0m")
if len(detections) == 0:
logger.info("No detections found.")
cv2.imwrite(opt.img_save_path, img)
return
detections = np.array(detections)
all_mask_coeffs = np.array(all_mask_coeffs)
# NMS
xywh = detections[:, :4].copy()
xywh[:, 2] = xywh[:, 2] - xywh[:, 0]
xywh[:, 3] = xywh[:, 3] - xywh[:, 1]
indices = cv2.dnn.NMSBoxes(xywh.tolist(), detections[:, 4].tolist(), opt.score_thres, opt.nms_thres)
if len(indices) == 0:
logger.info("No detections after NMS.")
cv2.imwrite(opt.img_save_path, img)
return
indices = indices.flatten()
detections = detections[indices]
all_mask_coeffs = all_mask_coeffs[indices]
# Get proto output (last output)
proto = ort_outputs[-1][0] # (nm, proto_h, proto_w)
logger.debug(f"Proto shape: {proto.shape}")
# Process masks (before scaling boxes!)
boxes_model = detections[:, :4].copy()
masks = process_mask(proto, all_mask_coeffs, boxes_model, (m_h, m_w), (orig_h, orig_w), scale)
# Scale boxes to original image coordinates
detections[:, :4] = detections[:, :4] / scale
detections[:, [0, 2]] = np.clip(detections[:, [0, 2]], 0, orig_w)
detections[:, [1, 3]] = np.clip(detections[:, [1, 3]], 0, orig_h)
# COCO class names
coco_names = [
"person", "bicycle", "car", "motorcycle", "airplane", "bus", "train", "truck", "boat", "traffic light",
"fire hydrant", "stop sign", "parking meter", "bench", "bird", "cat", "dog", "horse", "sheep", "cow",
"elephant", "bear", "zebra", "giraffe", "backpack", "umbrella", "handbag", "tie", "suitcase", "frisbee",
"skis", "snowboard", "sports ball", "kite", "baseball bat", "baseball glove", "skateboard", "surfboard",
"tennis racket", "bottle", "wine glass", "cup", "fork", "knife", "spoon", "bowl", "banana", "apple",
"sandwich", "orange", "broccoli", "carrot", "hot dog", "pizza", "donut", "cake", "chair", "couch",
"potted plant", "bed", "dining table", "toilet", "tv", "laptop", "mouse", "remote", "keyboard",
"cell phone", "microwave", "oven", "toaster", "sink", "refrigerator", "book", "clock", "vase",
"scissors", "teddy bear", "hair drier", "toothbrush"
]
base_colors = [
(255, 42, 4), (235, 219, 11), (243, 243, 243), (183, 223, 0), (104, 31, 17),
(221, 111, 255), (79, 68, 255), (0, 237, 204), (68, 243, 0), (255, 0, 189),
(255, 180, 0), (186, 0, 221), (255, 255, 0), (0, 192, 38), (179, 255, 1),
(255, 36, 125), (104, 0, 123), (108, 27, 255), (47, 109, 252), (11, 255, 162),
]
logger.info(f"\033[1;32mDraw Results ({len(detections)} objects): \033[0m")
for i, det in enumerate(detections):
box = det[:4].astype(int)
conf = det[4]
cls_id = int(det[5])
color = [int(c) for c in base_colors[cls_id % len(base_colors)]]
cls_name = coco_names[cls_id] if cls_id < len(coco_names) else str(cls_id)
logger.info(f"({box[0]}, {box[1]}, {box[2]}, {box[3]}) -> {cls_name}: {conf:.2f}")
if i < len(masks):
mask = masks[i]
mask_colored = np.zeros_like(img)
mask_colored[mask > 0] = color
img = cv2.addWeighted(img, 1.0, mask_colored, 0.5, 0)
cv2.rectangle(img, (box[0], box[1]), (box[2], box[3]), color, 2)
label = f"{cls_name} {conf:.2f}"
cv2.putText(img, label, (box[0], box[1] - 10),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, color, 2)
cv2.imwrite(opt.img_save_path, img)
logger.info(f"Saved to {opt.img_save_path}")
if __name__ == "__main__":
main()