DeepSORT算法原理

DeepSORT算法的核心在于其对目标的外观特征和运动特征的联合使用,以及对目标匹配问题的优化处理。该算法通过融合目标检测的结果,结合匈牙利算法和卡尔曼滤波等技术,实现对多个目标的持续跟踪。DeepSORT 的本质任务:在视频帧之间保持目标身份的一致性(ID一致)为同一目标生成持久化ID

DeepSORT算法的主要步骤
目标检测:DeepSORT算法依赖于目标检测器来确定视频中每一帧的目标位置。常用的目标检测器包括YOLO、Faster R-CNN等。检测器的输出通常包括目标的边界框(bounding box)和类别。
特征提取:DeepSORT使用深度学习模型来提取目标的外观特征。这些特征对于目标的再识别(re-identification,简称Re-ID)至关重要,因为即使目标在视频中被临时遮挡或丢失,这些特征也能帮助算法重新识别和关联目标。
匹配和跟踪:DeepSORT算法中的匹配过程涉及到计算检测框和预测框之间的相似度,并使用匈牙利算法来找到最优匹配。这个过程还包括卡尔曼滤波器的使用,它根据目标的历史运动信息来预测其在下一帧中的位置。
级联匹配:DeepSORT中的级联匹配是一种特殊的机制,它首先尝试将检测结果与高置信度的轨迹进行匹配,然后再与低置信度的轨迹进行匹配。这有助于提高匹配的准确性,尤其是在目标被遮挡或短暂消失时。
轨迹管理:DeepSORT维护每个目标的轨迹,并对新检测到的目标初始化新的轨迹。它还设置了确认状态(confirmed)和未确认状态(unconfirmed),以处理遮挡和临时丢失的情况。


                        

代码部分

deepsort的实现库有很多,生产工程方面常用:deep-sort-realtime

from deep_sort_realtime.deepsort_tracker import DeepSort


# 初始化跟踪器
deepsort = DeepSort(
    max_age=30,
    n_init=3,
    nn_budget=100,
    embedder="torchreid",
    half=True,
    max_cosine_distance=0.2,
    max_iou_distance=0.7,

)

# YOLO 推理
    results = model(frame)[0]  # batch 0
    boxes = results.boxes.xywh # xywh (中心点)
    confs = results.boxes.conf
    cls = results.boxes.cls

    # 构建 raw_detections,坐标转换成左上角
    raw_detections = []
    for box, conf, cl in zip(boxes, confs, cls):
        x_c, y_c, w, h = box
        x_left = (x_c - w / 2)
        y_top = (y_c - h / 2)
        raw_detections.append([[x_left, y_top, w, h], conf, cl])
#必须是这种[[x_left, y_top, w, h], conf, cl]

# 更新跟踪
tracks = deepsort.update_tracks(raw_detections, frame=frame)
#返回值跟踪对象列表

执行过程中可能会出现ModuleNotFoundError: No module named 'torchreid.utils'错误,将

yolo+deepsort

from ultralytics import YOLO
import cv2
from deep_sort_realtime.deepsort_tracker import DeepSort
import numpy as np
deepsort=DeepSort(embedder='torchreid')
CLASSES = {
    0: 'person', 1: 'bicycle', 2: 'car', 3: 'motorcycle', 4: 'airplane', 5: 'bus', 6: 'train', 7: 'truck',
    8: 'boat', 9: 'traffic light', 10: 'fire hydrant', 11: 'stop sign', 12: 'parking meter', 13: 'bench',
    14: 'bird', 15: 'cat', 16: 'dog', 17: 'horse', 18: 'sheep', 19: 'cow', 20: 'elephant', 21: 'bear',
    22: 'zebra', 23: 'giraffe', 24: 'backpack', 25: 'umbrella', 26: 'handbag', 27: 'tie', 28: 'suitcase',
    29: 'frisbee', 30: 'skis', 31: 'snowboard', 32: 'sports ball', 33: 'kite', 34: 'baseball bat',
    35: 'baseball glove', 36: 'skateboard', 37: 'surfboard', 38: 'tennis racket', 39: 'bottle',
    40: 'wine glass', 41: 'cup', 42: 'fork', 43: 'knife', 44: 'spoon', 45: 'bowl', 46: 'banana', 47: 'apple',
    48: 'sandwich', 49: 'orange', 50: 'broccoli', 51: 'carrot', 52: 'hot dog', 53: 'pizza', 54: 'donut',
    55: 'cake', 56: 'chair', 57: 'couch', 58: 'potted plant', 59: 'bed', 60: 'dining table', 61: 'toilet',
    62: 'tv', 63: 'laptop', 64: 'mouse', 65: 'remote', 66: 'keyboard', 67: 'cell phone', 68: 'microwave',
    69: 'oven', 70: 'toaster', 71: 'sink', 72: 'refrigerator', 73: 'book', 74: 'clock', 75: 'vase',
    76: 'scissors', 77: 'teddy bear', 78: 'hair drier', 79: 'toothbrush'
}
colors=np.random.uniform(0,255,size=(len(CLASSES),3))
#初始化yolo
model=YOLO('yolo11n.pt')
cap=cv2.VideoCapture('test_person.mp4')#读取视频

while cap.isOpened():
    success,frame=cap.read()#读取一帧
    if success==False:
        print("视频读取完成")
        break
    key=cv2.waitKey(1)&0xFF
    if key==ord('q'):
        break
    result=model(frame)
    result=result[0]
    boxs=result.boxes.xywh
    cons=result.boxes.conf
    clss=result.boxes.cls
    dections=[]
    for box,con,cl in zip(boxs,cons,clss):
        x_c,y_c,w,h=box
        x=x_c-w/2
        y=y_c-h/2
        dections.append([[x,y,w,h],con,int(cl)])
    track_results=deepsort.update_tracks(raw_detections=dections,frame=frame)
    node_features = []
    for track in track_results:
        if not track.is_confirmed():
            continue
        id=track.track_id
        x,y,w,h=track.to_ltwh()
        cl=track.det_class
        x_r=x+w
        y_r=y+h
        cv2.rectangle(frame,(int(x),int(y)),(int(x_r),int(y_r)),colors[int(cl)],2)
        cv2.putText(frame,text=f'{CLASSES[int(cl)]},{id}',org=(int(x-10),int(y-10)),fontFace=cv2.FONT_HERSHEY_SIMPLEX,fontScale=1,color=colors[int(cl)],thickness=1)
        cv2.imshow('1',frame)
        cv2.waitKey(1)


Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