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

如何在基于Faster-RCNN的VCS中结合质心跟踪实现过线车辆计数

How to Implement Centroid Tracking for Vehicle Counting with Your TensorFlow Object Detection Model

Hey, great job getting the detection part working with Faster-RCNN! You're right that centroid tracking is a perfect fit for your vehicle counting task—it lets you track individual vehicles across frames and reliably count them when they cross your defined line. Let's walk through the implementation, starting with the centroid tracker itself, then integrating it into your existing code.

Step 1: Implement the Centroid Tracker

First, we'll create a simple centroid tracker class that handles:

  • Assigning unique IDs to new detected objects
  • Matching centroids between consecutive frames using Euclidean distance
  • Removing objects that haven't been detected for a few frames (to handle temporary occlusions)

Here's the tracker code:

from scipy.spatial import distance as dist
from collections import OrderedDict

class CentroidTracker:
    def __init__(self, maxDisappeared=50):
        # Initialize next unique object ID and tracking dictionaries
        self.nextObjectID = 0
        self.objects = OrderedDict()
        self.disappeared = OrderedDict()
        
        # Max frames an object can be missing before deregistering
        self.maxDisappeared = maxDisappeared

    def register(self, centroid):
        # Assign ID to a new detected object
        self.objects[self.nextObjectID] = centroid
        self.disappeared[self.nextObjectID] = 0
        self.nextObjectID += 1

    def deregister(self, objectID):
        # Remove an object from tracking if it's gone too long
        del self.objects[objectID]
        del self.disappeared[objectID]

    def update(self, rects):
        # Handle case where no objects are detected
        if len(rects) == 0:
            for objectID in list(self.disappeared.keys()):
                self.disappeared[objectID] += 1
                if self.disappeared[objectID] > self.maxDisappeared:
                    self.deregister(objectID)
            return self.objects

        # Calculate centroids for current frame's detections
        inputCentroids = np.zeros((len(rects), 2), dtype="int")
        for (i, (startX, startY, endX, endY)) in enumerate(rects):
            cX = int((startX + endX) / 2.0)
            cY = int((startY + endY) / 2.0)
            inputCentroids[i] = (cX, cY)

        # Register all centroids if we're not tracking any objects yet
        if len(self.objects) == 0:
            for i in range(0, len(inputCentroids)):
                self.register(inputCentroids[i])
        else:
            # Match current centroids to existing tracked objects
            objectIDs = list(self.objects.keys())
            objectCentroids = list(self.objects.values())

            # Compute distance between all pairs of existing and new centroids
            D = dist.cdist(np.array(objectCentroids), inputCentroids)
            rows = D.min(axis=1).argsort()
            cols = D.argmin(axis=1)[rows]

            # Track used rows/columns to avoid duplicate matches
            usedRows = set()
            usedCols = set()

            for (row, col) in zip(rows, cols):
                if row in usedRows or col in usedCols:
                    continue

                # Update the centroid for the matched object
                objectID = objectIDs[row]
                self.objects[objectID] = inputCentroids[col]
                self.disappeared[objectID] = 0

                usedRows.add(row)
                usedCols.add(col)

            # Handle objects that weren't matched (possible occlusions)
            unusedRows = set(range(0, D.shape[0])).difference(usedRows)
            for row in unusedRows:
                objectID = objectIDs[row]
                self.disappeared[objectID] += 1
                if self.disappeared[objectID] > self.maxDisappeared:
                    self.deregister(objectID)

            # Register new objects that weren't matched to existing ones
            unusedCols = set(range(0, D.shape[1])).difference(usedCols)
            for col in unusedCols:
                self.register(inputCentroids[col])

        return self.objects

Step 2: Integrate Tracker and Counting Logic into Your Code

Now we'll modify your existing code to use the tracker, calculate centroids from detection boxes, and count vehicles when they cross your line. We'll also add a set to track which vehicle IDs have already been counted to avoid duplicates.

Here's the updated full code:

import os
import cv2
import numpy as np
import tensorflow as tf
import sys
from scipy.spatial import distance as dist
from collections import OrderedDict

