feat: per-class exemplar preview, class hide toggle, container NMS flag (REQ-182, REQ-183, REQ-184)

This commit is contained in:
asus committed 2026-10-02 11:11:33 +07:00
1 parent 1a99f2ffe4
commit 96a00267d9
18 files changed
+537 -124

No files matched your search

+3 -1
View File
@@ -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))
+7 -1
View File
@@ -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:
+7
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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: