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
-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:
|
||||
|
||||
Reference in new issue
Block a user