"""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(), }