基于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)
关键修改说明
- 追踪功能替换:用
model.track()替代model.predict(),开启persist=True保证追踪ID的连续性 - 追踪状态维护:
tracked_objects字典记录每个车辆的ID、最高置信度类别、计数状态,解决同一车辆类别误判问题 - 计数逻辑:设定画面底部1/3区域为计数触发区,当车辆中心点进入该区域且未被计数时,更新统计结果
- 计数展示:在侧边栏实时显示总车辆数及各车型计数
- 状态重置:每次开始新的视频/ webcam推理时,重置追踪和计数状态,避免干扰
注意事项
- 请根据你的实际数据集调整
vehicle_classes列表,确保与模型训练的类别顺序一致 - 计数区域(
frame_height * 2/3)可根据你的监控场景调整位置 - 若需要更精准的追踪,可调整YOLOv8追踪的参数(如
iou阈值)
内容的提问来源于stack exchange,提问作者saas
相关产品推荐
相关产品推荐

