如何用Python加载YOLO-NAS的TensorRT .engine并获取检测边界框?
加载YOLO-NAS TensorRT引擎并提取检测框坐标的Python脚本
由于YOLO-NAS的输出格式和传统YOLO系列存在差异,之前通用的TensorRT推理脚本可能因未适配其输出结构导致无法解析出边界框。以下是针对性的完整实现:
依赖安装
先确保安装必要依赖:
pip install tensorrt opencv-python numpy
完整推理脚本
import tensorrt as trt import cv2 import numpy as np class YOLONASTRTInfer: def __init__(self, engine_path, input_shape=(640, 640), conf_thres=0.5, iou_thres=0.5): self.input_shape = input_shape self.conf_thres = conf_thres self.iou_thres = iou_thres self.logger = trt.Logger(trt.Logger.WARNING) self.engine = self.load_engine(engine_path) self.context = self.engine.create_execution_context() # 获取输入输出张量信息 self.input_tensor_name = self.engine.get_tensor_name(0) self.output_tensor_name = self.engine.get_tensor_name(1) self.context.set_input_shape(self.input_tensor_name, (1, 3, *input_shape)) def load_engine(self, engine_path): with open(engine_path, 'rb') as f, trt.Runtime(self.logger) as runtime: return runtime.deserialize_cuda_engine(f.read()) def preprocess(self, img): # 图像预处理:Resize、归一化、转CHW格式 h, w = img.shape[:2] img_resized = cv2.resize(img, self.input_shape) img_rgb = cv2.cvtColor(img_resized, cv2.COLOR_BGR2RGB) img_normalized = img_rgb.astype(np.float32) / 255.0 img_transposed = np.transpose(img_normalized, (2, 0, 1)) return np.expand_dims(img_transposed, axis=0), (w, h) def postprocess(self, output, original_size): original_w, original_h = original_size # YOLO-NAS输出格式:[1, num_boxes, 4 + 1 + num_classes] # 4=框坐标(cx, cy, w, h),1=置信度,后面是类别概率 boxes = output[0][:, :4] confs = output[0][:, 4:5] class_probs = output[0][:, 5:] # 计算最高类别概率和对应的类别ID class_ids = np.argmax(class_probs, axis=1).reshape(-1, 1) max_probs = np.max(class_probs, axis=1).reshape(-1, 1) # 过滤低置信度框:置信度×类别概率 > 阈值 filter_mask = (confs * max_probs) > self.conf_thres filtered_boxes = boxes[filter_mask[:, 0]] filtered_scores = (confs * max_probs)[filter_mask[:, 0]] filtered_class_ids = class_ids[filter_mask[:, 0]] # 将相对坐标转换为原图绝对坐标 # cx, cy, w, h -> x1, y1, x2, y2 cx, cy, w, h = filtered_boxes[:, 0], filtered_boxes[:, 1], filtered_boxes[:, 2], filtered_boxes[:, 3] x1 = (cx - w/2) * original_w y1 = (cy - h/2) * original_h x2 = (cx + w/2) * original_w y2 = (cy + h/2) * original_h # 非极大值抑制(NMS)去除重复框 indices = cv2.dnn.NMSBoxes(x1.tolist(), y1.tolist(), x2.tolist(), y2.tolist(), filtered_scores.tolist(), self.conf_thres, self.iou_thres) # 整理最终结果 final_boxes = [] for i in indices: i = i[0] if isinstance(i, (list, np.ndarray)) else i final_boxes.append({ 'bbox': [x1[i], y1[i], x2[i], y2[i]], 'score': float(filtered_scores[i]), 'class_id': int(filtered_class_ids[i]) }) return final_boxes def infer(self, img): input_data, original_size = self.preprocess(img) # 分配输入输出内存 bindings = [] output_memory = None for binding in self.engine: binding_idx = self.engine.get_binding_index(binding) size = trt.volume(self.context.get_binding_shape(binding_idx)) dtype = trt.nptype(self.engine.get_binding_dtype(binding)) if self.engine.binding_is_input(binding): input_memory = np.ascontiguousarray(input_data) bindings.append(input_memory.ctypes.data) else: output_memory = np.zeros(size, dtype=dtype) bindings.append(output_memory.ctypes.data) # 执行推理 self.context.execute_v2(bindings) # 获取输出数据 output = np.reshape(output_memory, (1, -1, 4 + 1 + self.engine.num_classes)) # 后处理得到检测框 return self.postprocess(output, original_size) # 示例用法 if __name__ == '__main__': # 替换为你的.engine文件路径 engine_path = "yolo_nas.engine" inferencer = YOLONASTRTInfer(engine_path) # 读取图像(也可以替换为摄像头实时帧) img = cv2.imread("test.jpg") # 或者实时摄像头: # cap = cv2.VideoCapture(0) # while True: # ret, img = cap.read() # if not ret: # break # results = inferencer.infer(img) results = inferencer.infer(img) # 打印检测框信息 for idx, result in enumerate(results): print(f"检测目标{idx+1}:") print(f"边界框坐标:{result['bbox']}") print(f"置信度:{result['score']}") print(f"类别ID:{result['class_id']}\n") # 可视化检测框 for result in results: x1, y1, x2, y2 = map(int, result['bbox']) cv2.rectangle(img, (x1, y1), (x2, y2), (0, 255, 0), 2) cv2.putText(img, f"Class {result['class_id']}: {result['score']:.2f}", (x1, y1-10), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0,255,0), 2) cv2.imshow("YOLO-NAS TensorRT Detection", img) cv2.waitKey(0) cv2.destroyAllWindows()
关键注意事项
- 输入尺寸匹配:确保脚本中
input_shape与导出ONNX模型时设置的尺寸完全一致(比如YOLO-NAS默认640x640) - 输出结构适配:若你的YOLO-NAS版本输出格式不同(比如坐标是x1,y1,x2,y2而非cx,cy,w,h),需调整
postprocess中的坐标转换逻辑 - 置信度阈值:可根据实际场景修改
conf_thres和iou_thres参数,平衡检测精度和召回率
内容的提问来源于stack exchange,提问作者Naofel
相关产品推荐
相关产品推荐

