Files
reTraining/backend/counting_bench.py
T

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))