如何加载自定义YOLOv7训练模型并获取目标边界框坐标
加载自定义YOLOv7模型并获取检测框坐标
YOLOv7的加载和推理逻辑和YOLOv5高度相似,你可以沿用类似YOLOv5的编程式加载方法,无需通过detect.py命令行执行,直接在代码中加载模型并提取边界框坐标。
具体实现步骤
1. 加载自定义YOLOv7权重
YOLOv7同样提供了attempt_load函数用于加载预训练或自定义训练的权重,用法和YOLOv5一致:
import torch from models.experimental import attempt_load from utils.general import non_max_suppression, scale_coords from utils.torch_utils import select_device # 选择计算设备(优先GPU) device = select_device('0' if torch.cuda.is_available() else 'cpu') # 替换为你的YOLOv7自定义权重路径 yolov7_weight_file = r'runs/train/yolov7x-custom/weights/best.pt' model = attempt_load(yolov7_weight_file, map_location=device) model.to(device).eval() # 获取模型对应的类别名称 names = model.module.names if hasattr(model, 'module') else model.names
2. 编写检测函数处理图像帧
参考你YOLOv5的实现,修改为YOLOv7适配版本,重点处理图像预处理和结果解析:
import cv2 import numpy as np conf_set = 0.5 # 置信度阈值 iou_thres = 0.20 # IOU阈值 def object_detection(frame): # 图像预处理:转换为YOLO系列要求的RGB格式+张量格式 img = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB) img = torch.from_numpy(img).to(device) img = img.permute(2, 0, 1).float() / 255.0 # 转换为[C, H, W]并归一化到0-1 if img.ndimension() == 3: img = img.unsqueeze(0) # 添加batch维度 # 模型推理(关闭梯度计算加速) with torch.no_grad(): pred = model(img, augment=False)[0] # 非极大值抑制去除重复检测框 pred = non_max_suppression(pred, conf_set, iou_thres) detection_result = [] # 解析检测结果 for det in pred: if len(det): # 将模型输出的缩放后坐标映射回原始图像尺寸 det[:, :4] = scale_coords(img.shape[2:], det[:, :4], frame.shape).round() 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] detection_result.append([x1, y1, x2, y2, conf, c]) # 在图像上绘制边界框和标签 frame = cv2.rectangle(frame, (x1, y1), (x2, y2), (255, 0, 0), 1) if c != 1: # 按需跳过特定类别的标签绘制 frame = cv2.putText(frame, f'{detected_name} {conf}', (x1, y1), cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 0, 255), 1, cv2.LINE_AA) return frame, detection_result
3. 检测函数调用示例
# 读取测试图像 test_img = cv2.imread('test.jpg') # 执行检测 result_img, dets = object_detection(test_img) # 打印检测到的边界框信息 for det in dets: x1, y1, x2, y2, conf, cls_idx = det print(f"类别: {names[cls_idx]}, 置信度: {conf}, 边界框: ({x1}, {y1}), ({x2}, {y2})") # 显示结果图像 cv2.imshow('Detection Result', result_img) cv2.waitKey(0) cv2.destroyAllWindows()
关键注意事项
- 确保YOLOv7代码仓库结构完整,
models、utils文件夹未缺失,避免导入错误 scale_coords函数是必要步骤,用于将模型输出的缩放后坐标映射回原始图像尺寸- 开启
torch.no_grad()可以大幅提升推理速度
内容的提问来源于stack exchange,提问作者pavan
相关产品推荐
相关产品推荐

