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

如何加载自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.20 08:04:41