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

YOLOv5视频检测中跨帧目标计数重复问题求助

问题:YOLOv5目标计数重复统计虚高

使用YOLOv5对视频中穿过画面中线的自行车(class_id=1)和摩托车(class_id=3)进行计数时,出现同一目标每帧被重复计数的问题,统计结果虚高。

原代码尝试通过object_id = f"{int(class_id)}-{int(x1)}-{int(y1)}"生成唯一ID跟踪目标,但未解决重复计数问题:

import torch
import cv2
import numpy as np

# Load the YOLOv5 model (we are using the small version, yolov5s)
model = torch.hub.load('ultralytics/yolov5', 'yolov5s')

# Load the video
video_path = 'C:/Users/komyj/Downloads/yolov5-master/yolov5-master/data/images/test_2.mp4'
output_path = 'C:/Users/komyj/Downloads/yolov5-master/yolov5-master/data/images/output.mp4'
cap = cv2.VideoCapture(video_path)

if not cap.isOpened():
    print("Error opening the video.")
    exit()

# Get video dimensions
frame_width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
frame_height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
fps = cap.get(cv2.CAP_PROP_FPS)
line_position = frame_width // 2  # Vertical line in the middle of the video

# Initialize counters
count_bikes = 0
count_motorcycles = 0

# Dictionary to store tracked objects with their ID and position history
tracked_objects = {}

# Function to check if an object has crossed the line
def is_object_passing(line_position, current_x, previous_x):
    return previous_x < line_position and current_x >= line_position

# Setup VideoWriter to save the output video
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
out = cv2.VideoWriter(output_path, fourcc, fps, (frame_width, frame_height))

if not out.isOpened():
    print("Error opening VideoWriter.")
    cap.release()
    exit()

while cap.isOpened():
    ret, frame = cap.read()
    if not ret:
        break

    # Convert the image to RGB (OpenCV uses BGR)
    img_rgb = cv2.cvtColor(frame, cv2.COLOR_BGR2RGB)
    
    # Object detection
    results = model(img_rgb)
    
    # Render the detection results
    img_rendered = results.render()[0]
    
    # Convert back to BGR for OpenCV
    img_bgr = cv2.cvtColor(img_rendered, cv2.COLOR_RGB2BGR)
    
    # Display the vertical line
    cv2.line(img_bgr, (line_position, 0), (line_position, frame_height), (255, 0, 0), 2)

    # Process detected objects
    for detection in results.pred[0]:
        x1, y1, x2, y2, conf, class_id = detection
        center_x = int((x1 + x2) / 2)
        center_y = int((y1 + y2) / 2)

        # Generate a unique ID for each object based on its class and position
        object_id = f"{int(class_id)}-{int(x1)}-{int(y1)}"

        # If the object is already being tracked, update its position
        if object_id in tracked_objects:
            prev_center_x = tracked_objects[object_id]['center_x']
            tracked_objects[object_id]['center_x'] = center_x
        else:
            # If the object is not being tracked, add it with its ID and initial position
            tracked_objects[object_id] = {'class_id': class_id, 'center_x': center_x, 'counted': False}

        # Check if the object has already been counted and if it has crossed the line
        if not tracked_objects[object_id]['counted'] and is_object_passing(line_position, center_x, prev_center_x):
            if class_id == 1:  # ID for a bicycle
                count_bikes += 1
            elif class_id == 3:  # ID for a motorcycle
                count_motorcycles += 1
            tracked_objects[object_id]['counted'] = True  # Mark the object as counted

        # Draw a rectangle around the detected object
        if class_id == 1:  # ID for a bicycle
            cv2.rectangle(img_bgr, (int(x1), int(y1)), (int(x2), int(y2)), (0, 255, 0), 2)
            cv2.putText(img_bgr, "Bike", (int(x1), int(y1) - 10), cv2.FONT_HERSHEY_SIMPLEX, 0.9, (0, 255, 0), 2)
        
        elif class_id == 3:  # ID for a motorcycle
            cv2.rectangle(img_bgr, (int(x1), int(y1)), (int(x2), int(y2)), (0, 0, 255), 2)
            cv2.putText(img_bgr, "Motorcycle", (int(x1), int(y1) - 10), cv2.FONT_HERSHEY_SIMPLEX, 0.9, (0, 0, 255), 2)

    # Display the bike and motorcycle counts
    cv2.putText(img_bgr, f'Bikes Count: {count_bikes}', (50, 50), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2)
    cv2.putText(img_bgr, f'Motorcycles Count: {count_motorcycles}', (50, 100), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 0, 255), 2)
    
    # Save the frame to the new video
    out.write(img_bgr)

cap.release()
out.release()
cv2.destroyAllWindows()

# Print the final results
print(f'Final Bike Count: {count_bikes}')
print(f'Final Motorcycle Count: {count_motorcycles}')

