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.
532 lines
21 KiB
Python
532 lines
21 KiB
Python
"""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(),
|
|
}
|