feat: live count worker, tracking, and page updates
This commit is contained in:
1 parent
4606c0adeb
commit
1403a07b86
5 files changed
+77
-15
No files matched your search
@@ -22,6 +22,7 @@ class StartRequest(BaseModel):
|
||||
source_rel: Optional[str] = None
|
||||
model_path: Optional[str] = None
|
||||
model_version_id: Optional[int] = None
|
||||
count_classes: list = ["sack"]
|
||||
line_y: int = 266
|
||||
line_x_start: int = 469
|
||||
line_x_end: int = 910
|
||||
@@ -97,6 +98,7 @@ def start(project_id: int, request: StartRequest) -> dict:
|
||||
min_area_scale=request.min_area_scale,
|
||||
spatial_dedup=request.spatial_dedup,
|
||||
whep=whep,
|
||||
count_classes=tuple(request.count_classes),
|
||||
)
|
||||
except live_count.LiveCountError as exc:
|
||||
raise HTTPException(400, str(exc))
|
||||
|
||||
@@ -25,6 +25,7 @@ DEFAULTS = {
|
||||
"entry_travel_min": 60.0, "handoff_radius": 100.0,
|
||||
"unload_confirm_frames": 3, "min_area_scale": 1.0,
|
||||
"dedup_radius": 60.0, "spatial_dedup": False,
|
||||
"count_classes": ["sack"],
|
||||
}
|
||||
|
||||
|
||||
@@ -176,7 +177,8 @@ def count_video(path: str, model, params: dict, should_cancel=None,
|
||||
raise CountingBenchError(f"Could not open {path}")
|
||||
total_frames = int(capture.get(cv2.CAP_PROP_FRAME_COUNT) or 0)
|
||||
|
||||
tracker = ByteTrackTracker(model, settings["conf"])
|
||||
tracker = ByteTrackTracker(model, settings["conf"],
|
||||
class_filter=tuple(settings["count_classes"]))
|
||||
stabilizer = BboxStabilizer(ema_alpha=0.35, max_hold_frames=10,
|
||||
max_height_ratio=1.5, min_height_ratio=0.70)
|
||||
counter = LineCrossCounter(
|
||||
@@ -199,7 +201,8 @@ def count_video(path: str, model, params: dict, should_cancel=None,
|
||||
if not ok or frame is None:
|
||||
break
|
||||
frame = cv2.resize(frame, (1280, 720))
|
||||
detections = [d for d in tracker.update(frame, []) if d.class_name == "sack"]
|
||||
# The tracker only emits count_classes, so no post-filter is needed.
|
||||
detections = tracker.update(frame, [])
|
||||
stable = stabilizer.update(detections)
|
||||
inside = [
|
||||
d for d in stable
|
||||
|
||||
@@ -43,7 +43,8 @@ class Session:
|
||||
dedup_radius: float, margin: int, imgsz: int,
|
||||
entry_travel_min: float, handoff_radius: float,
|
||||
unload_confirm_frames: int, min_area_scale: float,
|
||||
spatial_dedup: bool, whep_url: str = ""):
|
||||
spatial_dedup: bool, whep_url: str = "",
|
||||
count_classes: tuple = ("sack",)):
|
||||
self.source = source
|
||||
# When the browser watches the camera over WebRTC it never asks for the
|
||||
# MJPEG, so encoding a JPEG per frame would be pure waste (REQ-177).
|
||||
@@ -61,6 +62,7 @@ class Session:
|
||||
self.unload_confirm_frames = unload_confirm_frames
|
||||
self.min_area_scale = min_area_scale
|
||||
self.spatial_dedup = spatial_dedup
|
||||
self.count_classes = tuple(count_classes)
|
||||
|
||||
self.started_at = time.time()
|
||||
self.error = ""
|
||||
@@ -167,7 +169,8 @@ class Session:
|
||||
# segfault without it (same reason predict.py does this).
|
||||
model(np.zeros((720, 1280, 3), dtype=np.uint8), imgsz=self.imgsz, verbose=False)
|
||||
|
||||
tracker = ByteTrackTracker(model, self.conf)
|
||||
tracker = ByteTrackTracker(model, self.conf,
|
||||
class_filter=tuple(self.count_classes))
|
||||
stabilizer = BboxStabilizer(ema_alpha=0.35, max_hold_frames=10,
|
||||
max_height_ratio=1.5, min_height_ratio=0.70)
|
||||
counter = LineCrossCounter(
|
||||
@@ -200,7 +203,8 @@ class Session:
|
||||
break
|
||||
|
||||
frame = cv2.resize(frame, (1280, 720))
|
||||
detections = [d for d in tracker.update(frame, []) if d.class_name == "sack"]
|
||||
# The tracker only emits count_classes, so no post-filter is needed.
|
||||
detections = tracker.update(frame, [])
|
||||
stable = stabilizer.update(detections)
|
||||
inside, outside = [], []
|
||||
small = 0
|
||||
@@ -338,7 +342,7 @@ def start(source: str, model_path: str, line_y: int, line_x_start: int, line_x_e
|
||||
imgsz: int = 640, entry_travel_min: float = 60.0,
|
||||
handoff_radius: float = 100.0, unload_confirm_frames: int = 3,
|
||||
min_area_scale: float = 1.0, spatial_dedup: bool = False,
|
||||
whep: str = "") -> dict:
|
||||
whep: str = "", count_classes: tuple = ("sack",)) -> dict:
|
||||
global _session
|
||||
with _guard:
|
||||
if _session is not None and _session.status()["running"]:
|
||||
@@ -350,7 +354,7 @@ def start(source: str, model_path: str, line_y: int, line_x_start: int, line_x_e
|
||||
_session = Session(source, model_path, line_y, line_x_start, line_x_end,
|
||||
conf, dedup_radius, margin, imgsz, entry_travel_min,
|
||||
handoff_radius, unload_confirm_frames, min_area_scale,
|
||||
spatial_dedup, whep)
|
||||
spatial_dedup, whep, count_classes)
|
||||
_session.start()
|
||||
time.sleep(0.4) # let an immediate failure surface in the response
|
||||
return _session.status()
|
||||
|
||||
Reference in new issue
Block a user