323 lines
13 KiB
Python
323 lines
13 KiB
Python
"""Batch counting bench: run the counter over archive videos and score it (REQ-150…153).
|
|
|
|
This is the offline twin of `live_count`. Same model, same tracker, stabiliser
|
|
and `LineCrossCounter`, same defaults — but no MJPEG stream, no annotated frame,
|
|
no JPEG encode. Rendering is most of the per-frame cost once the model is warm,
|
|
so dropping it is what makes counting a 30-minute video practical.
|
|
|
|
The point is measurement, not watching: a row per video, the counter's numbers
|
|
beside a ground truth you type in, and the signed error between them. A model
|
|
that counts 98 where you counted 100 is a different problem from one that counts
|
|
103, and a single accuracy percentage hides which of the two you have.
|
|
"""
|
|
|
|
import os
|
|
import time
|
|
from typing import List, Optional
|
|
|
|
from backend import archive_index, db, jobs, library, projects
|
|
|
|
# Defaults are the values dialled in against the real camera. A run records the
|
|
# parameters it used, so a row always says what produced it.
|
|
DEFAULTS = {
|
|
"line_y": 266, "line_x_start": 469, "line_x_end": 910,
|
|
"margin": 5, "conf": 0.35, "imgsz": 640,
|
|
"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"],
|
|
}
|
|
|
|
|
|
class CountingBenchError(Exception):
|
|
pass
|
|
|
|
|
|
def _split_rel(rel: str) -> tuple:
|
|
date_label = rel.split("/")[0] if "/" in rel else ""
|
|
return date_label, library.batch_label(os.path.basename(rel))
|
|
|
|
|
|
# ---- rows ----------------------------------------------------------------
|
|
|
|
def listing(project_id: int, date: Optional[str] = None) -> dict:
|
|
"""Every archive video with whatever has been measured for it.
|
|
|
|
Videos with no run yet are still rows — the table is the work list, so a
|
|
video nobody has counted has to be visible in it.
|
|
"""
|
|
project = projects.get(project_id)
|
|
if project is None:
|
|
raise CountingBenchError("No such project")
|
|
|
|
with db.cursor() as cur:
|
|
cur.execute("SELECT * FROM count_runs WHERE project_id = ?", (project_id,))
|
|
stored = {row["video_rel"]: dict(row) for row in cur.fetchall()}
|
|
|
|
dates = [d["date"] for d in library.list_dates(project["video_root"])]
|
|
|
|
# Folder names are not when a recording happened, so the grouping and the
|
|
# ordering both come from the timestamp index instead (REQ-160…163).
|
|
clock = archive_index.index(project_id)
|
|
|
|
rows = []
|
|
for day in dates:
|
|
# No `project_id`: that argument makes the library kick off an H.264
|
|
# preview transcode per video, and this table never plays anything.
|
|
for video in library.list_videos(project["video_root"], day):
|
|
rel = video["rel"]
|
|
run = stored.get(rel)
|
|
timing = clock.get(rel) or {}
|
|
rows.append({
|
|
"video_rel": rel,
|
|
"folder_date": day,
|
|
"date_label": timing.get("working_day") or day,
|
|
"working_day": timing.get("working_day") or "",
|
|
"started_at": timing.get("started_at"),
|
|
"clock_source": timing.get("source") or "",
|
|
"clock_trusted": bool(timing.get("trusted")),
|
|
"clock_error": timing.get("error") or "",
|
|
"batch_label": library.batch_label(os.path.basename(rel)),
|
|
"duration": video.get("duration"),
|
|
"loading": run["loading"] if run else None,
|
|
"unloading": run["unloading"] if run else None,
|
|
"net": run["net"] if run else None,
|
|
"ground_truth": run["ground_truth"] if run else None,
|
|
"frames": run["frames"] if run else 0,
|
|
"seconds": round(run["seconds"], 1) if run else 0,
|
|
"counted_at": run["counted_at"] if run else None,
|
|
"error": run["error"] if run else "",
|
|
"model_path": run["model_path"] if run else "",
|
|
})
|
|
|
|
rows = archive_index.assign_batch_numbers(rows)
|
|
if date:
|
|
rows = [r for r in rows if r["date_label"] == date]
|
|
return {"rows": rows, "totals": totals(rows),
|
|
"unindexed": sum(1 for r in rows if not r.get("started_at"))}
|
|
|
|
|
|
def totals(rows: List[dict]) -> dict:
|
|
"""Only rows with a ground truth score. An unmeasured video is not a
|
|
perfect one, and letting it into the denominator would say it was."""
|
|
scored = [r for r in rows
|
|
if r["ground_truth"] is not None and r["loading"] is not None]
|
|
predicted = sum(r["loading"] for r in scored)
|
|
truth = sum(r["ground_truth"] for r in scored)
|
|
return {
|
|
"counted_videos": sum(1 for r in rows if r["loading"] is not None),
|
|
"total_videos": len(rows),
|
|
"scored_videos": len(scored),
|
|
"predicted": predicted,
|
|
"ground_truth": truth,
|
|
"delta": predicted - truth,
|
|
"accuracy": round(100.0 * (1 - abs(predicted - truth) / truth), 2) if truth else None,
|
|
}
|
|
|
|
|
|
def set_ground_truth(project_id: int, video_rel: str, value: Optional[int]) -> dict:
|
|
"""Record what you actually counted. Kept even when no run exists yet, so
|
|
the truth can be entered while the recount is still queued."""
|
|
date_label, batch_label = _split_rel(video_rel)
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""INSERT INTO count_runs (project_id, video_rel, date_label, batch_label,
|
|
ground_truth)
|
|
VALUES (?, ?, ?, ?, ?)
|
|
ON CONFLICT(project_id, video_rel)
|
|
DO UPDATE SET ground_truth = excluded.ground_truth""",
|
|
(project_id, video_rel, date_label, batch_label, value),
|
|
)
|
|
cur.execute("SELECT * FROM count_runs WHERE project_id = ? AND video_rel = ?",
|
|
(project_id, video_rel))
|
|
return dict(cur.fetchone())
|
|
|
|
|
|
def _store(project_id: int, video_rel: str, result: dict, params: dict,
|
|
model_path: str) -> None:
|
|
import json
|
|
|
|
date_label, batch_label = _split_rel(video_rel)
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""INSERT INTO count_runs (project_id, video_rel, date_label, batch_label,
|
|
loading, unloading, net, frames, seconds, params,
|
|
model_path, error, counted_at)
|
|
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(project_id, video_rel) DO UPDATE SET
|
|
loading = excluded.loading, unloading = excluded.unloading,
|
|
net = excluded.net, frames = excluded.frames,
|
|
seconds = excluded.seconds, params = excluded.params,
|
|
model_path = excluded.model_path, error = excluded.error,
|
|
counted_at = excluded.counted_at""",
|
|
(project_id, video_rel, date_label, batch_label,
|
|
result.get("loading"), result.get("unloading"), result.get("net"),
|
|
result.get("frames", 0), result.get("seconds", 0.0),
|
|
json.dumps(params), model_path, result.get("error", ""), time.time()),
|
|
)
|
|
|
|
|
|
# ---- the run itself ------------------------------------------------------
|
|
|
|
def count_video(path: str, model, params: dict, should_cancel=None,
|
|
on_progress=None) -> dict:
|
|
"""Count one video end to end. No drawing, no encoding — just the numbers."""
|
|
import cv2
|
|
|
|
# Imported first: `live_count` is what puts `algoritma-batch` on sys.path, so
|
|
# the `src.*` imports below only resolve once it has been loaded.
|
|
from backend.live_count import _too_small
|
|
from src.counting import LineCrossCounter
|
|
from src.stabilizer import BboxStabilizer
|
|
from src.tracking import ByteTrackTracker
|
|
|
|
settings = {**DEFAULTS, **params}
|
|
capture = cv2.VideoCapture(path)
|
|
if not capture.isOpened():
|
|
raise CountingBenchError(f"Could not open {path}")
|
|
total_frames = int(capture.get(cv2.CAP_PROP_FRAME_COUNT) or 0)
|
|
|
|
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(
|
|
line_y=settings["line_y"], line_x_start=settings["line_x_start"],
|
|
line_x_end=settings["line_x_end"], margin=settings["margin"],
|
|
dedup_radius=settings["dedup_radius"],
|
|
entry_travel_min=settings["entry_travel_min"],
|
|
handoff_radius=settings["handoff_radius"],
|
|
unload_confirm_frames=settings["unload_confirm_frames"],
|
|
spatial_dedup=settings["spatial_dedup"],
|
|
)
|
|
|
|
started = time.time()
|
|
frames = 0
|
|
try:
|
|
while True:
|
|
if should_cancel is not None and should_cancel():
|
|
break
|
|
ok, frame = capture.read()
|
|
if not ok or frame is None:
|
|
break
|
|
frame = cv2.resize(frame, (1280, 720))
|
|
# 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
|
|
if not _too_small(d.bbox, settings["min_area_scale"])
|
|
and settings["line_x_start"] <= (d.bbox[0] + d.bbox[2]) / 2 <= settings["line_x_end"]
|
|
]
|
|
counter.update(inside)
|
|
counter.drain_traces() # bounded memory; traces are the live view's job
|
|
frames += 1
|
|
if on_progress is not None and frames % 50 == 0:
|
|
on_progress(frames, total_frames)
|
|
finally:
|
|
capture.release()
|
|
|
|
return {
|
|
"loading": counter.loading_count,
|
|
"unloading": counter.unloading_count,
|
|
"net": counter.net_count,
|
|
"frames": frames,
|
|
"seconds": round(time.time() - started, 1),
|
|
}
|
|
|
|
|
|
def queue(project_id: int, video_rels: List[str], model_path: str,
|
|
params: Optional[dict] = None, recount: bool = False) -> dict:
|
|
"""Queue one job for the whole selection (REQ-152).
|
|
|
|
One job rather than one per video: they share a model load, and the GPU can
|
|
only run them one at a time anyway.
|
|
"""
|
|
project = projects.get(project_id)
|
|
if project is None:
|
|
raise CountingBenchError("No such project")
|
|
if not video_rels:
|
|
raise CountingBenchError("Pick at least one video to count")
|
|
if not os.path.isfile(model_path):
|
|
raise CountingBenchError(f"Model not found: {model_path}")
|
|
|
|
if not recount:
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT video_rel FROM count_runs
|
|
WHERE project_id = ? AND loading IS NOT NULL""",
|
|
(project_id,),
|
|
)
|
|
done = {row[0] for row in cur.fetchall()}
|
|
video_rels = [rel for rel in video_rels if rel not in done]
|
|
if not video_rels:
|
|
raise CountingBenchError(
|
|
"Every video in this selection has already been counted — "
|
|
"tick 'recount' to run them again")
|
|
|
|
job = jobs.create(
|
|
"count",
|
|
params={"video_rels": video_rels, "model_path": model_path,
|
|
"params": {**DEFAULTS, **(params or {})}},
|
|
project_id=project_id,
|
|
message=f"{len(video_rels)} video(s)",
|
|
)
|
|
return job.to_dict()
|
|
|
|
|
|
@jobs.handler("count")
|
|
def _run_count(job) -> None:
|
|
from ultralytics import YOLO
|
|
import numpy as np
|
|
|
|
project = projects.get(job.project_id)
|
|
rels = job.params["video_rels"]
|
|
settings = job.params.get("params") or DEFAULTS
|
|
model_path = job.params["model_path"]
|
|
|
|
job.progress(0, len(rels))
|
|
job.log(f"Counting {len(rels)} video(s) with {os.path.basename(model_path)}")
|
|
|
|
model = YOLO(model_path)
|
|
# Same warm-up as the live path: the first CUDA call inside the tracker has
|
|
# been seen to segfault without it.
|
|
model(np.zeros((720, 1280, 3), dtype=np.uint8), imgsz=settings.get("imgsz", 640),
|
|
verbose=False)
|
|
|
|
for index, rel in enumerate(rels):
|
|
if job.cancelled:
|
|
job.log(f"Cancelled after {index} video(s)")
|
|
return
|
|
try:
|
|
path = library.resolve(project["video_root"], rel)
|
|
except library.LibraryError as exc:
|
|
_store(job.project_id, rel, {"error": str(exc)}, settings, model_path)
|
|
job.log(f"{rel}: {exc}")
|
|
job.progress(index + 1, len(rels))
|
|
continue
|
|
|
|
def report(done, total, rel=rel, index=index):
|
|
job.progress(index, len(rels), f"{rel} — {done}/{total or '?'} frames")
|
|
|
|
try:
|
|
result = count_video(path, model, settings,
|
|
should_cancel=lambda: job.cancelled,
|
|
on_progress=report)
|
|
except Exception as exc:
|
|
_store(job.project_id, rel, {"error": f"{type(exc).__name__}: {exc}"},
|
|
settings, model_path)
|
|
job.log(f"{rel}: FAILED {exc}")
|
|
job.progress(index + 1, len(rels))
|
|
continue
|
|
|
|
# A cancel mid-video leaves a partial count, which would read as a real
|
|
# measurement of a video that was never finished.
|
|
if job.cancelled:
|
|
job.log(f"Cancelled during {rel} — its count was not saved")
|
|
return
|
|
_store(job.project_id, rel, result, settings, model_path)
|
|
rate = result["frames"] / result["seconds"] if result["seconds"] else 0
|
|
job.log(f"{rel}: in {result['loading']} out {result['unloading']} "
|
|
f"net {result['net']} ({result['frames']} frames, {rate:.0f} fps)")
|
|
job.progress(index + 1, len(rels))
|