在OpenVINO Runtime运行YOLOv8输出形状异常,求检测框提取方案
问题
我用官方的yolov8n.pt作为基础模型训练了目标检测模型,但在Intel CPU上推理速度太慢,没法落地。于是用Intel的OpenVINO Runtime优化,通过mo命令把ONNX模型转成IR文件后加载运行,但转换后模型输出形状不符合预期,对比yolov8n.yaml和IR文件也搞不懂输出含义。换了yolov8s.pt和不同数据集后,OpenVINO优化后的输出形状还是固定的:
- 输入形状:
[1,3,640,640] - 输出形状:
[1,6,8400]
运行代码如下:
import cv2 import matplotlib.pyplot as plt import numpy as np from openvino.runtime import Core from ultralytics import YOLO ie = Core() model = ie.read_model(model=r"\models\IR\model_fit.xml") compiled_model = ie.compile_model(model=model, device_name="CPU") input_layer_ir = compiled_model.input(0) output_layer_ir = compiled_model.output() image = cv2.imread(r'\Resources\test_image.jpg') N, C, H, W = input_layer_ir.shape resized_image = cv2.resize(image, (W, H)) input_image = np.expand_dims(resized_image.transpose(2, 0, 1), 0) plt.imshow(cv2.cvtColor(image, cv2.COLOR_BGR2RGB)); boxes = compiled_model([input_image])[output_layer_ir] print(boxes.shape)
输出结果:
<Output: names[output0] shape[1,6,8400] type: f32>
我是计算机视觉新手,该怎么处理才能获取检测框坐标?
解决方案
1. 先理解输出形状的含义
[1,6,8400]的维度定义:
1:批量大小(当前为单张输入图)6:每个预测框的6个参数,顺序为 x中心坐标、y中心坐标、框宽度、框高度、目标置信度、类别ID8400:YOLOv8固定的锚框预测总数,不同模型结构下该值不变
2. 解析输出并处理检测框
步骤1:调整输出张量形状
先把输出转为更易处理的维度格式:
# 去掉批量维度,转置后得到(8400,6)的形状 predictions = boxes[0].transpose(1, 0)
步骤2:过滤低置信度预测框
设置置信度阈值(比如0.5),筛掉置信度不足的无效框:
conf_threshold = 0.5 valid_predictions = predictions[predictions[:, 4] > conf_threshold]
步骤3:将归一化坐标转成原图像素坐标
YOLOv8输出的坐标是相对于640×640输入图的归一化值,需要先转成输入图的像素坐标,再映射回原图尺寸:
orig_h, orig_w = image.shape[:2] input_w, input_h = W, H # 计算输入图到原图的缩放比例 scale_x = orig_w / input_w scale_y = orig_h / input_h processed_boxes = [] for pred in valid_predictions: x_center, y_center, w, h, conf, cls_id = pred # 把归一化坐标转为输入图的像素坐标,再映射到原图 x1 = int((x_center - w/2) * input_w * scale_x) y1 = int((y_center - h/2) * input_h * scale_y) x2 = int((x_center + w/2) * input_w * scale_x) y2 = int((y_center + h/2) * input_h * scale_y) processed_boxes.append([x1, y1, x2, y2, conf, cls_id])
步骤4:非极大值抑制(NMS)去除重复框
同一目标可能被多个框重复检测,用NMS去掉高重叠的冗余框:
nms_threshold = 0.5 if processed_boxes: boxes_np = np.array([b[:4] for b in processed_boxes]) confidences = np.array([b[4] for b in processed_boxes]) # 调用OpenCV的NMS函数 indices = cv2.dnn.NMSBoxes(boxes_np.tolist(), confidences.tolist(), conf_threshold, nms_threshold) # 得到最终有效检测框 final_boxes = [processed_boxes[i] for i in indices.flatten()]
步骤5:完整可运行代码
整合所有步骤,最终可以在原图上绘制检测框:
import cv2 import matplotlib.pyplot as plt import numpy as np from openvino.runtime import Core ie = Core() model = ie.read_model(model=r"\models\IR\model_fit.xml") compiled_model = ie.compile_model(model=model, device_name="CPU") input_layer_ir = compiled_model.input(0) output_layer_ir = compiled_model.output() image = cv2.imread(r'\Resources\test_image.jpg') N, C, H, W = input_layer_ir.shape resized_image = cv2.resize(image, (W, H)) input_image = np.expand_dims(resized_image.transpose(2, 0, 1), 0) # 推理获取输出 boxes = compiled_model([input_image])[output_layer_ir] # 调整输出形状 predictions = boxes[0].transpose(1, 0) # 过滤低置信度框 conf_threshold = 0.5 valid_predictions = predictions[predictions[:, 4] > conf_threshold] # 坐标转换与映射 orig_h, orig_w = image.shape[:2] input_w, input_h = W, H scale_x = orig_w / input_w scale_y = orig_h / input_h processed_boxes = [] for pred in valid_predictions: x_center, y_center, w, h, conf, cls_id = pred x1 = int((x_center - w/2) * input_w * scale_x) y1 = int((y_center - h/2) * input_h * scale_y) x2 = int((x_center + w/2) * input_w * scale_x) y2 = int((y_center + h/2) * input_h * scale_y) processed_boxes.append([x1, y1, x2, y2, conf, cls_id]) # NMS去重 nms_threshold = 0.5 if processed_boxes: boxes_np = np.array([b[:4] for b in processed_boxes]) confidences = np.array([b[4] for b in processed_boxes]) indices = cv2.dnn.NMSBoxes(boxes_np.tolist(), confidences.tolist(), conf_threshold, nms_threshold) final_boxes = [processed_boxes[i] for i in indices.flatten()] # 在原图绘制检测框与标签 for box in final_boxes: x1, y1, x2, y2, conf, cls_id = box cv2.rectangle(image, (x1, y1), (x2, y2), (0, 255, 0), 2) label = f"Class {int(cls_id)}: {conf:.2f}" cv2.putText(image, label, (x1, y1-10), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0,255,0), 2) # 显示结果 plt.imshow(cv2.cvtColor(image, cv2.COLOR_BGR2RGB)) plt.show()
内容的提问来源于stack exchange,提问作者wizer_102
相关产品推荐
相关产品推荐

