feat: per-class autolabel params + frame-scoped class clear (REQ-180, REQ-181)
- REQ-181: class_params {name: {threshold?, iou_threshold?, min_box_frac?}}
on /preview and /autolabel, per-class override table in both modals;
empty overrides take the unchanged global path
- REQ-180: x button on each review sidebar class row clears that class on
the current frame only via bulk-delete, no confirmation
- includes REQ-178 empty date-folder cycle fix (archive_index.py)
This commit is contained in:
1 parent
c244ab9795
commit
1a99f2ffe4
16 files changed
+372
-45
No files matched your search
@@ -35,6 +35,7 @@ class AutolabelRequest(BaseModel):
|
||||
resume: bool = False
|
||||
append: bool = False
|
||||
custom_model_path: Optional[str] = None
|
||||
class_params: Optional[dict[str, dict[str, float]]] = None
|
||||
|
||||
class Exemplar(BaseModel):
|
||||
box: list[float] # [cx, cy, w, h], normalized 0..1
|
||||
@@ -53,6 +54,7 @@ class PreviewRequest(BaseModel):
|
||||
# Preview-only — `autolabel.start` deliberately has no equivalent.
|
||||
exemplars: Optional[list[Exemplar]] = None
|
||||
exemplar_class_name: Optional[str] = None
|
||||
class_params: Optional[dict[str, dict[str, float]]] = None
|
||||
|
||||
|
||||
@router.post("/api/projects/{project_id}/batches")
|
||||
@@ -116,7 +118,8 @@ def start_autolabel(batch_id: int, request: AutolabelRequest) -> dict:
|
||||
engine=request.engine, engines=engine_list, class_ids=request.class_ids,
|
||||
engine_classes=request.engine_classes,
|
||||
target_class_names=request.target_class_names,
|
||||
custom_model_path=request.custom_model_path)
|
||||
custom_model_path=request.custom_model_path,
|
||||
class_params=request.class_params)
|
||||
except batch_store.BatchError as exc:
|
||||
raise HTTPException(400, str(exc))
|
||||
@router.post("/api/batches/inspect-model")
|
||||
@@ -189,6 +192,7 @@ def preview_autolabel(batch_id: int, request: PreviewRequest) -> dict:
|
||||
custom_model_path=request.custom_model_path,
|
||||
exemplars=[e.model_dump() for e in request.exemplars or []],
|
||||
exemplar_class_name=request.exemplar_class_name,
|
||||
class_params=request.class_params,
|
||||
)
|
||||
return {"shapes": shapes}
|
||||
except Exception as exc:
|
||||
|
||||
@@ -167,7 +167,14 @@ def cycles(project_id: int) -> List[dict]:
|
||||
|
||||
buckets: dict = {}
|
||||
for day in library.list_dates(project["video_root"]):
|
||||
for name in _video_names(project, day["date"]):
|
||||
names = _video_names(project, day["date"])
|
||||
if not names:
|
||||
# Folder still holds no recording: it must still appear in the
|
||||
# list, or a folder just created with "Folder baru" is invisible.
|
||||
buckets.setdefault(day["date"], {"cycle": day["date"], "video_count": 0,
|
||||
"flagged": 0, "first_start": None})
|
||||
continue
|
||||
for name in names:
|
||||
rel = f"{day['date']}/{name}"
|
||||
timing = _sidecar(project, rel) or known.get(rel) or {}
|
||||
cycle = timing.get("working_day") or day["date"]
|
||||
|
||||
+59
-5
@@ -17,6 +17,26 @@ DEFAULT_THRESHOLD = 0.35
|
||||
DEFAULT_IOU = 0.0
|
||||
|
||||
|
||||
def _parse_class_params(raw) -> dict:
|
||||
"""Normalize `class_params` (REQ-181): lowercased class name → overrides.
|
||||
|
||||
Unknown keys and non-numeric values are dropped, so a malformed payload
|
||||
degrades to the global values instead of failing the job."""
|
||||
out: dict = {}
|
||||
for name, values in (raw or {}).items():
|
||||
if not isinstance(values, dict):
|
||||
continue
|
||||
entry = {}
|
||||
for key in ("threshold", "iou_threshold", "min_box_frac"):
|
||||
if key in values:
|
||||
try:
|
||||
entry[key] = float(values[key])
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
if entry:
|
||||
out[str(name).strip().lower()] = entry
|
||||
return out
|
||||
|
||||
|
||||
|
||||
def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
||||
@@ -26,7 +46,8 @@ def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
||||
class_ids: Optional[List[int]] = None,
|
||||
engine_classes: Optional[dict[str, List[str]]] = None,
|
||||
custom_model_path: Optional[str] = None,
|
||||
target_class_names: Optional[List[str]] = None) -> dict:
|
||||
target_class_names: Optional[List[str]] = None,
|
||||
class_params: Optional[dict] = None) -> dict:
|
||||
batch = batches.get(batch_id)
|
||||
if batch is None:
|
||||
raise batches.BatchError("No such batch")
|
||||
@@ -41,7 +62,8 @@ def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
||||
"iou_threshold": iou_threshold, "min_box_frac": min_box_frac,
|
||||
"resume": resume, "append": append, "engine": active_engines[0], "engines": active_engines,
|
||||
"class_ids": class_ids, "engine_classes": engine_classes,
|
||||
"custom_model_path": custom_model_path, "target_class_names": target_class_names},
|
||||
"custom_model_path": custom_model_path, "target_class_names": target_class_names,
|
||||
"class_params": class_params},
|
||||
project_id=batch["project_id"],
|
||||
batch_id=batch_id,
|
||||
message=f"{batch['date_label']}/{batch['batch_label']} ({'+'.join(e.upper() for e in active_engines)})",
|
||||
@@ -85,6 +107,10 @@ def _run_autolabel(job) -> None:
|
||||
|
||||
conf = job.params.get("threshold", DEFAULT_THRESHOLD)
|
||||
iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
|
||||
per_class = _parse_class_params(job.params.get("class_params"))
|
||||
# Predict at the LOWEST threshold in play so a class with a lower override
|
||||
# can still see its boxes; each box then passes its own class gate below.
|
||||
predict_conf = min([conf] + [v["threshold"] for v in per_class.values() if "threshold" in v])
|
||||
|
||||
yolo_model = None
|
||||
sam3_target_classes = []
|
||||
@@ -159,7 +185,7 @@ def _run_autolabel(job) -> None:
|
||||
all_raw_detections = []
|
||||
|
||||
if yolo_model is not None:
|
||||
results = yolo_model.predict(frame_file, conf=conf, verbose=False)
|
||||
results = yolo_model.predict(frame_file, conf=predict_conf, verbose=False)
|
||||
if results and len(results) > 0:
|
||||
model_names = results[0].names
|
||||
for box in results[0].boxes:
|
||||
@@ -189,6 +215,15 @@ def _run_autolabel(job) -> None:
|
||||
|
||||
score = float(box.conf[0].item())
|
||||
xyxyn = box.xyxyn[0].tolist()
|
||||
|
||||
cp = per_class.get(proj_cls_name) or per_class.get(raw_cls_name)
|
||||
if cp:
|
||||
if score < cp.get("threshold", conf):
|
||||
continue
|
||||
mb = cp.get("min_box_frac")
|
||||
if mb and mb > 0 and (xyxyn[2]-xyxyn[0])*(xyxyn[3]-xyxyn[1]) < mb:
|
||||
continue
|
||||
|
||||
all_raw_detections.append(labeling.Detection(
|
||||
class_id=target_class_id,
|
||||
class_name=proj_cls_name or raw_cls_name,
|
||||
@@ -199,9 +234,21 @@ def _run_autolabel(job) -> None:
|
||||
|
||||
if selected_engine == "sam3" and sam3_target_classes:
|
||||
prompts = [(c.get("prompt") or c["name"]).strip() for c in sam3_target_classes]
|
||||
thr_list = iou_list = mb_list = None
|
||||
if per_class:
|
||||
names = [c["name"].strip().lower() for c in sam3_target_classes]
|
||||
if any("threshold" in v for v in per_class.values()):
|
||||
thr_list = [per_class.get(n, {}).get("threshold", conf) for n in names]
|
||||
if any("iou_threshold" in v for v in per_class.values()):
|
||||
iou_list = [per_class.get(n, {}).get("iou_threshold", iou_thresh) for n in names]
|
||||
if any("min_box_frac" in v for v in per_class.values()):
|
||||
mb_list = [per_class.get(n, {}).get("min_box_frac", job.params.get("min_box_frac", 0.0)) for n in names]
|
||||
res = labeling.label_image(
|
||||
frame_file, frame["filename"], prompts, conf,
|
||||
iou_threshold=iou_thresh, min_box_frac=job.params.get("min_box_frac", 0.0)
|
||||
iou_threshold=iou_thresh, min_box_frac=job.params.get("min_box_frac", 0.0),
|
||||
thresholds=thr_list,
|
||||
iou_by_class=dict(enumerate(iou_list)) if iou_list else None,
|
||||
min_box_fracs=mb_list,
|
||||
)
|
||||
if not res.error and res.detections:
|
||||
for det in res.detections:
|
||||
@@ -213,7 +260,14 @@ def _run_autolabel(job) -> None:
|
||||
elif res.error:
|
||||
job.log(f"[SAM3 ERROR] {frame['filename']}: {res.error}")
|
||||
|
||||
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_thresh)
|
||||
iou_by_class_proj = {
|
||||
c["class_id"]: per_class[c["name"].strip().lower()]["iou_threshold"]
|
||||
for c in project["classes"]
|
||||
if c["name"].strip().lower() in per_class
|
||||
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)
|
||||
items = []
|
||||
for det in kept:
|
||||
if project["label_type"] == "bbox" or det.mask is None:
|
||||
|
||||
+33
-11
@@ -10,7 +10,7 @@ calls — see the domain invariants in `../AGENTS.md`.
|
||||
"""
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional
|
||||
from typing import Dict, List, Optional
|
||||
|
||||
from PIL import Image
|
||||
|
||||
@@ -41,9 +41,12 @@ 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) -> List[Detection]:
|
||||
"""Greedy NMS per class: highest score wins within the SAME class."""
|
||||
if iou_threshold <= 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.
|
||||
|
||||
`iou_by_class` overrides the threshold per class id (REQ-181)."""
|
||||
if not iou_by_class and iou_threshold <= 0.0:
|
||||
return detections
|
||||
|
||||
by_class: dict[int, List[Detection]] = {}
|
||||
@@ -52,10 +55,14 @@ def deduplicate(detections: List[Detection], iou_threshold: float = 0.8) -> List
|
||||
by_class.setdefault(det.class_id, []).append(det)
|
||||
|
||||
kept: List[Detection] = []
|
||||
for cls_dets in by_class.values():
|
||||
for cls, cls_dets in by_class.items():
|
||||
iou = (iou_by_class or {}).get(cls, iou_threshold)
|
||||
if iou <= 0.0:
|
||||
kept.extend(cls_dets)
|
||||
continue
|
||||
cls_kept: List[Detection] = []
|
||||
for det in sorted(cls_dets, key=lambda d: d.score, reverse=True):
|
||||
if all(_iou(det.box, k.box) < iou_threshold for k in cls_kept):
|
||||
if all(_iou(det.box, k.box) < iou for k in cls_kept):
|
||||
cls_kept.append(det)
|
||||
kept.extend(cls_kept)
|
||||
return kept
|
||||
@@ -70,11 +77,16 @@ def label_image(
|
||||
min_box_frac: float = 0.0,
|
||||
exemplar_index: int = -1,
|
||||
exemplars: Optional[List[dict]] = None,
|
||||
thresholds: Optional[List[float]] = None,
|
||||
iou_by_class: Optional[Dict[int, float]] = None,
|
||||
min_box_fracs: Optional[List[float]] = 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."""
|
||||
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."""
|
||||
try:
|
||||
image = Image.open(image_path).convert("RGB")
|
||||
except Exception as exc: # unreadable/corrupt frame: report, don't abort the job
|
||||
@@ -84,14 +96,24 @@ def label_image(
|
||||
try:
|
||||
if exemplars and 0 <= exemplar_index < len(prompts):
|
||||
detections = get_engine().detect_with_exemplars(
|
||||
image, prompts, threshold, exemplar_index, exemplars
|
||||
image, prompts, threshold, exemplar_index, exemplars,
|
||||
thresholds=thresholds
|
||||
)
|
||||
else:
|
||||
detections = get_engine().detect(image, prompts, threshold)
|
||||
detections = get_engine().detect(image, prompts, threshold,
|
||||
thresholds=thresholds)
|
||||
except Exception as exc:
|
||||
return ImageResult(image_path, rel_path, width, height, error=str(exc))
|
||||
|
||||
if min_box_frac > 0:
|
||||
if min_box_fracs is not None:
|
||||
def _keep(det) -> bool:
|
||||
frac = min_box_fracs[det.class_id] if det.class_id < len(min_box_fracs) else min_box_frac
|
||||
if frac <= 0:
|
||||
return True
|
||||
floor = width * height * frac
|
||||
return (det.box[2] - det.box[0]) * (det.box[3] - det.box[1]) >= floor
|
||||
detections = [d for d in detections if _keep(d)]
|
||||
elif min_box_frac > 0:
|
||||
floor = width * height * min_box_frac
|
||||
detections = [
|
||||
d for d in detections
|
||||
@@ -99,4 +121,4 @@ def label_image(
|
||||
]
|
||||
|
||||
return ImageResult(image_path, rel_path, width, height,
|
||||
deduplicate(detections, iou_threshold))
|
||||
deduplicate(detections, iou_threshold, iou_by_class=iou_by_class))
|
||||
+35
-4
@@ -12,7 +12,7 @@ import os
|
||||
from typing import List, Optional
|
||||
|
||||
from backend import batches, db, labeling, projects, review
|
||||
from backend.autolabel import DEFAULT_IOU, DEFAULT_THRESHOLD, _geometries
|
||||
from backend.autolabel import DEFAULT_IOU, DEFAULT_THRESHOLD, _geometries, _parse_class_params
|
||||
|
||||
def preview_frame(
|
||||
batch_id: int,
|
||||
@@ -25,6 +25,7 @@ def preview_frame(
|
||||
custom_model_path: Optional[str] = None,
|
||||
exemplars: Optional[List[dict]] = None,
|
||||
exemplar_class_name: Optional[str] = None,
|
||||
class_params: Optional[dict] = None,
|
||||
) -> List[dict]:
|
||||
batch = batches.get(batch_id)
|
||||
if not batch:
|
||||
@@ -69,11 +70,13 @@ def preview_frame(
|
||||
|
||||
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
|
||||
per_class = _parse_class_params(class_params)
|
||||
predict_conf = min([threshold] + [v["threshold"] for v in per_class.values() if "threshold" in v])
|
||||
|
||||
all_raw_detections = []
|
||||
|
||||
if yolo_model is not None:
|
||||
results = yolo_model.predict(frame_file, conf=threshold, verbose=False)
|
||||
results = yolo_model.predict(frame_file, conf=predict_conf, verbose=False)
|
||||
if results and len(results) > 0:
|
||||
model_names = results[0].names
|
||||
for box in results[0].boxes:
|
||||
@@ -103,6 +106,15 @@ def preview_frame(
|
||||
|
||||
score = float(box.conf[0].item())
|
||||
xyxyn = box.xyxyn[0].tolist()
|
||||
|
||||
cp = per_class.get(proj_cls_name) or per_class.get(raw_cls_name)
|
||||
if cp:
|
||||
if score < cp.get("threshold", threshold):
|
||||
continue
|
||||
mb = cp.get("min_box_frac")
|
||||
if mb and mb > 0 and (xyxyn[2]-xyxyn[0])*(xyxyn[3]-xyxyn[1]) < mb:
|
||||
continue
|
||||
|
||||
all_raw_detections.append(labeling.Detection(
|
||||
class_id=target_class_id,
|
||||
class_name=proj_cls_name or raw_cls_name,
|
||||
@@ -125,10 +137,22 @@ def preview_frame(
|
||||
if c["name"].strip().lower() == wanted),
|
||||
-1,
|
||||
)
|
||||
thr_list = iou_list = mb_list = None
|
||||
if per_class:
|
||||
names = [c["name"].strip().lower() for c in sam3_target_classes]
|
||||
if any("threshold" in v for v in per_class.values()):
|
||||
thr_list = [per_class.get(n, {}).get("threshold", threshold) for n in names]
|
||||
if any("iou_threshold" in v for v in per_class.values()):
|
||||
iou_list = [per_class.get(n, {}).get("iou_threshold", iou_threshold) for n in names]
|
||||
if any("min_box_frac" in v for v in per_class.values()):
|
||||
mb_list = [per_class.get(n, {}).get("min_box_frac", min_box_frac) for n in names]
|
||||
res = labeling.label_image(
|
||||
frame_file, frame["filename"], prompts, threshold,
|
||||
iou_threshold=iou_threshold, min_box_frac=min_box_frac,
|
||||
exemplar_index=exemplar_index, exemplars=exemplars
|
||||
exemplar_index=exemplar_index, exemplars=exemplars,
|
||||
thresholds=thr_list,
|
||||
iou_by_class=dict(enumerate(iou_list)) if iou_list else None,
|
||||
min_box_fracs=mb_list,
|
||||
)
|
||||
if not res.error and res.detections:
|
||||
for det in res.detections:
|
||||
@@ -138,7 +162,14 @@ def preview_frame(
|
||||
det.class_name = real_cls["name"]
|
||||
all_raw_detections.append(det)
|
||||
|
||||
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_threshold)
|
||||
iou_by_class_proj = {
|
||||
c["class_id"]: per_class[c["name"].strip().lower()]["iou_threshold"]
|
||||
for c in project["classes"]
|
||||
if c["name"].strip().lower() in per_class
|
||||
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)
|
||||
items = []
|
||||
for det in kept:
|
||||
if project["label_type"] == "bbox" or det.mask is None:
|
||||
|
||||
+11
-2
@@ -67,8 +67,12 @@ class Sam3Engine:
|
||||
)
|
||||
self.processor = Sam3Processor(self.model, device=self.device)
|
||||
|
||||
def detect(self, image: Image.Image, prompts: List[str], threshold: float) -> List[Detection]:
|
||||
"""Run every prompt against one image; prompt index becomes the class id."""
|
||||
def detect(self, image: Image.Image, prompts: List[str], threshold: float,
|
||||
thresholds: Optional[List[float]] = None) -> List[Detection]:
|
||||
"""Run every prompt against one image; prompt index becomes the class id.
|
||||
|
||||
`thresholds` overrides the confidence per prompt (REQ-181); still one
|
||||
`set_image` for the whole call — only the grounding head sees the change."""
|
||||
processor = Sam3Processor(self.model, device=self.device)
|
||||
processor.confidence_threshold = threshold
|
||||
|
||||
@@ -76,6 +80,8 @@ class Sam3Engine:
|
||||
with torch.autocast(self.device, dtype=self.autocast_dtype):
|
||||
state = processor.set_image(image)
|
||||
for class_id, prompt in enumerate(prompts):
|
||||
if thresholds is not None and class_id < len(thresholds):
|
||||
processor.confidence_threshold = thresholds[class_id]
|
||||
output = processor.set_text_prompt(prompt=prompt, state=state)
|
||||
masks, boxes, scores = output["masks"], output["boxes"], output["scores"]
|
||||
if masks.shape[0] == 0:
|
||||
@@ -110,6 +116,7 @@ class Sam3Engine:
|
||||
threshold: float,
|
||||
exemplar_index: int,
|
||||
exemplars: List[dict],
|
||||
thresholds: Optional[List[float]] = None,
|
||||
) -> List[Detection]:
|
||||
"""`detect()`, but one prompt also carries drawn box exemplars (REQ-172).
|
||||
|
||||
@@ -124,6 +131,8 @@ class Sam3Engine:
|
||||
with torch.autocast(self.device, dtype=self.autocast_dtype):
|
||||
state = processor.set_image(image)
|
||||
for class_id, prompt in enumerate(prompts):
|
||||
if thresholds is not None and class_id < len(thresholds):
|
||||
processor.confidence_threshold = thresholds[class_id]
|
||||
processor.reset_all_prompts(state)
|
||||
output = processor.set_text_prompt(prompt=prompt, state=state)
|
||||
if class_id == exemplar_index:
|
||||
|
||||
Reference in new issue
Block a user