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}')
修复说明
- 启用YOLOv5跟踪:通过
model(frame, track=True)启用内置跟踪,每个目标会获得唯一且持续的track_id,避免因框位置波动导致的ID变化。 - 避免重复计数:用
counted_ids集合存储已计数的目标ID,确保每个目标仅被统计一次。 - 优化过线判断:通过
prev_centers存储目标上一帧的中心位置,准确判断是否从左到右穿过中线。
内容的提问来源于stack exchange,提问作者Jiří Komínek
相关产品推荐
相关产品推荐

