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

基于YOLOv8实现目标追踪与车辆分类计数的技术求助

YOLOv8目标追踪与车辆准确计数实现方案

核心问题解决思路

同一车辆因距离变化被模型误识别为不同类别,本质是缺乏目标身份绑定。通过YOLOv8内置的追踪功能为每个车辆分配唯一ID,再基于ID维护其最高置信度的类别,即可避免重复计数,实现各车型及总车辆数的准确统计。

修改后完整代码

from ultralytics import YOLO
import streamlit as st
import cv2
from PIL import Image
import tempfile
import numpy as np

# 全局变量:存储追踪对象状态,key为track_id,value为{'class': 类别名, 'max_conf': 最高置信度, 'counted': 是否已计数}
tracked_objects = {}
# 计数统计结果
count_stats = {"total": 0}
# 初始化11种车辆类别(根据你的数据集调整)
vehicle_classes = ["Volkswagen", "BMW", "Audi", "Toyota", "Honda", "Ford", "Chevrolet", "Hyundai", "Kia", "Mercedes", "Tesla"]
for cls in vehicle_classes:
    count_stats[cls] = 0

def _display_detected_frames(conf, model, st_frame, image, count_placeholder):
    """
    带追踪和计数的帧处理函数
    """
    # 调整画面尺寸
    image = cv2.resize(image, (720, int(720 * (9 / 16))))
    frame_height, frame_width = image.shape[:2]
    
    # 使用YOLOv8追踪功能,persist=True保证ID连续
    res = model.track(image, conf=conf, persist=True, classes=list(range(len(vehicle_classes))))
    
    # 处理追踪结果
    if res[0].boxes.id is not None:
        boxes = res[0].boxes.xyxy.cpu().numpy()
        track_ids = res[0].boxes.id.cpu().numpy().astype(int)
        classes = res[0].boxes.cls.cpu().numpy().astype(int)
        confs = res[0].boxes.conf.cpu().numpy()
        
        for box, track_id, cls_idx, conf in zip(boxes, track_ids, classes, confs):
            cls_name = vehicle_classes[cls_idx]
            x1, y1, x2, y2 = box
            center_y = (y1 + y2) / 2
            
            # 更新追踪对象的类别(保留最高置信度的结果)
            if track_id not in tracked_objects:
                tracked_objects[track_id] = {
                    "class": cls_name,
                    "max_conf": conf,
                    "counted": False
                }
            else:
                if conf > tracked_objects[track_id]["max_conf"]:
                    tracked_objects[track_id]["class"] = cls_name
                    tracked_objects[track_id]["max_conf"] = conf
            
            # 计数逻辑:当目标中心点进入画面底部1/3区域且未被计数时
            if center_y > frame_height * 2/3 and not tracked_objects[track_id]["counted"]:
                tracked_objects[track_id]["counted"] = True
                count_stats["total"] += 1
                count_stats[cls_name] += 1
    
    # 绘制计数区域线
    cv2.line(image, (0, int(frame_height*2/3)), (frame_width, int(frame_height*2/3)), (0, 255, 0), 2)
    # 绘制追踪结果
    res_plotted = res[0].plot()
    st_frame.image(res_plotted, caption='Detected Video', channels="BGR", use_column_width=True)
    
    # 更新计数展示
    count_text = "计数结果:\n"
    count_text += f"总车辆数: {count_stats['total']}\n"
    for cls in vehicle_classes:
        count_text += f"{cls}: {count_stats[cls]}\n"
    count_placeholder.text(count_text)

@st.cache_resource
def load_model(model_path):
    model = YOLO("best.pt")
    return model

