如何在YOLOv8+SORT目标检测跟踪中获取带ID的类别信息?
问题分析与解决方案
你当前的核心问题在于:SORT算法仅接收检测框坐标+置信度的数组,完全不知道类别信息;且SORT输出的跟踪结果顺序和YOLO原检测框顺序没有对应关系,直接用索引匹配类别自然会出错。下面是具体的修正方案:
解决思路
- 保存检测框完整信息:每帧处理时,同时存储符合条件的检测框的坐标、置信度、类别,而不是只生成SORT需要的5列数组。
- 跟踪框与检测框匹配:拿到SORT的跟踪结果后,通过IOU计算找到每个跟踪框对应的最匹配检测框,从而获取类别。
- 维护跟踪ID的类别字典:用字典记录每个跟踪ID对应的最新类别,方便后续统计未穿戴装备的人员数量。
修改后的代码
import cv2 import numpy as np import math from sort import Sort from ultralytics import YOLO class_names = ['Hardhat', 'Mask', 'NO-Hardhat', 'NO-Mask', 'NO-Safety Vest', 'Person', 'Safety Cone', 'Safety Vest', 'machinery', 'vehicle'] def iou(box1, box2): # 计算两个框的IOU x1 = max(box1[0], box2[0]) y1 = max(box1[1], box2[1]) x2 = min(box1[2], box2[2]) y2 = min(box1[3], box2[3]) inter_area = max(0, x2 - x1) * max(0, y2 - y1) box1_area = (box1[2] - box1[0]) * (box1[3] - box1[1]) box2_area = (box2[2] - box2[0]) * (box2[3] - box2[1]) return inter_area / (box1_area + box2_area - inter_area) def process_video(video_path: str, model: YOLO): # 调整SORT参数,默认参数更合理 tracker = Sort(max_age=30, min_hits=3, iou_threshold=0.3) cap = cv2.VideoCapture(video_path) # 存储每个跟踪ID对应的最新类别 id_to_class = {} while True: ret, img = cap.read() if not ret: break results = model(img, stream=True) detections = np.empty((0,5)) # 保存当前帧所有符合条件的检测框完整信息(框坐标、置信度、类别) current_detections_with_cls = [] for r in results: boxes = r.boxes for box in boxes: x1, y1, x2, y2 = box.xyxy[0] x1, y1, x2, y2 = int(x1), int(y1), int(x2), int(y2) conf = math.ceil((box.conf[0] * 100)) cls_idx = int(box.cls[0]) cls = class_names[cls_idx] if cls in ['Person','NO-Hardhat', 'NO-Mask', 'NO-Safety Vest']: currentArray = np.array([x1,y1,x2,y2,conf]) detections = np.vstack((detections,currentArray)) current_detections_with_cls.append({ "box": [x1,y1,x2,y2], "conf": conf, "class": cls }) resultsTracker = tracker.update(detections) for result in resultsTracker: x1, y1, x2, y2, id = result x1, y1, x2, y2, id = int(x1), int(y1), int(x2), int(y2), int(id) # 找到当前跟踪框匹配的检测框,获取类别 max_iou = 0 matched_cls = "Person" # 默认类别 for det in current_detections_with_cls: current_iou = iou([x1,y1,x2,y2], det["box"]) if current_iou > max_iou: max_iou = current_iou matched_cls = det["class"] # 更新ID对应的类别 id_to_class[id] = matched_cls # 绘制跟踪框和ID+类别 (text_width, text_height), _ = cv2.getTextSize(f'{id}:{matched_cls}', cv2.FONT_HERSHEY_PLAIN, fontScale=0.5, thickness=1) text_offset_x = x1 text_offset_y = y1 - text_height cv2.rectangle(img,(x1,y1),(x2,y2),(0,0,255),1) cv2.putText(img, f'{id}:{matched_cls}', (text_offset_x, text_offset_y+6), cv2.FONT_HERSHEY_PLAIN, fontScale=0.5, color=(255, 255, 255), thickness=1) # 示例:统计当前未穿戴防护装备的人员数量 unsafe_count = sum(1 for cls in id_to_class.values() if cls.startswith("NO-")) cv2.putText(img, f'Unsafe: {unsafe_count}', (20,20), cv2.FONT_HERSHEY_SIMPLEX, 0.7, (0,0,255), 2) cv2.imshow('video',img) if cv2.waitKey(1) & 0xFF == ord('q'): break cap.release() cv2.destroyAllWindows()
额外建议
- SORT参数调整:你原代码的
iou_threshold=0.01过低,会导致跟踪逻辑混乱,建议设为0.3-0.5;max_age=2000过大,会保留大量已消失的跟踪ID,建议改为30-50。 - 类别判断优化:直接用类别ID判断比字符串匹配更高效,比如
if cls_idx in [2,3,4,5](对应NO-Hardhat、NO-Mask、NO-Safety Vest、Person)。 - 统计逻辑优化:如果需要统计全程未穿戴装备的人员,可在字典中记录每个ID的历史类别,只要出现过NO-开头的类别就计入统计。
内容的提问来源于stack exchange,提问作者Deepak Pawade
相关产品推荐
相关产品推荐

