feat: add counting bench, triage, and dataset modules

This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
This commit is contained in:
asus committed 2026-08-14 16:28:52 +07:00
1 parent 8285400254
commit 5c7c122105
80 files changed
+20074 -1412

No files matched your search

+531
View File
@@ -0,0 +1,531 @@
"""Triage: deciding what each SAM3 shape is actually worth (REQ-100…108).
A shape is never rewritten. Its verdict is *resolved* every time it is needed:
manual override > first matching rule > keep
so `annotations.class_id` keeps whatever SAM3 said, and any rule can be re-cut
later against the original output. That is the whole reason rules are evaluated
at training time rather than baked in at merge (REQ-102).
A verdict is one of:
keep the shape trains as its own class
reclass -> class_id the shape trains as a different class (REQ-105)
ignore the box is dropped; its image still trains (REQ-104)
`ignore` drops the box rather than the image because these frames are dense —
around 44 shapes each. Excluding the whole image was measured against a real
batch and cost 96% of it (1,882 of 1,950 frames) to remove 10% of the boxes.
Dropping four boxes out of forty-four leaves the image overwhelmingly correct;
dropping the image leaves nothing to train on.
The exception is a frame that loses *every* shape it had: an empty label file
says "there is nothing here", and for a frame that was full of sacks that is a
lie the model will learn. Those images are excluded.
"""
import hashlib
import json
import time
from typing import List, Optional
from backend import db, review
class TriageError(Exception):
pass
# ---- rules ---------------------------------------------------------------
def rules(project_id: int, stage: str = "dataprep") -> List[dict]:
with db.cursor() as cur:
cur.execute(
"""SELECT * FROM triage_rules WHERE project_id = ? AND stage = ?
ORDER BY position""",
(project_id, stage),
)
return [_rule_dict(row) for row in cur.fetchall()]
def _rule_dict(row) -> dict:
return {
"id": row["id"],
"stage": row["stage"],
"position": row["position"],
"name": row["name"],
"predicate": json.loads(row["predicate"]),
"action": row["action"],
"target_class": row["target_class"],
}
RANGE_FIELDS = ("score", "area_pct", "aspect")
def _validate_predicate(name, predicate: dict) -> None:
"""Reject a malformed rule here rather than inside a merge job.
A bad predicate used to be stored happily and only raise when the resolver
reached it — by which time the batch had been flipped to `approved` and the
user was looking at a failed job quoting a Python unpacking error.
"""
label = name or "this rule"
for field, bounds in predicate.items():
if field == "class_id":
if bounds is not None and not isinstance(bounds, int):
raise TriageError(f"{label}: class_id must be a class number")
continue
if field not in RANGE_FIELDS:
raise TriageError(
f"{label}: '{field}' is not something a rule can test "
f"(use {', '.join(RANGE_FIELDS)} or class_id)")
if bounds is None:
continue
if not isinstance(bounds, (list, tuple)) or len(bounds) != 2:
raise TriageError(f"{label}: '{field}' needs a [minimum, maximum] pair")
low, high = bounds
for edge in (low, high):
if edge is not None and not isinstance(edge, (int, float)):
raise TriageError(f"{label}: '{field}' bounds must be numbers or blank")
if low is not None and high is not None and low > high:
raise TriageError(
f"{label}: '{field}' minimum {low} is above its maximum {high}, "
"so the rule can never match")
def replace_rules(project_id: int, incoming: List[dict], stage: str = "dataprep") -> List[dict]:
"""Store the whole ordered list — the UI edits it as one thing (REQ-100)."""
for item in incoming:
if item.get("action") not in ("keep", "ignore", "reclass"):
raise TriageError(f"Unknown action: {item.get('action')}")
if item["action"] == "reclass" and item.get("target_class") is None:
raise TriageError(f"Rule '{item.get('name')}' reclassifies but names no target class")
_validate_predicate(item.get("name"), item.get("predicate") or {})
with db.cursor() as cur:
cur.execute("DELETE FROM triage_rules WHERE project_id = ? AND stage = ?",
(project_id, stage))
for position, item in enumerate(incoming):
cur.execute(
"""INSERT INTO triage_rules (project_id, stage, position, name, predicate,
action, target_class, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
(project_id, stage, position, item.get("name") or f"rule {position + 1}",
json.dumps(item.get("predicate") or {}), item["action"],
item.get("target_class"), time.time()),
)
return rules(project_id, stage)
# ---- overrides -----------------------------------------------------------
def set_overrides(annotation_ids: List[int], verdict: str,
target_class: Optional[int] = None) -> int:
"""A hand decision outranks every rule, now and after any rule edit (REQ-103)."""
if verdict not in ("keep", "ignore", "reclass"):
raise TriageError(f"Unknown verdict: {verdict}")
if verdict == "reclass" and target_class is None:
raise TriageError("A reclass override needs a target class")
now = time.time()
with db.cursor() as cur:
cur.executemany(
"""INSERT INTO annotation_overrides (annotation_id, verdict, target_class, decided_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(annotation_id) DO UPDATE SET
verdict = excluded.verdict,
target_class = excluded.target_class,
decided_at = excluded.decided_at""",
[(aid, verdict, target_class, now) for aid in annotation_ids],
)
return cur.rowcount
def clear_overrides(annotation_ids: List[int]) -> int:
if not annotation_ids:
return 0
with db.cursor() as cur:
placeholders = ",".join("?" for _ in annotation_ids)
cur.execute(f"DELETE FROM annotation_overrides WHERE annotation_id IN ({placeholders})",
annotation_ids)
return cur.rowcount
# ---- resolution ----------------------------------------------------------
def metrics(geometry: dict) -> dict:
"""The signals a rule can test, all derived from the box."""
x0, y0, x1, y1 = review.to_box(geometry)
width = max(0.0, x1 - x0)
height = max(0.0, y1 - y0)
return {
"area_pct": round(width * height * 100.0, 4),
"aspect": round(width / height, 4) if height > 0 else 0.0,
}
def _in_range(value: float, bounds) -> bool:
low, high = bounds
return (low is None or value >= low) and (high is None or value <= high)
def _matches(predicate: dict, shape: dict) -> bool:
if "class_id" in predicate and predicate["class_id"] is not None:
if shape["class_id"] != predicate["class_id"]:
return False
for field in ("score", "area_pct", "aspect"):
if predicate.get(field) is not None and not _in_range(shape[field], predicate[field]):
return False
return True
class Resolver:
"""Holds a project's rules and overrides so a whole dataset can be resolved
without re-reading them per shape."""
def __init__(self, project_id: int, stage: str = "dataprep",
frozen: Optional[List[dict]] = None):
# `frozen` is a dataset's snapshot (REQ-132): the merge that cut it runs
# under those rules, not under whatever the project says today.
self.rules = list(frozen) if frozen is not None else rules(project_id, stage)
with db.cursor() as cur:
cur.execute("SELECT annotation_id, verdict, target_class FROM annotation_overrides")
self.overrides = {row[0]: (row[1], row[2]) for row in cur.fetchall()}
def verdict(self, shape: dict) -> dict:
"""Resolve one shape. `shape` needs id, class_id, score, area_pct, aspect."""
override = self.overrides.get(shape["id"])
if override is not None:
verdict, target = override
return {"verdict": verdict, "target_class": target, "source": "manual"}
for rule in self.rules:
if _matches(rule["predicate"], shape):
return {"verdict": rule["action"], "target_class": rule["target_class"],
"source": rule["name"]}
return {"verdict": "keep", "target_class": None, "source": "default"}
def resolve_shapes(self, annotations: list) -> Optional[list]:
"""Apply verdicts to one frame's annotations.
Returns the surviving annotations with their effective class, or None
when the frame must not train at all — which now happens only if every
shape was dropped.
"""
kept = []
for item in annotations:
shape = {"id": item["id"], "class_id": item["class_id"],
"score": float(item.get("score") or 1.0),
**metrics(item["geometry"])}
effective = self.effective_class(shape)
if effective is None:
continue
kept.append({**item, "class_id": effective})
if annotations and not kept:
return None
return kept
def effective_class(self, shape: dict) -> Optional[int]:
"""The class this shape trains as, or None when the box is dropped."""
resolved = self.verdict(shape)
if resolved["verdict"] == "ignore":
return None
if resolved["verdict"] == "reclass":
return resolved["target_class"]
return shape["class_id"]
def version(self) -> str:
"""A short hash of what this resolver would do (REQ-107).
Two runs with the same version measured the same thing; two runs with
different versions did not, because a rule edit can change which images
are in the val set.
"""
payload = json.dumps(
{"rules": self.rules, "overrides": sorted(self.overrides.items())},
sort_keys=True, default=str,
)
return hashlib.sha1(payload.encode("utf-8")).hexdigest()[:12]
def shapes_for_frames(frame_ids: List[int]) -> List[dict]:
"""Every annotation on these frames, with its signals and resolved verdict."""
if not frame_ids:
return []
with db.cursor() as cur:
placeholders = ",".join("?" for _ in frame_ids)
cur.execute(
f"""SELECT a.id, a.frame_id, a.class_id, a.score, a.source, a.geometry
FROM annotations a WHERE a.frame_id IN ({placeholders})""",
frame_ids,
)
rows = cur.fetchall()
shapes = []
for row in rows:
geometry = json.loads(row["geometry"])
shape = {
"id": row["id"],
"frame_id": row["frame_id"],
"class_id": row["class_id"],
"score": round(float(row["score"]), 4),
"origin": row["source"],
"box": review.to_box(geometry),
**metrics(geometry),
}
shapes.append(shape)
return shapes
SCATTER_POINTS = 4000
"""How many dots the scatter gets. A real batch runs to ~85k shapes; every one of
them as an SVG circle locks the browser, and a boundary between two clusters is
just as visible in a few thousand points. The verdict tallies are still counted
over every shape, so the numbers are never a sample."""
def as_ids(batch_ids) -> List[int]:
"""One batch or many — Data Prep now tunes a whole selection at once (REQ-130)."""
if isinstance(batch_ids, int):
return [batch_ids]
if isinstance(batch_ids, str):
return [int(part) for part in batch_ids.split(",") if part.strip().lstrip("-").isdigit()]
return list(batch_ids)
def _resolved_shapes(batch_ids):
from backend import batches
ids = as_ids(batch_ids)
found = [batches.get(bid) for bid in ids]
if not ids or any(batch is None for batch in found):
raise TriageError("No such batch")
if len({batch["project_id"] for batch in found}) > 1:
raise TriageError("Those batches are not all in the same project")
batch = found[0]
with db.cursor() as cur:
placeholders = ",".join("?" for _ in ids)
cur.execute(
f"SELECT id FROM frames WHERE batch_id IN ({placeholders}) ORDER BY batch_id, idx",
ids,
)
frame_ids = [row[0] for row in cur.fetchall()]
resolver = Resolver(batch["project_id"])
shapes = shapes_for_frames(frame_ids)
for shape in shapes:
shape.update(resolver.verdict(shape))
return batch, frame_ids, shapes, resolver
def batch_summary(batch_ids) -> dict:
"""Verdict tallies over the whole selection, plus a sample to plot (REQ-106)."""
ids = as_ids(batch_ids)
batch, frame_ids, shapes, resolver = _resolved_shapes(ids)
counts = {"keep": 0, "ignore": 0, "reclass": 0, "manual": 0}
per_frame = {}
for shape in shapes:
counts[shape["verdict"]] += 1
if shape["source"] == "manual":
counts["manual"] += 1
total, dropped = per_frame.get(shape["frame_id"], (0, 0))
per_frame[shape["frame_id"]] = (total + 1, dropped + (shape["verdict"] == "ignore"))
# Only a frame that loses everything is held back; the rest keep training
# with their surviving boxes.
ignored_frames = {fid for fid, (total, dropped) in per_frame.items() if total == dropped}
# An even stride rather than a random draw: the sample is stable across
# reloads, so points do not jump around while the user is reading the plot.
stride = max(1, len(shapes) // SCATTER_POINTS)
sample = [
{k: shape[k] for k in ("id", "class_id", "score", "area_pct", "aspect", "verdict", "source")}
for shape in shapes[::stride][:SCATTER_POINTS]
]
return {
"batch_ids": ids,
"project_id": batch["project_id"],
"status": batch["status"],
"merged": batch["status"] == "merged",
"frame_count": len(frame_ids),
"total_shapes": len(shapes),
"counts": counts,
# What merging this batch would do right now (REQ-104).
"frames_held_back": len(ignored_frames),
"frames_would_merge": len(frame_ids) - len(ignored_frames),
"sample": sample,
"sampled": len(sample) < len(shapes),
"rule_version": resolver.version(),
}
def batch_page(batch_ids, sort: str = "score", offset: int = 0, limit: int = 120) -> dict:
"""One page of shapes for the crop grid, sorted server-side so the client
never holds the whole batch."""
if sort not in ("score", "area_pct"):
raise TriageError(f"Cannot sort by {sort}")
_, _, shapes, _ = _resolved_shapes(batch_ids)
shapes.sort(key=lambda shape: shape[sort])
page = shapes[offset:offset + limit]
for shape in page:
shape.pop("box", None)
return {"total": len(shapes), "offset": offset, "limit": limit, "shapes": page}
def _percentile(values: list, fraction: float) -> float:
if not values:
return 0.0
return values[min(len(values) - 1, int(len(values) * fraction))]
def suggest(batch_ids) -> dict:
"""Presets with thresholds read off this batch's own distribution.
Asking someone to invent "score below 0.45" from nothing is guesswork. The
same question is easy when the number comes from their data and the effect
is stated: "the weakest 10% of detections — 8,552 shapes".
"""
_, frame_ids, shapes, _ = _resolved_shapes(batch_ids)
if not shapes:
return {"presets": [], "stats": {}}
scores = sorted(shape["score"] for shape in shapes)
areas = sorted(shape["area_pct"] for shape in shapes)
total = len(shapes)
def impact(predicate: dict) -> dict:
matched = [s for s in shapes if _matches(predicate, s)]
frames = {s["frame_id"] for s in matched}
return {"shapes": len(matched), "frames": len(frames)}
presets = []
weak = round(_percentile(scores, 0.10), 3)
presets.append({
"key": "drop-weakest",
"title": "Ignore the weakest detections",
"blurb": f"SAM3 scored these below {weak} — the bottom 10% of this batch.",
"rule": {"name": "low confidence", "predicate": {"score": [None, weak]}, "action": "ignore"},
"impact": impact({"score": [None, weak]}),
})
specks = round(_percentile(areas, 0.05), 3)
presets.append({
"key": "drop-specks",
"title": "Ignore tiny specks",
"blurb": f"Boxes smaller than {specks}% of the frame — usually noise, not objects.",
"rule": {"name": "specks", "predicate": {"area_pct": [None, specks]}, "action": "ignore"},
"impact": impact({"area_pct": [None, specks]}),
})
median_area = round(_percentile(areas, 0.50), 3)
presets.append({
"key": "split-by-size",
"title": "Split by size into a second class",
"blurb": f"Everything under {median_area}% area (half this batch) becomes another class — "
"pick which one. Size tracks distance from the camera as much as object type, "
"so check the crops before trusting it.",
"rule": {"name": "small ones", "predicate": {"area_pct": [None, median_area]},
"action": "reclass", "target_class": None},
"impact": impact({"area_pct": [None, median_area]}),
"needs_target": True,
})
tall = round(_percentile(sorted(s["aspect"] for s in shapes), 0.15), 3)
presets.append({
"key": "odd-shapes",
"title": "Ignore oddly-shaped boxes",
"blurb": f"Aspect ratio under {tall} — long thin slivers, usually a bad mask.",
"rule": {"name": "slivers", "predicate": {"aspect": [None, tall]}, "action": "ignore"},
"impact": impact({"aspect": [None, tall]}),
})
return {
"presets": presets,
"stats": {
"total_shapes": total,
"total_frames": len(frame_ids),
"score": {"p05": round(_percentile(scores, 0.05), 3),
"p50": round(_percentile(scores, 0.50), 3),
"p95": round(_percentile(scores, 0.95), 3)},
"area_pct": {"p05": round(_percentile(areas, 0.05), 3),
"p50": round(_percentile(areas, 0.50), 3),
"p95": round(_percentile(areas, 0.95), 3)},
},
}
def simulate(batch_ids, candidate_rules: List[dict]) -> dict:
"""What these rules would do, without saving them.
Editing a threshold and seeing the number move is the whole difference
between tuning a filter and guessing at one.
"""
_, frame_ids, shapes, _ = _resolved_shapes(batch_ids)
with db.cursor() as cur:
cur.execute("SELECT annotation_id, verdict, target_class FROM annotation_overrides")
manual = {row[0]: (row[1], row[2]) for row in cur.fetchall()}
counts = {"keep": 0, "ignore": 0, "reclass": 0}
per_rule = [0] * len(candidate_rules)
per_frame = {}
for shape in shapes:
if shape["id"] in manual:
verdict = manual[shape["id"]][0]
else:
verdict = "keep"
for index, rule in enumerate(candidate_rules):
if _matches(rule.get("predicate") or {}, shape):
verdict = rule["action"]
per_rule[index] += 1
break
counts[verdict] += 1
total, dropped = per_frame.get(shape["frame_id"], (0, 0))
per_frame[shape["frame_id"]] = (total + 1, dropped + (verdict == "ignore"))
ignored_frames = {fid for fid, (total, dropped) in per_frame.items() if total == dropped}
return {
"total_shapes": len(shapes),
"counts": counts,
"per_rule": per_rule,
"frames_held_back": len(ignored_frames),
"frames_would_merge": len(frame_ids) - len(ignored_frames),
}
def preview(project_id: int) -> dict:
"""What the current rules would do to the whole merged dataset."""
with db.cursor() as cur:
cur.execute(
"SELECT DISTINCT frame_id FROM dataset_items WHERE project_id = ?", (project_id,))
frame_ids = [row[0] for row in cur.fetchall()]
resolver = Resolver(project_id)
shapes = shapes_for_frames(frame_ids)
counts = {"keep": 0, "ignore": 0, "reclass": 0}
per_class: dict = {}
per_frame = {}
for shape in shapes:
resolved = resolver.verdict(shape)
counts[resolved["verdict"]] += 1
total, dropped = per_frame.get(shape["frame_id"], (0, 0))
per_frame[shape["frame_id"]] = (total + 1, dropped + (resolved["verdict"] == "ignore"))
if resolved["verdict"] == "ignore":
continue
effective = resolver.effective_class(shape)
per_class[effective] = per_class.get(effective, 0) + 1
excluded_images = {fid for fid, (total, dropped) in per_frame.items() if total == dropped}
return {
"total_shapes": len(shapes),
"total_images": len(frame_ids),
**counts,
"excluded_images": len(excluded_images),
"trainable_images": len(frame_ids) - len(excluded_images),
"per_class": per_class,
"rule_version": resolver.version(),
}