You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.21 12:52:49