如何在基于Faster-RCNN的VCS中结合质心跟踪实现过线车辆计数
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_linefunction to check the sign of the cross product change. - Detection Threshold: If you're missing valid detections, you can lower
min_score_threshslightly (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

