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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 06:54:57