问题根源

  • 目标ID生成逻辑失效:原代码用class_id+x1+y1生成object_id,但每帧检测的目标框位置会因检测波动产生微小变化,导致同一目标每帧生成不同ID,无法被持续跟踪。
  • 未初始化变量报错:当目标是新对象时,prev_center_x未定义,执行判断时会抛出NameError,代码实际无法正常运行。

修复方案

使用YOLOv5自带的目标跟踪功能(通过track=True启用),该功能会为同一目标分配固定的跟踪ID,避免重复计数。同时优化过线判断逻辑,确保每个目标仅被计数一次。

修复后代码

import torch
import cv2
import numpy as np

# 加载YOLOv5模型并启用跟踪功能
model = torch.hub.load('ultralytics/yolov5', 'yolov5s', pretrained=True)

# 视频路径设置
video_path = 'C:/Users/komyj/Downloads/yolov5-master/yolov5-master/data/images/test_2.mp4'
output_path = 'C:/Users/komyj/Downloads/yolov5-master/yolov5-master/data/images/output.mp4'
cap = cv2.VideoCapture(video_path)

if not cap.isOpened():
    print("Error opening the video.")
    exit()

# 获取视频参数
frame_width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
frame_height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
fps = cap.get(cv2.CAP_PROP_FPS)
line_position = frame_width // 2  # 画面中线位置

# 计数器初始化
count_bikes = 0
count_motorcycles = 0

# 存储已计数的目标ID,避免重复统计
counted_ids = set()

# 初始化视频写入器
fourcc = cv2.VideoWriter_fourcc(*'mp4v')
out = cv2.VideoWriter(output_path, fourcc, fps, (frame_width, frame_height))

if not out.isOpened():
    print("Error opening VideoWriter.")
    cap.release()
    exit()

# 过线判断函数:判断目标是否从左到右穿过中线
def is_crossed_line(line_pos, current_center_x, prev_center_x):
    return prev_center_x < line_pos and current_center_x >= line_pos

# 存储目标上一帧的中心位置
prev_centers = {}

while cap.isOpened():
    ret, frame = cap.read()
    if not ret:
        break

    # 启用YOLOv5跟踪模式,获取带跟踪ID的结果
    results = model(frame, track=True)
    
    # 渲染检测框
    img_bgr = results.render()[0]
    
    # 绘制中线
    cv2.line(img_bgr, (line_position, 0), (line_position, frame_height), (255, 0, 0), 2)

    # 处理检测结果(results.pred[0]包含:x1,y1,x2,y2,conf,class_id,track_id)
    for detection in results.pred[0]:
        x1, y1, x2, y2, conf, class_id, track_id = detection
        track_id = int(track_id)
        class_id = int(class_id)
        center_x = int((x1 + x2) / 2)
        center_y = int((y1 + y2) / 2)

        # 仅处理自行车和摩托车
        if class_id not in [1, 3]:
            continue

        # 记录当前目标中心位置,用于下一帧对比
        current_center = center_x
        prev_center = prev_centers.get(track_id, None)

        # 判断是否过线且未被计数
        if prev_center is not None and track_id not in counted_ids:
            if is_crossed_line(line_position, current_center, prev_center):
                if class_id == 1:
                    count_bikes += 1
                elif class_id == 3:
                    count_motorcycles += 1
                counted_ids.add(track_id)

        # 更新目标上一帧的中心位置
        prev_centers[track_id] = current_center

        # 绘制目标框和标签
        if class_id == 1:
            color = (0, 255, 0)
            label = f"Bike ID:{track_id}"
        else:
            color = (0, 0, 255)
            label = f"Motorcycle ID:{track_id}"
        
        cv2.rectangle(img_bgr, (int(x1), int(y1)), (int(x2), int(y2)), color, 2)
        cv2.putText(img_bgr, label, (int(x1), int(y1)-10), cv2.FONT_HERSHEY_SIMPLEX, 0.9, color, 2)

    # 显示计数结果
    cv2.putText(img_bgr, f'Bikes Count: {count_bikes}', (50, 50), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 255, 0), 2)
    cv2.putText(img_bgr, f'Motorcycles Count: {count_motorcycles}', (50, 100), cv2.FONT_HERSHEY_SIMPLEX, 1, (0, 0, 255), 2)
    
    # 写入视频帧
    out.write(img_bgr)

# 释放资源
cap.release()
out.release()
cv2.destroyAllWindows()

# 打印最终结果
print(f'Final Bike Count: {count_bikes}')
print(f'Final Motorcycle Count: {count_motorcycles}')

修复说明

  1. 启用YOLOv5跟踪:通过model(frame, track=True)启用内置跟踪,每个目标会获得唯一且持续的track_id,避免因框位置波动导致的ID变化。
  2. 避免重复计数:用counted_ids集合存储已计数的目标ID,确保每个目标仅被统计一次。
  3. 优化过线判断:通过prev_centers存储目标上一帧的中心位置,准确判断是否从左到右穿过中线。

内容的提问来源于stack exchange,提问作者Jiří Komínek

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.19 10:17:04