基于OpenCV的视频目标跟踪与跨线计数优化需求
基于OpenCV与ByteTrack的车辆越线计数优化方案
我需要优化现有基于OpenCV和ByteTrack的视频目标跟踪代码,实现车辆跨越限制线时的目标计数功能,同时要避免重复统计。以下是当前的检测结果和代码,请帮忙改进:
当前检测结果
Detections(xyxy=array([[342.3293 , 335.03476 , 465.3704 , 427.93594 ], [122.26098 , 531.9414 , 263.60587 , 574.8308 ], [209.58714 , 167.03366 , 405.6145 , 365.10098 ], [ 21.653362, 405.5883 , 167.43967 , 511.19382 ]], dtype=float32), confidence=array([0.72265625, 0.7578125 , 0.70703125, 0.60546875], dtype=float32),
当前代码
import os import cv2 from dotenv import load_dotenv from base_camera import BaseCamera from detector import Detector from utils import visualize from detection import Detections from tracker.byte_track.core import ByteTrack load_dotenv(".env") class Camera(BaseCamera): """ OpenCV video stream """ video_source = 0 detector = Detector() tracker = ByteTrack() def __init__(self, enable_detection: bool = False): video_source = os.environ.get("VIDEO_SOURCE") try: video_source = int(video_source) except Exception as exp: # pylint: disable=broad-except if not video_source: raise EnvironmentError("Cannot open the video source!") from exp finally: Camera.set_video_source(video_source) super().__init__() self.enable_detection = enable_detection @staticmethod def set_video_source(source): """Set video source""" Camera.video_source = source @classmethod def frames(cls): """ Get video frame """ camera = cv2.VideoCapture(Camera.video_source) if not camera.isOpened(): raise RuntimeError("Could not start camera.") while True: # read current frame ret, img = camera.read() # Loop back if not ret: camera.set(cv2.CAP_PROP_POS_FRAMES, 0) continue # Object detection results = cls.detector(image=img) tensorflow_results = results.detections if not tensorflow_results: yield cv2.imencode(".jpg", img)[1].tobytes() continue detections = Detections.from_tensorflow(tensorflow_results=tensorflow_results) detections = cls.tracker.update_with_detections(detections=detections) # print(detections) restrict_line=((100,300),(500,550)) img = visualize(image=img, detections=detections) count=0 cv2.line(img, restrict_line[0], restrict_line[1], (255, 0, 0), 5) # Check if any detected object crosses the restrict line # encode as a jpeg image and return it yield cv2.imencode(".jpg", img)[1].tobytes()
优化思路
- 跟踪ID持久化:利用ByteTrack输出的跟踪ID,记录已计数的车辆ID,避免重复统计
- 越线判断逻辑:计算车辆中心坐标,通过帧间中心位置与限制线的相对位置变化,判断是否跨越限制线
- 计数状态维护:为每个跟踪ID维护越线状态,确保仅在首次跨越时计数
优化后的代码
import os import cv2 from dotenv import load_dotenv from base_camera import BaseCamera from detector import Detector from utils import visualize from detection import Detections from tracker.byte_track.core import ByteTrack load_dotenv(".env") class Camera(BaseCamera): """ OpenCV video stream with vehicle crossing count """ video_source = 0 detector = Detector() tracker = ByteTrack() # 存储已计数的车辆ID,避免重复统计 counted_ids = set() # 存储每个ID的上一帧中心位置,用于判断越线 prev_centers = {} def __init__(self, enable_detection: bool = False): video_source = os.environ.get("VIDEO_SOURCE") try: video_source = int(video_source) except Exception as exp: # pylint: disable=broad-except if not video_source: raise EnvironmentError("Cannot open the video source!") from exp finally: Camera.set_video_source(video_source) super().__init__() self.enable_detection = enable_detection @staticmethod def set_video_source(source): """Set video source""" Camera.video_source = source @staticmethod def point_line_side(point, line_start, line_end): """判断点在直线的哪一侧,返回值>0表示一侧,<0表示另一侧,=0表示在直线上""" return (line_end[0] - line_start[0]) * (point[1] - line_start[1]) - (line_end[1] - line_start[1]) * (point[0] - line_start[0]) @classmethod def frames(cls): """ Get video frame with vehicle crossing count """ camera = cv2.VideoCapture(Camera.video_source) if not camera.isOpened(): raise RuntimeError("Could not start camera.") # 定义限制线,可根据实际场景调整 restrict_line = ((100, 300), (500, 550)) total_count = 0 while True: # read current frame ret, img = camera.read() # Loop back if not ret: camera.set(cv2.CAP_PROP_POS_FRAMES, 0) # 循环播放时重置计数状态 cls.counted_ids.clear() cls.prev_centers.clear() total_count = 0 continue # Object detection results = cls.detector(image=img) tensorflow_results = results.detections if not tensorflow_results: # 绘制计数文本 cv2.putText(img, f"Crossed Vehicles: {total_count}", (20, 50), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2) cv2.line(img, restrict_line[0], restrict_line[1], (255, 0, 0), 5) yield cv2.imencode(".jpg", img)[1].tobytes() continue detections = Detections.from_tensorflow(tensorflow_results=tensorflow_results) detections = cls.tracker.update_with_detections(detections=detections) img = visualize(image=img, detections=detections) cv2.line(img, restrict_line[0], restrict_line[1], (255, 0, 0), 5) # 遍历每个检测到的目标 for idx in range(len(detections.xyxy)): track_id = detections.tracker_id[idx] # 计算目标中心坐标 x1, y1, x2, y2 = detections.xyxy[idx] center_x = (x1 + x2) / 2 center_y = (y1 + y2) / 2 current_center = (center_x, center_y) # 如果是首次跟踪该ID,记录初始位置 if track_id not in cls.prev_centers: cls.prev_centers[track_id] = current_center continue # 获取上一帧的中心位置 prev_center = cls.prev_centers[track_id] # 判断前后两帧中心在限制线的两侧,且未被计数过 prev_side = cls.point_line_side(prev_center, restrict_line[0], restrict_line[1]) current_side = cls.point_line_side(current_center, restrict_line[0], restrict_line[1]) # 当两侧符号不同,说明跨越了限制线 if (prev_side * current_side < 0) and (track_id not in cls.counted_ids): total_count += 1 cls.counted_ids.add(track_id) # 更新当前ID的中心位置 cls.prev_centers[track_id] = current_center # 绘制计数文本到画面 cv2.putText(img, f"Crossed Vehicles: {total_count}", (20, 50), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2) # encode as a jpeg image and return it yield cv2.imencode(".jpg", img)[1].tobytes()
关键优化点说明
counted_ids集合:存储已经计数过的车辆跟踪ID,确保每个车辆只被统计一次prev_centers字典:记录每个跟踪ID上一帧的中心坐标,用于对比位置变化point_line_side静态方法:通过叉积计算点与直线的相对位置,判断是否跨越限制线- 循环播放重置逻辑:当视频循环时,清空计数状态,避免累计错误
内容的提问来源于stack exchange,提问作者SARON RAVUTH
相关产品推荐
相关产品推荐

