Vehicle-Counting / core /counter.py
riezqidr's picture
Add initial project structure with Streamlit UI and utility functions
ab2f940
Raw
History Blame Contribute Delete
6.48 kB
"""
Vehicle Counting Module
Dua metode counting:
1. VirtualLineCounter - hitung kendaraan yang melewati garis virtual
2. PolygonRegionCounter - hitung kendaraan yang masuk ke area polygon
Kedua counter melacak ID unik supaya satu kendaraan tidak dihitung dua kali.
"""
import numpy as np
from collections import defaultdict
class VirtualLineCounter:
"""
Menghitung kendaraan yang melewati garis virtual horizontal.
Cara kerja:
- Definisikan garis horizontal pada posisi y tertentu
- Untuk setiap tracked object, cek apakah centroid-nya
berpindah dari atas ke bawah (atau sebaliknya) melewati garis
- Setiap kendaraan hanya dihitung sekali berdasarkan track_id
"""
def __init__(self, line_position_ratio=0.5, frame_height=720):
"""
Args:
line_position_ratio: posisi garis sebagai rasio dari tinggi frame (0.0 - 1.0)
frame_height: tinggi frame video
"""
self.line_position_ratio = line_position_ratio
self.frame_height = frame_height
self.line_y = int(frame_height * line_position_ratio)
# set untuk menyimpan ID yang sudah dihitung
self.counted_ids = set()
# counter per kelas
self.class_counts = defaultdict(int)
self.total_count = 0
def update_line(self, line_position_ratio, frame_height):
"""Update posisi garis jika berubah."""
self.line_position_ratio = line_position_ratio
self.frame_height = frame_height
self.line_y = int(frame_height * line_position_ratio)
def update(self, tracked_objects):
"""
Cek setiap tracked object apakah melewati garis.
Args:
tracked_objects: list of dict dari ByteTracker.update()
harus punya: track_id, class_name, center_history
"""
for obj in tracked_objects:
track_id = obj["track_id"]
# skip kalau sudah pernah dihitung
if track_id in self.counted_ids:
continue
history = obj.get("center_history", [])
if len(history) < 2:
continue
# ambil posisi y sekarang dan sebelumnya
prev_y = history[-2][1]
curr_y = history[-1][1]
# cek crossing: dari atas ke bawah ATAU bawah ke atas
crossed = False
if prev_y < self.line_y and curr_y >= self.line_y:
crossed = True # atas ke bawah
elif prev_y > self.line_y and curr_y <= self.line_y:
crossed = True # bawah ke atas
if crossed:
self.counted_ids.add(track_id)
self.total_count += 1
self.class_counts[obj["class_name"]] += 1
def get_counts(self):
"""
Return hasil counting.
Returns:
dict dengan keys:
- total: int
- per_class: dict {class_name: count}
"""
return {
"total": self.total_count,
"per_class": dict(self.class_counts)
}
def get_line_coordinates(self, frame_width):
"""
Return koordinat garis untuk drawing.
Returns:
tuple: ((x1, y1), (x2, y2))
"""
return ((0, self.line_y), (frame_width, self.line_y))
def reset(self):
"""Reset semua counter."""
self.counted_ids = set()
self.class_counts = defaultdict(int)
self.total_count = 0
class PolygonRegionCounter:
"""
Menghitung kendaraan yang masuk ke area polygon.
Cara kerja:
- Definisikan polygon region (list of points)
- Cek apakah centroid kendaraan berada di dalam polygon
- Setiap kendaraan hanya dihitung sekali berdasarkan track_id
"""
def __init__(self, polygon_points=None, frame_width=1280, frame_height=720):
"""
Args:
polygon_points: list of (x, y) tuples, definisikan vertices polygon
Jika None, akan dibuat default rectangle di tengah frame
frame_width: lebar frame
frame_height: tinggi frame
"""
self.frame_width = frame_width
self.frame_height = frame_height
if polygon_points is None:
# default: rectangle di area tengah-bawah frame
margin_x = int(frame_width * 0.15)
margin_top = int(frame_height * 0.4)
margin_bottom = int(frame_height * 0.1)
self.polygon = [
(margin_x, margin_top),
(frame_width - margin_x, margin_top),
(frame_width - margin_x, frame_height - margin_bottom),
(margin_x, frame_height - margin_bottom)
]
else:
self.polygon = polygon_points
self.counted_ids = set()
self.class_counts = defaultdict(int)
self.total_count = 0
def _point_in_polygon(self, x, y):
"""
Cek apakah point (x, y) berada di dalam polygon.
Menggunakan ray casting algorithm.
"""
n = len(self.polygon)
inside = False
j = n - 1
for i in range(n):
xi, yi = self.polygon[i]
xj, yj = self.polygon[j]
if ((yi > y) != (yj > y)) and (x < (xj - xi) * (y - yi) / (yj - yi) + xi):
inside = not inside
j = i
return inside
def update(self, tracked_objects):
"""
Cek setiap tracked object apakah centroid-nya di dalam polygon.
Args:
tracked_objects: list of dict dari ByteTracker.update()
"""
for obj in tracked_objects:
track_id = obj["track_id"]
if track_id in self.counted_ids:
continue
cx, cy = obj["center"]
if self._point_in_polygon(cx, cy):
self.counted_ids.add(track_id)
self.total_count += 1
self.class_counts[obj["class_name"]] += 1
def get_counts(self):
"""Return hasil counting."""
return {
"total": self.total_count,
"per_class": dict(self.class_counts)
}
def get_polygon_points(self):
"""Return polygon points untuk drawing."""
return self.polygon
def reset(self):
"""Reset semua counter."""
self.counted_ids = set()
self.class_counts = defaultdict(int)
self.total_count = 0