feat: per-class exemplar preview, class hide toggle, container NMS flag (REQ-182, REQ-183, REQ-184)
This commit is contained in:
1 parent
1a99f2ffe4
commit
96a00267d9
18 files changed
+537
-124
No files matched your search
@@ -37,6 +37,7 @@ class ProjectPatch(BaseModel):
|
||||
prompts: Optional[Dict[int, str]] = None
|
||||
val_every: Optional[int] = None
|
||||
video_root: Optional[str] = None
|
||||
containers: Optional[Dict[int, bool]] = None
|
||||
|
||||
|
||||
@router.get("")
|
||||
@@ -68,7 +69,8 @@ def patch_project(project_id: int, request: ProjectPatch) -> dict:
|
||||
try:
|
||||
return project_store.update(project_id, prompts=request.prompts,
|
||||
val_every=request.val_every,
|
||||
video_root=request.video_root)
|
||||
video_root=request.video_root,
|
||||
containers=request.containers)
|
||||
except project_store.ProjectError as exc:
|
||||
raise HTTPException(400, str(exc))
|
||||
|
||||
|
||||
@@ -166,6 +166,10 @@ def _run_autolabel(job) -> None:
|
||||
|
||||
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
|
||||
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
||||
# REQ-184: classes marked container keep boxes that sit inside them. The
|
||||
# prompt pass above still sees prompt-index ids, so it gets the index-aligned set.
|
||||
container_ids = {c["class_id"] for c in project["classes"] if c.get("container")}
|
||||
prompt_container_ids = {i for i, c in enumerate(sam3_target_classes) if c.get("container")}
|
||||
|
||||
job.log(f"Starting auto-labeling with {selected_engine}...")
|
||||
|
||||
@@ -249,6 +253,7 @@ def _run_autolabel(job) -> None:
|
||||
thresholds=thr_list,
|
||||
iou_by_class=dict(enumerate(iou_list)) if iou_list else None,
|
||||
min_box_fracs=mb_list,
|
||||
container_ids=prompt_container_ids,
|
||||
)
|
||||
if not res.error and res.detections:
|
||||
for det in res.detections:
|
||||
@@ -267,7 +272,8 @@ def _run_autolabel(job) -> None:
|
||||
and "iou_threshold" in per_class[c["name"].strip().lower()]
|
||||
}
|
||||
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_thresh,
|
||||
iou_by_class=iou_by_class_proj or None)
|
||||
iou_by_class=iou_by_class_proj or None,
|
||||
container_ids=container_ids)
|
||||
items = []
|
||||
for det in kept:
|
||||
if project["label_type"] == "bbox" or det.mask is None:
|
||||
|
||||
@@ -262,6 +262,13 @@ def migrate() -> None:
|
||||
# REQ-110: augmentation settings, null until the user changes them.
|
||||
if "augment" not in cols:
|
||||
cur.execute("ALTER TABLE projects ADD COLUMN augment TEXT")
|
||||
# REQ-184: a class flagged container keeps boxes that sit inside it.
|
||||
cur.execute("PRAGMA table_info(project_classes)")
|
||||
class_cols = [column[1] for column in cur.fetchall()]
|
||||
if "container" not in class_cols:
|
||||
cur.execute(
|
||||
"ALTER TABLE project_classes ADD COLUMN container INTEGER NOT NULL DEFAULT 0"
|
||||
)
|
||||
# REQ-107: what rule set a run's numbers were measured under.
|
||||
cur.execute("PRAGMA table_info(model_versions)")
|
||||
version_cols = [column[1] for column in cur.fetchall()]
|
||||
|
||||
+49
-8
@@ -10,7 +10,7 @@ calls — see the domain invariants in `../AGENTS.md`.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Dict, List, Optional
|
||||
from typing import Dict, List, Optional, Set
|
||||
|
||||
from PIL import Image
|
||||
|
||||
@@ -41,11 +41,29 @@ def _iou(box_a: List[float], box_b: List[float]) -> float:
|
||||
return inter / union if union > 0 else 0.0
|
||||
|
||||
|
||||
def deduplicate(detections: List[Detection], iou_threshold: float = 0.8,
|
||||
iou_by_class: Optional[Dict[int, float]] = None) -> List[Detection]:
|
||||
"""Greedy NMS per class: highest score wins within the SAME class.
|
||||
def _containment_fraction(box_a: List[float], box_b: List[float]) -> float:
|
||||
"""Intersection over the smaller box's area: how much of the smaller box
|
||||
sits inside the other one (REQ-184's containment test)."""
|
||||
ax0, ay0, ax1, ay1 = box_a
|
||||
bx0, by0, bx1, by1 = box_b
|
||||
inter_w = max(0.0, min(ax1, bx1) - max(ax0, bx0))
|
||||
inter_h = max(0.0, min(ay1, by1) - max(ay0, by0))
|
||||
area_a = max(0.0, ax1 - ax0) * max(0.0, ay1 - ay0)
|
||||
area_b = max(0.0, bx1 - bx0) * max(0.0, by1 - by0)
|
||||
small = min(area_a, area_b)
|
||||
return inter_w * inter_h / small if small > 0 else 0.0
|
||||
|
||||
`iou_by_class` overrides the threshold per class id (REQ-181)."""
|
||||
|
||||
def deduplicate(detections: List[Detection], iou_threshold: float = 0.8,
|
||||
iou_by_class: Optional[Dict[int, float]] = None,
|
||||
container_ids: Optional[Set[int]] = None) -> List[Detection]:
|
||||
"""Greedy NMS: highest score wins within the SAME class, then across
|
||||
classes (REQ-031) — unless the kept box is a container class and the
|
||||
candidate is at least 90% inside it, which is containment, not overlap
|
||||
(REQ-184).
|
||||
|
||||
`iou_by_class` overrides the threshold per class id (REQ-181);
|
||||
`container_ids` are the class ids marked container."""
|
||||
if not iou_by_class and iou_threshold <= 0.0:
|
||||
return detections
|
||||
|
||||
@@ -65,7 +83,26 @@ def deduplicate(detections: List[Detection], iou_threshold: float = 0.8,
|
||||
if all(_iou(det.box, k.box) < iou for k in cls_kept):
|
||||
cls_kept.append(det)
|
||||
kept.extend(cls_kept)
|
||||
return kept
|
||||
|
||||
# Cross-class pass (REQ-031), greedy score-desc. A pair whose resolved
|
||||
# threshold is <= 0 is never suppressed — zero means NMS off for that
|
||||
# class, mirroring the within-class pass above.
|
||||
cross_kept: List[Detection] = []
|
||||
for det in sorted(kept, key=lambda d: d.score, reverse=True):
|
||||
drop = False
|
||||
for k in cross_kept:
|
||||
if det.class_id == k.class_id:
|
||||
continue
|
||||
if (container_ids and k.class_id in container_ids
|
||||
and _containment_fraction(det.box, k.box) >= 0.9):
|
||||
continue
|
||||
iou = (iou_by_class or {}).get(det.class_id, iou_threshold)
|
||||
if iou > 0.0 and _iou(det.box, k.box) >= iou:
|
||||
drop = True
|
||||
break
|
||||
if not drop:
|
||||
cross_kept.append(det)
|
||||
return cross_kept
|
||||
|
||||
|
||||
def label_image(
|
||||
@@ -80,13 +117,16 @@ def label_image(
|
||||
thresholds: Optional[List[float]] = None,
|
||||
iou_by_class: Optional[Dict[int, float]] = None,
|
||||
min_box_fracs: Optional[List[float]] = None,
|
||||
container_ids: Optional[Set[int]] = None,
|
||||
) -> ImageResult:
|
||||
"""Detect every prompt in one image and return the surviving instances.
|
||||
|
||||
When `exemplars` are given, the prompt at `exemplar_index` also carries them
|
||||
as drawn box exemplars (REQ-172); every other prompt runs on text alone.
|
||||
`thresholds`, `iou_by_class` and `min_box_fracs` are per-prompt overrides
|
||||
aligned with `prompts` (REQ-181); classes without one use the global values."""
|
||||
aligned with `prompts` (REQ-181); classes without one use the global values.
|
||||
`container_ids` are prompt indices marked container (REQ-184) — the caller
|
||||
maps them, because detections still carry prompt-index class ids here."""
|
||||
try:
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
except Exception as exc: # unreadable/corrupt frame: report, don't abort the job
|
||||
@@ -121,4 +161,5 @@ def label_image(
|
||||
]
|
||||
|
||||
return ImageResult(image_path, rel_path, width, height,
|
||||
deduplicate(detections, iou_threshold, iou_by_class=iou_by_class))
|
||||
deduplicate(detections, iou_threshold, iou_by_class=iou_by_class,
|
||||
container_ids=container_ids))
|
||||
+7
-1
@@ -72,6 +72,10 @@ def preview_frame(
|
||||
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
||||
per_class = _parse_class_params(class_params)
|
||||
predict_conf = min([threshold] + [v["threshold"] for v in per_class.values() if "threshold" in v])
|
||||
# REQ-184: classes marked container keep boxes that sit inside them. The
|
||||
# prompt pass below still sees prompt-index ids, so it gets the index-aligned set.
|
||||
container_ids = {c["class_id"] for c in project["classes"] if c.get("container")}
|
||||
prompt_container_ids = {i for i, c in enumerate(sam3_target_classes) if c.get("container")}
|
||||
|
||||
all_raw_detections = []
|
||||
|
||||
@@ -153,6 +157,7 @@ def preview_frame(
|
||||
thresholds=thr_list,
|
||||
iou_by_class=dict(enumerate(iou_list)) if iou_list else None,
|
||||
min_box_fracs=mb_list,
|
||||
container_ids=prompt_container_ids,
|
||||
)
|
||||
if not res.error and res.detections:
|
||||
for det in res.detections:
|
||||
@@ -169,7 +174,8 @@ def preview_frame(
|
||||
and "iou_threshold" in per_class[c["name"].strip().lower()]
|
||||
}
|
||||
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_threshold,
|
||||
iou_by_class=iou_by_class_proj or None)
|
||||
iou_by_class=iou_by_class_proj or None,
|
||||
container_ids=container_ids)
|
||||
items = []
|
||||
for det in kept:
|
||||
if project["label_type"] == "bbox" or det.mask is None:
|
||||
|
||||
+25
-15
@@ -13,7 +13,7 @@ import os
|
||||
import re
|
||||
import shutil
|
||||
import time
|
||||
from typing import List, Optional
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from backend import config, db
|
||||
|
||||
@@ -124,15 +124,17 @@ def create(name: str, label_type: str, video_root: str, classes: Optional[List[d
|
||||
def _write_classes(cur, project_id: int, classes: List[dict]) -> None:
|
||||
cur.execute("DELETE FROM project_classes WHERE project_id = ?", (project_id,))
|
||||
cur.executemany(
|
||||
"INSERT INTO project_classes (project_id, class_id, name, prompt) VALUES (?, ?, ?, ?)",
|
||||
[(project_id, index, item["name"], item["prompt"])
|
||||
"INSERT INTO project_classes (project_id, class_id, name, prompt, container) "
|
||||
"VALUES (?, ?, ?, ?, ?)",
|
||||
[(project_id, index, item["name"], item["prompt"], item.get("container", 0))
|
||||
for index, item in enumerate(classes)],
|
||||
)
|
||||
|
||||
|
||||
def _row_to_dict(cur, row) -> dict:
|
||||
cur.execute(
|
||||
"SELECT class_id, name, prompt FROM project_classes WHERE project_id = ? ORDER BY class_id",
|
||||
"SELECT class_id, name, prompt, container FROM project_classes "
|
||||
"WHERE project_id = ? ORDER BY class_id",
|
||||
(row["id"],),
|
||||
)
|
||||
classes = [dict(item) for item in cur.fetchall()]
|
||||
@@ -205,9 +207,10 @@ def listing() -> List[dict]:
|
||||
|
||||
|
||||
def update(project_id: int, prompts: Optional[dict] = None, val_every: Optional[int] = None,
|
||||
video_root: Optional[str] = None) -> dict:
|
||||
"""Edit the things that are safe to change: prompts, split ratio, archive root.
|
||||
Class names and label type are not among them."""
|
||||
video_root: Optional[str] = None,
|
||||
containers: Optional[Dict[int, bool]] = None) -> dict:
|
||||
"""Edit the things that are safe to change: prompts, container flags,
|
||||
split ratio, archive root. Class names and label type are not among them."""
|
||||
if get(project_id) is None:
|
||||
raise ProjectError("No such project")
|
||||
|
||||
@@ -226,6 +229,11 @@ def update(project_id: int, prompts: Optional[dict] = None, val_every: Optional[
|
||||
"UPDATE project_classes SET prompt = ? WHERE project_id = ? AND class_id = ?",
|
||||
(str(prompt).strip(), project_id, int(class_id)),
|
||||
)
|
||||
for class_id, flag in (containers or {}).items():
|
||||
cur.execute(
|
||||
"UPDATE project_classes SET container = ? WHERE project_id = ? AND class_id = ?",
|
||||
(1 if flag else 0, project_id, int(class_id)),
|
||||
)
|
||||
return get(project_id)
|
||||
|
||||
|
||||
@@ -344,9 +352,7 @@ def set_base_model(project_id: int, weights_path: str) -> dict:
|
||||
"UPDATE projects SET base_model_path = ?, base_model_kind = 'uploaded' WHERE id = ?",
|
||||
(stored, project_id),
|
||||
)
|
||||
_write_classes(cur, project_id,
|
||||
[{"name": n, "prompt": p}
|
||||
for n, p in zip(names, _kept_prompts(project, names))])
|
||||
_write_classes(cur, project_id, _kept_classes(project, names))
|
||||
return get(project_id)
|
||||
|
||||
|
||||
@@ -372,11 +378,15 @@ def set_secondary_model(project_id: int, weights_path: str, name: str = "") -> d
|
||||
return get(project_id)
|
||||
|
||||
|
||||
def _kept_prompts(project: dict, names: List[str]) -> List[str]:
|
||||
"""Keep the prompt the user already wrote for a class that survives a
|
||||
base-model swap; fall back to the class name for new ones."""
|
||||
known = {item["name"]: item["prompt"] for item in project["classes"]}
|
||||
return [known.get(name, name) for name in names]
|
||||
def _kept_classes(project: dict, names: List[str]) -> List[dict]:
|
||||
"""Keep the stored per-class attributes — prompt (REQ-171) and container
|
||||
flag (REQ-184) — for a class that survives a base-model swap; new classes
|
||||
fall back to the class name as prompt and no flag."""
|
||||
known = {item["name"]: item for item in project["classes"]}
|
||||
return [{"name": name,
|
||||
"prompt": known[name]["prompt"] if name in known else name,
|
||||
"container": known[name].get("container", 0) if name in known else 0}
|
||||
for name in names]
|
||||
|
||||
|
||||
def training_start_point(project: dict) -> str:
|
||||
|
||||
Reference in new issue
Block a user