def infer_uploaded_image(conf, model):
    # 单张图片无需追踪,保留原逻辑,若需要可扩展
    source_img = st.sidebar.file_uploader(label="Choose an image...", type=("jpg", "jpeg", "png", 'bmp', 'webp'))
    col1, col2 = st.columns(2)

    with col1:
        if source_img:
            st.image(image=source_img, caption="Uploaded Image", use_column_width=True)

    if source_img:
        if st.button("Execution"):
            with st.spinner("Running..."):
                res = model.predict(source_img, conf=conf)
                boxes = res[0].boxes
                res_plotted = res[0].plot()[:, :, ::-1]

                with col2:
                    st.image(res_plotted, caption="Detected Image", use_column_width=True)
                    with st.expander("Detection Results"):
                        for box in boxes:
                            cls_idx = int(box.cls)
                            st.write(f"类别: {vehicle_classes[cls_idx]}, 位置: {box.xywh}")

def infer_uploaded_video(conf, model):
    source_video = st.sidebar.file_uploader(label="Choose a video...")
    count_placeholder = st.sidebar.empty()

    if source_video:
        st.video(source_video)

    if source_video:
        if st.button("Execution"):
            # 重置计数状态
            global tracked_objects, count_stats
            tracked_objects = {}
            count_stats = {"total": 0}
            for cls in vehicle_classes:
                count_stats[cls] = 0
            
            with st.spinner("Running..."):
                try:
                    tfile = tempfile.NamedTemporaryFile()
                    tfile.write(source_video.read())
                    vid_cap = cv2.VideoCapture(tfile.name)
                    st_frame = st.empty()
                    while vid_cap.isOpened():
                        success, image = vid_cap.read()
                        if success:
                            _display_detected_frames(conf, model, st_frame, image, count_placeholder)
                        else:
                            vid_cap.release()
                            break
                except Exception as e:
                    st.error(f"Error loading video: {e}")

def infer_uploaded_webcam(conf, model):
    count_placeholder = st.sidebar.empty()
    try:
        flag = st.button(label="Stop running")
        # 重置计数状态
        global tracked_objects, count_stats
        tracked_objects = {}
        count_stats = {"total": 0}
        for cls in vehicle_classes:
            count_stats[cls] = 0
        
        vid_cap = cv2.VideoCapture(0)
        st_frame = st.empty()
        while not flag:
            success, image = vid_cap.read()
            if success:
                _display_detected_frames(conf, model, st_frame, image, count_placeholder)
            else:
                vid_cap.release()
                break
    except Exception as e:
        st.error(f"Error loading video: {str(e)}")

# 主程序入口
if __name__ == "__main__":
    st.title("YOLOv8 Vehicle Detection & Tracking")
    conf = st.sidebar.slider("Confidence Threshold", 0.1, 1.0, 0.3, 0.05)
    model = load_model("best.pt")
    
    st.sidebar.title("Source Selection")
    source_option = st.sidebar.radio("Choose Source", ["Image", "Video", "Webcam"])
    
    if source_option == "Image":
        infer_uploaded_image(conf, model)
    elif source_option == "Video":
        infer_uploaded_video(conf, model)
    elif source_option == "Webcam":
        infer_uploaded_webcam(conf, model)

关键修改说明

  1. 追踪功能替换:用model.track()替代model.predict(),开启persist=True保证追踪ID的连续性
  2. 追踪状态维护:tracked_objects字典记录每个车辆的ID、最高置信度类别、计数状态,解决同一车辆类别误判问题
  3. 计数逻辑:设定画面底部1/3区域为计数触发区,当车辆中心点进入该区域且未被计数时,更新统计结果
  4. 计数展示:在侧边栏实时显示总车辆数及各车型计数
  5. 状态重置:每次开始新的视频/ webcam推理时,重置追踪和计数状态,避免干扰

注意事项

  • 请根据你的实际数据集调整vehicle_classes列表,确保与模型训练的类别顺序一致
  • 计数区域(frame_height * 2/3)可根据你的监控场景调整位置
  • 若需要更精准的追踪,可调整YOLOv8追踪的参数(如iou阈值)

内容的提问来源于stack exchange,提问作者saas

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 18:42:04