# This is needed since the notebook is stored in the object_detection folder.
sys.path.append("..")

# Import utilities
from utils import label_map_util
from utils import visualization_utils as vis_util

# Centroid Tracker Class
class CentroidTracker:
    def __init__(self, maxDisappeared=50):
        self.nextObjectID = 0
        self.objects = OrderedDict()
        self.disappeared = OrderedDict()
        self.maxDisappeared = maxDisappeared

    def register(self, centroid):
        self.objects[self.nextObjectID] = centroid
        self.disappeared[self.nextObjectID] = 0
        self.nextObjectID += 1

    def deregister(self, objectID):
        del self.objects[objectID]
        del self.disappeared[objectID]

    def update(self, rects):
        if len(rects) == 0:
            for objectID in list(self.disappeared.keys()):
                self.disappeared[objectID] += 1
                if self.disappeared[objectID] > self.maxDisappeared:
                    self.deregister(objectID)
            return self.objects

        inputCentroids = np.zeros((len(rects), 2), dtype="int")
        for (i, (startX, startY, endX, endY)) in enumerate(rects):
            cX = int((startX + endX) / 2.0)
            cY = int((startY + endY) / 2.0)
            inputCentroids[i] = (cX, cY)

        if len(self.objects) == 0:
            for i in range(0, len(inputCentroids)):
                self.register(inputCentroids[i])
        else:
            objectIDs = list(self.objects.keys())
            objectCentroids = list(self.objects.values())

            D = dist.cdist(np.array(objectCentroids), inputCentroids)
            rows = D.min(axis=1).argsort()
            cols = D.argmin(axis=1)[rows]

            usedRows = set()
            usedCols = set()

            for (row, col) in zip(rows, cols):
                if row in usedRows or col in usedCols:
                    continue

                objectID = objectIDs[row]
                self.objects[objectID] = inputCentroids[col]
                self.disappeared[objectID] = 0

                usedRows.add(row)
                usedCols.add(col)

            unusedRows = set(range(0, D.shape[0])).difference(usedRows)
            for row in unusedRows:
                objectID = objectIDs[row]
                self.disappeared[objectID] += 1
                if self.disappeared[objectID] > self.maxDisappeared:
                    self.deregister(objectID)

            unusedCols = set(range(0, D.shape[1])).difference(usedCols)
            for col in unusedCols:
                self.register(inputCentroids[col])

        return self.objects

# Project paths and setup
MODEL_NAME = 'inference_graph'
VIDEO_NAME = 'Video_105.mp4'
CWD_PATH = os.getcwd()
PATH_TO_CKPT = os.path.join(CWD_PATH, MODEL_NAME, 'frozen_inference_graph.pb')
PATH_TO_LABELS = os.path.join(CWD_PATH, 'training', 'labelmap.pbtxt')
PATH_TO_VIDEO = os.path.join(CWD_PATH, VIDEO_NAME)
NUM_CLASSES = 7

# Load label map
label_map = label_map_util.load_labelmap(PATH_TO_LABELS)
categories = label_map_util.convert_label_map_to_categories(label_map, max_num_classes=NUM_CLASSES, use_display_name=True)
category_index = label_map_util.create_category_index(categories)

# Load TensorFlow model
detection_graph = tf.Graph()
with detection_graph.as_default():
    od_graph_def = tf.GraphDef()
    with tf.gfile.GFile(PATH_TO_CKPT, 'rb') as fid:
        serialized_graph = fid.read()
        od_graph_def.ParseFromString(serialized_graph)
        tf.import_graph_def(od_graph_def, name='')

    sess = tf.Session(graph=detection_graph)

# Define detection tensors
image_tensor = detection_graph.get_tensor_by_name('image_tensor:0')
detection_boxes = detection_graph.get_tensor_by_name('detection_boxes:0')
detection_scores = detection_graph.get_tensor_by_name('detection_scores:0')
detection_classes = detection_graph.get_tensor_by_name('detection_classes:0')
num_detections = detection_graph.get_tensor_by_name('num_detections:0')

# Initialize tracking and counting variables
ct = CentroidTracker(maxDisappeared=20)  # Adjust based on your video frame rate
counted_ids = set()
total_count = 0
prev_centroids = {}

