如何修改PyTorch车牌检测代码适配OpenVINO的XML/BIN模型
适配OpenVINO模型的车牌检测代码修改方案
核心修改步骤
1. 替换依赖库导入
移除PyTorch相关导入,新增OpenVINO运行时库,保留原YOLO后处理所需的非极大值抑制(NMS)函数:
import cv2 import numpy as np import torch # 仅用于转换推理结果格式以复用原NMS函数 from yolov5.utils.general import non_max_suppression # 若无YOLO依赖可自行实现NMS import openvino.runtime as ov
2. 替换模型加载逻辑
删除PyTorch的模型加载代码,改用OpenVINO加载xml/bin格式模型:
# OpenVINO模型文件路径 model_xml = r'D:\deepak\Helmet-Detection-final\model\rider_helmet_number_medium.xml' model_bin = r'D:\deepak\Helmet-Detection-final\model\rider_helmet_number_medium.bin' # 初始化OpenVINO核心并加载编译模型 core = ov.Core() model = core.read_model(model=model_xml, weights=model_bin) # 编译到指定设备:CPU/GPU/MYRIAD(根据硬件调整) compiled_model = core.compile_model(model=model, device_name="CPU") # 获取输入输出节点信息 input_layer = compiled_model.input(0) output_layer = compiled_model.output(0) # 类别名称:需与原PyTorch模型的类别顺序完全一致,可从原模型导出的txt读取 names = ["rider", "helmet", "license_plate"] # 替换为实际类别列表
3. 调整图像预处理逻辑
将原PyTorch张量预处理转换为numpy操作,适配OpenVINO输入格式:
def license_plate(frame): try: # 图像格式转换:OpenCV默认BGR转RGB,HWC转NCHW(与原PyTorch预处理对齐) img = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) img = img.transpose(2, 0, 1) # HWC -> CHW img = img.astype(np.float32) / 255.0 # 归一化到0-1区间 img = np.expand_dims(img, axis=0) # 增加batch维度,变为NCHW格式 # 执行OpenVINO推理 pred = compiled_model([img])[output_layer] # 转换为PyTorch张量,复用原NMS函数 pred = torch.from_numpy(pred) # 复用原NMS及结果处理逻辑 conf_set = 0.25 # 替换为原代码使用的置信度阈值 pred = non_max_suppression(pred, conf_set, 0.30) detection_result = [] for i, det in enumerate(pred): if len(det): for d in det: # d = (x1, y1, x2, y2, conf, cls) x1 = int(d[0].item()) y1 = int(d[1].item()) x2 = int(d[2].item()) y2 = int(d[3].item()) conf = round(d[4].item(), 2) c = int(d[5].item()) detected_name = names[c] print(f'Detected: {detected_name} conf: {conf} bbox: x1:{x1} y1:{y1} x2:{x2} y2:{y2}') detection_result.append([x1, y1, x2, y2, conf, c]) if c == 0 or c == 2: frame = cv2.rectangle(frame, (x1, y1), (x2, y2), (255, 0, 0), 1) frame = cv2.putText(frame, f'{[c]} {str(conf)}', (x1, y1), cv2.FONT_HERSHEY_SIMPLEX, 0.5,(0, 0, 255), 1, cv2.LINE_AA) return frame except Exception as ex: print(ex)
4. 关键注意事项
- 类别名称一致性:确保
names列表的顺序与原PyTorch模型完全匹配,可从原模型导出的names.txt文件读取,避免类别识别错误。 - 输入格式适配:若转换OpenVINO模型时指定了NHWC格式输入,需移除
transpose(2,0,1)步骤。 - 设备优化:根据硬件情况调整
device_name,如支持GPU则改为"GPU",可提升推理速度。 - NMS替代:若无YOLO依赖,可自行实现NMS函数,核心逻辑为过滤低置信度框、计算IOU并去除重叠框。
内容的提问来源于stack exchange,提问作者Deepak s
相关产品推荐
相关产品推荐