# Define your counting line (use your existing coordinates)
line_start = (1144, 568)
line_end = (1723, 664)

# Helper function to check if a point crossed the line
def point_crossed_line(prev_centroid, curr_centroid, line_start, line_end):
    def cross(o, a, b):
        return (a[0]-o[0])*(b[1]-o[1]) - (a[1]-o[1])*(b[0]-o[0])
    
    prev_side = cross(line_start, line_end, prev_centroid)
    curr_side = cross(line_start, line_end, curr_centroid)
    
    # Return True if the point moved from one side of the line to the other
    return (prev_side * curr_side) < 0

# Process video
video = cv2.VideoCapture(PATH_TO_VIDEO)

while(video.isOpened()):
    ret, frame = video.read()
    if not ret:
        break
        
    frame_expanded = np.expand_dims(frame, axis=0)

    # Run detection
    (boxes, scores, classes, num) = sess.run(
        [detection_boxes, detection_scores, detection_classes, num_detections],
        feed_dict={image_tensor: frame_expanded})

    # Filter high-confidence detections
    min_score_thresh = 0.90
    valid_indices = np.where(np.squeeze(scores) > min_score_thresh)[0]
    valid_boxes = np.squeeze(boxes)[valid_indices]
    im_height, im_width = frame.shape[:2]

    # Convert normalized boxes to pixel coordinates
    rects = []
    for box in valid_boxes:
        ymin, xmin, ymax, xmax = box
        startX = int(xmin * im_width)
        startY = int(ymin * im_height)
        endX = int(xmax * im_width)
        endY = int(ymax * im_height)
        rects.append((startX, startY, endX, endY))

    # Update tracker with current detections
    objects = ct.update(rects)

    # Draw detection results
    vis_util.visualize_boxes_and_labels_on_image_array(
        frame,
        np.squeeze(boxes),
        np.squeeze(classes).astype(np.int32),
        np.squeeze(scores),
        category_index,
        use_normalized_coordinates=True,
        line_thickness=8,
        min_score_thresh=0.90)

    # Draw counting line
    cv2.line(frame, line_start, line_end, (0,0,255), 2)

    # Check for line crossings and update count
    for (objectID, centroid) in objects.items():
        # Draw centroid and object ID
        cv2.circle(frame, (centroid[0], centroid[1]), 4, (0, 255, 0), -1)
        cv2.putText(frame, f"ID {objectID}", (centroid[0]-10, centroid[1]-10),
                    cv2.FONT_HERSHEY_SIMPLEX, 0.5, (0, 255, 0), 2)

        # Check if the object crossed the line since last frame
        if objectID in prev_centroids:
            prev_cent = prev_centroids[objectID]
            if point_crossed_line(prev_cent, centroid, line_start, line_end) and objectID not in counted_ids:
                total_count += 1
                counted_ids.add(objectID)
        
        # Save current centroid for next frame
        prev_centroids[objectID] = centroid

    # Display total count
    cv2.putText(frame, f"Total Count: {total_count}", (50, 50),
                cv2.FONT_HERSHEY_SIMPLEX, 1.5, (0, 255, 0), 3)

    # Show output frame
    cv2.imshow('Object detector', frame)

    if cv2.waitKey(1) == ord('q'):
        break

# Cleanup
video.release()
cv2.destroyAllWindows()

Key Adjustments for Your Project

  • maxDisappeared: Tweak this based on your video's frame rate. For 30fps video, setting it to 20 means we wait ~0.6 seconds before dropping a temporarily occluded vehicle.
  • Line Crossing Direction: The current logic counts any crossing, but if you only want to count vehicles moving in one direction, you can modify the point_crossed_line function to check the sign of the cross product change.
  • Detection Threshold: If you're missing valid detections, you can lower min_score_thresh slightly (e.g., to 0.85) and let the tracker handle temporary misses.

This implementation will reliably track each vehicle and count it exactly once when it crosses your defined line. Let me know if you need help tweaking any part for your specific video!

内容的提问来源于stack exchange,提问作者Guna

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.14 08:44:49