update from asus 106
This commit is contained in:
1 parent
6637fb1302
commit
8285400254
28 files changed
+3215
-459
No files matched your search
+121
-6
@@ -1,13 +1,14 @@
|
||||
"""Batch, frame and auto-annotation routes (REQ-020…034)."""
|
||||
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import tempfile
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Response
|
||||
from fastapi import APIRouter, File, Form, HTTPException, Response, UploadFile
|
||||
from fastapi.responses import FileResponse
|
||||
from pydantic import BaseModel
|
||||
|
||||
from backend import autolabel, dataset, library
|
||||
from backend import autolabel, dataset, library, projects
|
||||
from backend import batches as batch_store
|
||||
from backend import review as review_store
|
||||
from backend.api.common import project_or_404, thumbnail
|
||||
@@ -27,10 +28,12 @@ class AutolabelRequest(BaseModel):
|
||||
engines: Optional[list[str]] = None
|
||||
class_ids: Optional[list[int]] = None
|
||||
engine_classes: Optional[dict[str, list[str]]] = None
|
||||
target_class_names: Optional[list[str]] = None
|
||||
threshold: float = autolabel.DEFAULT_THRESHOLD
|
||||
iou_threshold: float = autolabel.DEFAULT_IOU
|
||||
min_box_frac: float = 0.0
|
||||
resume: bool = False
|
||||
append: bool = False
|
||||
|
||||
|
||||
@router.post("/api/projects/{project_id}/batches")
|
||||
@@ -90,11 +93,115 @@ def start_autolabel(batch_id: int, request: AutolabelRequest) -> dict:
|
||||
try:
|
||||
engine_list = request.engines if (request.engines and len(request.engines) > 0) else [request.engine]
|
||||
return autolabel.start(batch_id, request.threshold, request.iou_threshold,
|
||||
request.min_box_frac, resume=request.resume,
|
||||
request.min_box_frac, resume=request.resume, append=request.append,
|
||||
engines=engine_list, class_ids=request.class_ids,
|
||||
engine_classes=request.engine_classes)
|
||||
engine_classes=request.engine_classes,
|
||||
target_class_names=request.target_class_names)
|
||||
except batch_store.BatchError as exc:
|
||||
raise HTTPException(400, str(exc))
|
||||
@router.post("/api/batches/inspect-model")
|
||||
async def inspect_model(file: UploadFile = File(...)) -> dict:
|
||||
if not (file.filename or "").endswith(".pt"):
|
||||
raise HTTPException(400, "Model must be a .pt file")
|
||||
with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
|
||||
shutil.copyfileobj(file.file, staged)
|
||||
staged_path = staged.name
|
||||
await file.close()
|
||||
try:
|
||||
classes = projects.read_model_classes(staged_path)
|
||||
return {"filename": file.filename, "classes": classes, "staged_path": staged_path}
|
||||
except Exception as exc:
|
||||
if os.path.exists(staged_path):
|
||||
os.unlink(staged_path)
|
||||
raise HTTPException(400, f"Could not inspect model: {exc}")
|
||||
|
||||
|
||||
@router.post("/api/batches/{batch_id}/autolabel-with-model")
|
||||
async def autolabel_with_model(
|
||||
batch_id: int,
|
||||
file: UploadFile = File(...),
|
||||
threshold: float = Form(0.35),
|
||||
iou_threshold: float = Form(0.8),
|
||||
selected_classes: str = Form("[]"),
|
||||
append: bool = Form(True),
|
||||
) -> dict:
|
||||
if not (file.filename or "").endswith(".pt"):
|
||||
raise HTTPException(400, "Model must be a .pt file")
|
||||
with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
|
||||
shutil.copyfileobj(file.file, staged)
|
||||
staged_path = staged.name
|
||||
await file.close()
|
||||
|
||||
try:
|
||||
target_classes = json.loads(selected_classes) if selected_classes else None
|
||||
return autolabel.start(
|
||||
batch_id,
|
||||
threshold=threshold,
|
||||
iou_threshold=iou_threshold,
|
||||
append=append,
|
||||
custom_model_path=staged_path,
|
||||
target_class_names=target_classes,
|
||||
)
|
||||
except Exception as exc:
|
||||
if os.path.exists(staged_path):
|
||||
os.unlink(staged_path)
|
||||
raise HTTPException(400, f"Auto-annotation failed to start: {exc}")
|
||||
|
||||
|
||||
@router.post("/api/sam3/playground-test")
|
||||
async def sam3_playground_test(
|
||||
file: UploadFile = File(...),
|
||||
prompts: str = Form(...),
|
||||
threshold: float = Form(0.35),
|
||||
iou_threshold: float = Form(0.8),
|
||||
) -> dict:
|
||||
from PIL import Image
|
||||
from backend import labeling
|
||||
from backend.sam3_engine import get_engine
|
||||
|
||||
try:
|
||||
image = Image.open(file.file).convert("RGB")
|
||||
except Exception as exc:
|
||||
raise HTTPException(400, f"Could not read image: {exc}")
|
||||
|
||||
width, height = image.size
|
||||
prompt_list = [p.strip() for p in prompts.split(",") if p.strip()]
|
||||
if not prompt_list:
|
||||
raise HTTPException(400, "At least one text prompt is required")
|
||||
|
||||
try:
|
||||
engine = get_engine()
|
||||
raw_dets = engine.detect(image, prompt_list, threshold)
|
||||
kept_dets = labeling.deduplicate(raw_dets, iou_threshold=iou_threshold)
|
||||
except Exception as exc:
|
||||
raise HTTPException(500, f"SAM3 inference failed: {exc}")
|
||||
|
||||
results = []
|
||||
for det in kept_dets:
|
||||
norm_box = [
|
||||
det.box[0] / width,
|
||||
det.box[1] / height,
|
||||
det.box[2] / width,
|
||||
det.box[3] / height,
|
||||
]
|
||||
polys = []
|
||||
if det.mask is not None:
|
||||
raw_polys = review_store.mask_to_polygons(det.mask)
|
||||
polys = [[[float(pt[0]), float(pt[1])] for pt in poly] for poly in raw_polys]
|
||||
|
||||
results.append({
|
||||
"class_id": det.class_id,
|
||||
"prompt": det.class_name,
|
||||
"score": round(float(det.score), 4),
|
||||
"box": [round(v, 5) for v in norm_box],
|
||||
"polygons": polys,
|
||||
})
|
||||
|
||||
return {
|
||||
"width": width,
|
||||
"height": height,
|
||||
"detections": results,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -150,3 +257,11 @@ def clear_batch_class_annotations(batch_id: int, class_id: int) -> dict:
|
||||
deleted = review_store.clear_batch_class_annotations(batch_id, class_id)
|
||||
return {"deleted": deleted}
|
||||
|
||||
|
||||
@router.post("/api/batches/{batch_id}/reset-auto-annotations")
|
||||
def reset_batch_auto_annotations(batch_id: int) -> dict:
|
||||
if batch_store.get(batch_id) is None:
|
||||
raise HTTPException(404, "No such batch")
|
||||
deleted = review_store.clear_batch_auto_annotations(batch_id)
|
||||
return {"deleted": deleted}
|
||||
|
||||
@@ -19,6 +19,7 @@ class TrainRequest(BaseModel):
|
||||
imgsz: Optional[int] = None
|
||||
device: Optional[Union[int, str]] = None
|
||||
batch_ids: Optional[list] = None
|
||||
class_ids: Optional[list] = None
|
||||
|
||||
|
||||
@router.get("/api/hardware")
|
||||
@@ -34,6 +35,7 @@ def start_training(project_id: int, request: TrainRequest) -> dict:
|
||||
project_id, request.epochs,
|
||||
{"batch": request.batch, "imgsz": request.imgsz, "device": request.device},
|
||||
batch_ids=request.batch_ids,
|
||||
class_ids=request.class_ids,
|
||||
)
|
||||
except training.TrainingError as exc:
|
||||
raise HTTPException(400, str(exc))
|
||||
|
||||
+64
-69
@@ -18,10 +18,12 @@ DEFAULT_IOU = 0.8
|
||||
|
||||
def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
||||
iou_threshold: float = DEFAULT_IOU, min_box_frac: float = 0.0,
|
||||
resume: bool = False, engine: str = "sam3",
|
||||
resume: bool = False, append: bool = False, engine: str = "base_model",
|
||||
engines: Optional[List[str]] = None,
|
||||
class_ids: Optional[List[int]] = None,
|
||||
engine_classes: Optional[dict[str, List[str]]] = None) -> dict:
|
||||
engine_classes: Optional[dict[str, List[str]]] = None,
|
||||
custom_model_path: Optional[str] = None,
|
||||
target_class_names: Optional[List[str]] = None) -> dict:
|
||||
batch = batches.get(batch_id)
|
||||
if batch is None:
|
||||
raise batches.BatchError("No such batch")
|
||||
@@ -34,8 +36,9 @@ def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
||||
"autolabel",
|
||||
params={"batch_id": batch_id, "threshold": threshold,
|
||||
"iou_threshold": iou_threshold, "min_box_frac": min_box_frac,
|
||||
"resume": resume, "engine": active_engines[0], "engines": active_engines,
|
||||
"class_ids": class_ids, "engine_classes": engine_classes},
|
||||
"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},
|
||||
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)})",
|
||||
@@ -62,16 +65,7 @@ def _run_autolabel(job) -> None:
|
||||
raise batches.BatchError("The batch disappeared before labeling started")
|
||||
project = projects.get(batch["project_id"])
|
||||
|
||||
raw_active = job.params.get("engines") or [job.params.get("engine", "sam3")]
|
||||
expanded_engines = []
|
||||
for eng in raw_active:
|
||||
if eng == "both":
|
||||
expanded_engines.extend(["base_model", "secondary_model"])
|
||||
elif eng == "sam3+model1":
|
||||
expanded_engines.extend(["sam3", "base_model"])
|
||||
else:
|
||||
expanded_engines.append(eng)
|
||||
expanded_engines = list(dict.fromkeys(expanded_engines))
|
||||
selected_engine = job.params.get("engine", "base_model")
|
||||
|
||||
frames = batches.frames(batch["id"])
|
||||
batches.set_status(batch["id"], "labeling")
|
||||
@@ -86,42 +80,34 @@ def _run_autolabel(job) -> None:
|
||||
attempted = 0
|
||||
failures = []
|
||||
|
||||
from ultralytics import YOLO
|
||||
m1_path = projects.training_start_point(project)
|
||||
with db.cursor() as cur:
|
||||
cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
|
||||
row = cur.fetchone()
|
||||
if row and os.path.isfile(row[0]):
|
||||
m1_path = row[0]
|
||||
|
||||
yolo_models = {}
|
||||
if "base_model" in expanded_engines or "yolo" in expanded_engines:
|
||||
job.log(f"Loading Base/Trained Model: {os.path.basename(m1_path)}...")
|
||||
yolo_models["base_model"] = YOLO(m1_path)
|
||||
|
||||
if "secondary_model" in expanded_engines:
|
||||
m2_path = project["secondary_model_path"] if (project.get("secondary_model_path") and os.path.isfile(project["secondary_model_path"])) else m1_path
|
||||
label_name = project.get("secondary_model_name") or os.path.basename(m2_path)
|
||||
job.log(f"Loading Secondary Model: {label_name}...")
|
||||
yolo_models["secondary_model"] = YOLO(m2_path)
|
||||
|
||||
allowed_class_ids = set(job.params["class_ids"]) if job.params.get("class_ids") is not None else None
|
||||
engine_classes = job.params.get("engine_classes") or {}
|
||||
conf = job.params.get("threshold", DEFAULT_THRESHOLD)
|
||||
iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
|
||||
|
||||
yolo_model = None
|
||||
sam3_target_classes = []
|
||||
if "sam3" in expanded_engines:
|
||||
sam3_classes = engine_classes.get("sam3")
|
||||
if sam3_classes is not None:
|
||||
allowed_set = {c.strip().lower() for c in sam3_classes}
|
||||
sam3_target_classes = [
|
||||
c for c in project["classes"]
|
||||
if c["name"].strip().lower() in allowed_set or c["prompt"].strip().lower() in allowed_set
|
||||
]
|
||||
|
||||
custom_path = job.params.get("custom_model_path")
|
||||
target_class_names = job.params.get("target_class_names")
|
||||
|
||||
if selected_engine == "sam3" and not custom_path:
|
||||
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
||||
if allowed_classes_set:
|
||||
sam3_target_classes = [c for c in project["classes"] if c["name"].strip().lower() in allowed_classes_set or c["prompt"].strip().lower() in allowed_classes_set]
|
||||
# Add any new target class names that aren't in project classes yet
|
||||
existing_names = {c["name"].strip().lower() for c in project["classes"]}
|
||||
for name in target_class_names:
|
||||
if name.strip().lower() not in existing_names:
|
||||
try:
|
||||
updated_proj = projects.add_class(project["id"], name=name.strip(), prompt=name.strip())
|
||||
project["classes"] = updated_proj["classes"]
|
||||
for new_c in project["classes"]:
|
||||
if new_c["name"].strip().lower() == name.strip().lower() and new_c not in sam3_target_classes:
|
||||
sam3_target_classes.append(new_c)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
sam3_target_classes = [
|
||||
c for c in project["classes"]
|
||||
if (allowed_class_ids is None or c["class_id"] in allowed_class_ids)
|
||||
]
|
||||
sam3_target_classes = [c for c in project["classes"]]
|
||||
|
||||
prompts = [c["prompt"] for c in sam3_target_classes]
|
||||
if prompts:
|
||||
from backend.sam3_engine import engine_is_loaded, get_engine
|
||||
@@ -130,13 +116,27 @@ def _run_autolabel(job) -> None:
|
||||
engine = get_engine()
|
||||
job.log(f"SAM3 ready on {engine.device}; prompts: {', '.join(prompts)}")
|
||||
else:
|
||||
job.log("SAM3 selected but 0 prompts match class filter.")
|
||||
job.log("SAM3 selected but 0 prompts match project classes.")
|
||||
else:
|
||||
from ultralytics import YOLO
|
||||
if custom_path and os.path.isfile(custom_path):
|
||||
m_path = custom_path
|
||||
job.log(f"Loading Custom Model: {os.path.basename(m_path)}...")
|
||||
else:
|
||||
m_path = projects.training_start_point(project)
|
||||
with db.cursor() as cur:
|
||||
cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
|
||||
row = cur.fetchone()
|
||||
if row and os.path.isfile(row[0]):
|
||||
m_path = row[0]
|
||||
job.log(f"Loading Base Model: {os.path.basename(m_path)}...")
|
||||
|
||||
yolo_model = YOLO(m_path)
|
||||
|
||||
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
|
||||
conf = job.params.get("threshold", DEFAULT_THRESHOLD)
|
||||
iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
|
||||
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
||||
|
||||
job.log(f"Starting multi-engine auto-labeling ({', '.join(expanded_engines)})...")
|
||||
job.log(f"Starting auto-labeling with {selected_engine}...")
|
||||
|
||||
for index, frame in enumerate(frames):
|
||||
if job.cancelled:
|
||||
@@ -151,29 +151,21 @@ def _run_autolabel(job) -> None:
|
||||
frame_file = os.path.join(directory, frame["filename"])
|
||||
all_raw_detections = []
|
||||
|
||||
for eng_key, y_model in yolo_models.items():
|
||||
allowed_for_eng = engine_classes.get(eng_key)
|
||||
if allowed_for_eng is not None and len(allowed_for_eng) == 0:
|
||||
continue
|
||||
results = y_model.predict(frame_file, conf=conf, verbose=False)
|
||||
if yolo_model is not None:
|
||||
results = yolo_model.predict(frame_file, conf=conf, verbose=False)
|
||||
if results and len(results) > 0:
|
||||
model_names = results[0].names
|
||||
for box in results[0].boxes:
|
||||
cls_idx = int(box.cls[0].item())
|
||||
cls_name = str(model_names.get(cls_idx, cls_idx)).strip().lower()
|
||||
if allowed_for_eng is not None and cls_name not in [c.strip().lower() for c in allowed_for_eng]:
|
||||
|
||||
if allowed_classes_set is not None and cls_name not in allowed_classes_set:
|
||||
continue
|
||||
if cls_name not in name_to_class_id:
|
||||
try:
|
||||
updated_proj = projects.add_class(project["id"], {"name": cls_name, "prompt": cls_name})
|
||||
project["classes"] = updated_proj["classes"]
|
||||
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
|
||||
except Exception:
|
||||
pass
|
||||
if cls_name in name_to_class_id:
|
||||
target_class_id = name_to_class_id[cls_name]
|
||||
else:
|
||||
|
||||
target_class_id = name_to_class_id.get(cls_name)
|
||||
if target_class_id is None:
|
||||
continue
|
||||
|
||||
score = float(box.conf[0].item())
|
||||
xyxyn = box.xyxyn[0].tolist()
|
||||
all_raw_detections.append(labeling.Detection(
|
||||
@@ -184,7 +176,7 @@ def _run_autolabel(job) -> None:
|
||||
mask=None
|
||||
))
|
||||
|
||||
if "sam3" in expanded_engines and sam3_target_classes:
|
||||
elif selected_engine == "sam3" and sam3_target_classes:
|
||||
prompts = [c["prompt"] for c in sam3_target_classes]
|
||||
res = labeling.label_image(
|
||||
frame_file, frame["filename"], prompts, conf,
|
||||
@@ -208,7 +200,10 @@ def _run_autolabel(job) -> None:
|
||||
for geometry in _geometries(det, frame["width"], frame["height"], project["label_type"]):
|
||||
items.append({"class_id": det.class_id, "geometry": geometry, "score": det.score})
|
||||
|
||||
review.replace_auto(frame["id"], items)
|
||||
if job.params.get("append"):
|
||||
review.append_auto(frame["id"], items)
|
||||
else:
|
||||
review.replace_auto(frame["id"], items)
|
||||
written += len(items)
|
||||
job.progress(index + 1, len(frames), f"{frame['filename']}: {len(items)} shape(s)")
|
||||
except Exception as exc:
|
||||
|
||||
+46
-2
@@ -14,6 +14,7 @@ Label files are plain YOLO:
|
||||
import os
|
||||
import shutil
|
||||
import time
|
||||
from typing import List, Optional
|
||||
|
||||
from backend import batches, config, db, jobs, projects, review
|
||||
|
||||
@@ -74,12 +75,55 @@ def _next_split(cur, project_id: int, val_every: int) -> str:
|
||||
return "val" if position % val_every == val_every - 1 else "train"
|
||||
|
||||
|
||||
def write_data_yaml(project: dict, batch_ids: list = None) -> str:
|
||||
def sync_labels(project_id: int, selected_class_ids: Optional[List[int]] = None) -> dict:
|
||||
"""Re-sync label files on disk for all merged frames in the project dataset."""
|
||||
project = projects.get(project_id)
|
||||
root = dataset_dir(project["slug"])
|
||||
with db.cursor() as cur:
|
||||
cur.execute(
|
||||
"SELECT d.frame_id, d.label_rel FROM dataset_items d WHERE d.project_id = ?",
|
||||
(project_id,),
|
||||
)
|
||||
items = cur.fetchall()
|
||||
|
||||
class_map = None
|
||||
if selected_class_ids is not None and len(selected_class_ids) > 0:
|
||||
class_map = {cid: idx for idx, cid in enumerate(sorted(selected_class_ids))}
|
||||
|
||||
synced_files = 0
|
||||
total_lines = 0
|
||||
for frame_id, label_rel in items:
|
||||
annotations = review.listing(frame_id)
|
||||
if class_map is not None:
|
||||
annotations = [a for a in annotations if a["class_id"] in class_map]
|
||||
|
||||
lines = []
|
||||
for item in annotations:
|
||||
mapped_cid = class_map[item["class_id"]] if class_map is not None else item["class_id"]
|
||||
lines.append(_label_line(mapped_cid, item["geometry"], project["label_type"]))
|
||||
|
||||
path = os.path.join(root, label_rel)
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
with open(path, "w", encoding="utf-8") as f:
|
||||
f.write("\n".join(lines) + ("\n" if lines else ""))
|
||||
synced_files += 1
|
||||
total_lines += len(lines)
|
||||
|
||||
return {"synced_files": synced_files, "total_lines": total_lines}
|
||||
|
||||
|
||||
def write_data_yaml(project: dict, batch_ids: list = None, selected_class_ids: Optional[List[int]] = None) -> str:
|
||||
"""Rebuild data.yaml from the project's classes (REQ-051)."""
|
||||
sync_labels(project["id"], selected_class_ids=selected_class_ids)
|
||||
root = dataset_dir(project["slug"])
|
||||
os.makedirs(root, exist_ok=True)
|
||||
counts = summary(project["id"])["splits"]
|
||||
names = ", ".join(f"'{item['name']}'" for item in project["classes"])
|
||||
|
||||
target_classes = project["classes"]
|
||||
if selected_class_ids is not None and len(selected_class_ids) > 0:
|
||||
target_classes = [c for c in project["classes"] if c["class_id"] in selected_class_ids]
|
||||
|
||||
names = ", ".join(f"'{item['name']}'" for item in target_classes)
|
||||
|
||||
if batch_ids:
|
||||
with db.cursor() as cur:
|
||||
|
||||
+3
-3
@@ -44,9 +44,9 @@ def defaults(epochs: int = 50) -> dict:
|
||||
if info["device"] == "cpu":
|
||||
settings = {"batch": 4, "imgsz": 512, "device": "cpu", "workers": 2}
|
||||
note = "No GPU visible — training on CPU will be very slow."
|
||||
elif vram < 8:
|
||||
settings = {"batch": 8, "imgsz": 640, "device": 0, "workers": 2}
|
||||
note = f"{vram} GB of VRAM: small batches, 640 px."
|
||||
elif vram < 6:
|
||||
settings = {"batch": 16, "imgsz": 640, "device": 0, "workers": 4}
|
||||
note = f"{vram} GB of VRAM: batch 16, 640 px."
|
||||
elif vram <= 16:
|
||||
settings = {"batch": 32, "imgsz": 640, "device": 0, "workers": 8}
|
||||
note = f"{vram} GB of VRAM: optimized batch 32, 640 px."
|
||||
|
||||
+11
-4
@@ -42,11 +42,18 @@ def _iou(box_a: List[float], box_b: List[float]) -> float:
|
||||
|
||||
|
||||
def deduplicate(detections: List[Detection], iou_threshold: float = 0.8) -> List[Detection]:
|
||||
"""Greedy NMS across all prompts: highest score wins an overlapping region."""
|
||||
"""Greedy NMS per class: highest score wins within the SAME class."""
|
||||
by_class: dict[int, List[Detection]] = {}
|
||||
for det in detections:
|
||||
by_class.setdefault(det.class_id, []).append(det)
|
||||
|
||||
kept: List[Detection] = []
|
||||
for det in sorted(detections, key=lambda d: d.score, reverse=True):
|
||||
if all(_iou(det.box, k.box) < iou_threshold for k in kept):
|
||||
kept.append(det)
|
||||
for cls_dets in by_class.values():
|
||||
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):
|
||||
cls_kept.append(det)
|
||||
kept.extend(cls_kept)
|
||||
return kept
|
||||
|
||||
|
||||
|
||||
@@ -384,6 +384,11 @@ def _kept_prompts(project: dict, names: List[str]) -> List[str]:
|
||||
def training_start_point(project: dict) -> str:
|
||||
"""The weights a training run should start from (REQ-060, REQ-004)."""
|
||||
path = project["base_model_path"]
|
||||
if path and not os.path.isfile(path):
|
||||
if path.startswith("/data/"):
|
||||
alt_path = os.path.join(config.DATA_DIR, path[6:])
|
||||
if os.path.isfile(alt_path):
|
||||
path = alt_path
|
||||
if path and os.path.isfile(path):
|
||||
return path
|
||||
return PRETRAINED[project["label_type"]]
|
||||
|
||||
@@ -234,6 +234,38 @@ def replace_auto(frame_id: int, items: List[dict]) -> int:
|
||||
return len(items)
|
||||
|
||||
|
||||
def append_auto(frame_id: int, items: List[dict]) -> int:
|
||||
"""Add new automatic shapes to this frame without duplicating existing ones."""
|
||||
if not items:
|
||||
return 0
|
||||
existing = listing(frame_id)
|
||||
filtered_items = []
|
||||
for item in items:
|
||||
is_dup = False
|
||||
item_box = to_box(item["geometry"])
|
||||
for ex in existing:
|
||||
if ex["class_id"] == item["class_id"]:
|
||||
ex_box = to_box(ex["geometry"])
|
||||
from backend.labeling import _iou
|
||||
if _iou(item_box, ex_box) >= 0.85:
|
||||
is_dup = True
|
||||
break
|
||||
if not is_dup:
|
||||
filtered_items.append(item)
|
||||
|
||||
if not filtered_items:
|
||||
return 0
|
||||
|
||||
with db.cursor() as cur:
|
||||
cur.executemany(
|
||||
"""INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
|
||||
VALUES (?, ?, ?, ?, 'auto', ?)""",
|
||||
[(frame_id, item["class_id"], json.dumps(item["geometry"]),
|
||||
item.get("score", 1.0), time.time()) for item in filtered_items],
|
||||
)
|
||||
return len(filtered_items)
|
||||
|
||||
|
||||
def frames_with_auto(batch_id: int) -> set:
|
||||
"""Frame ids that already carry automatic shapes — the resume skip-list
|
||||
for REQ-035."""
|
||||
@@ -339,6 +371,24 @@ def clear_batch_class_annotations(batch_id: int, class_id: int) -> int:
|
||||
return cur.rowcount
|
||||
|
||||
|
||||
def clear_batch_auto_annotations(batch_id: int) -> int:
|
||||
"""Delete all automatic annotations (source = 'auto') for a batch and reset frame statuses."""
|
||||
with db.cursor() as cur:
|
||||
cur.execute(
|
||||
"""DELETE FROM annotations
|
||||
WHERE source = 'auto' AND frame_id IN (
|
||||
SELECT id FROM frames WHERE batch_id = ?
|
||||
)""",
|
||||
(batch_id,),
|
||||
)
|
||||
deleted = cur.rowcount
|
||||
cur.execute(
|
||||
"UPDATE frames SET review_status = 'pending' WHERE batch_id = ?",
|
||||
(batch_id,),
|
||||
)
|
||||
return deleted
|
||||
|
||||
|
||||
def _check_class(project_id: int, class_id: int) -> None:
|
||||
with db.cursor() as cur:
|
||||
cur.execute(
|
||||
|
||||
+9
-4
@@ -25,7 +25,7 @@ def models_dir(project_slug: str) -> str:
|
||||
return os.path.join(config.project_dir(project_slug), "models")
|
||||
|
||||
|
||||
def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, batch_ids: Optional[list] = None) -> dict:
|
||||
def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, batch_ids: Optional[list] = None, class_ids: Optional[list] = None) -> dict:
|
||||
project = projects.get(project_id)
|
||||
if project is None:
|
||||
raise TrainingError("No such project")
|
||||
@@ -38,7 +38,7 @@ def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, b
|
||||
settings = hardware.resolve(overrides, epochs)
|
||||
job = jobs.create(
|
||||
"train",
|
||||
params={"project_id": project_id, "settings": settings, "batch_ids": batch_ids},
|
||||
params={"project_id": project_id, "settings": settings, "batch_ids": batch_ids, "class_ids": class_ids},
|
||||
project_id=project_id,
|
||||
message=f"{counts['train']} train / {counts['val']} val",
|
||||
)
|
||||
@@ -102,7 +102,8 @@ def _run_train(job) -> None:
|
||||
project = projects.get(job.params["project_id"])
|
||||
settings = job.params["settings"]
|
||||
batch_ids = job.params.get("batch_ids")
|
||||
data_yaml = dataset.write_data_yaml(project, batch_ids=batch_ids)
|
||||
class_ids = job.params.get("class_ids")
|
||||
data_yaml = dataset.write_data_yaml(project, batch_ids=batch_ids, selected_class_ids=class_ids)
|
||||
|
||||
# SAM3 and a training run must not hold VRAM at the same time (REQ-065).
|
||||
from backend.sam3_engine import release_engine
|
||||
@@ -133,6 +134,10 @@ def _run_train(job) -> None:
|
||||
model.add_callback("on_fit_epoch_end", on_epoch)
|
||||
job.progress(0, settings["epochs"])
|
||||
|
||||
import torch
|
||||
if torch.cuda.is_available():
|
||||
torch.backends.cudnn.benchmark = True
|
||||
|
||||
keep_run_dir = False
|
||||
try:
|
||||
model.train(
|
||||
@@ -142,7 +147,7 @@ def _run_train(job) -> None:
|
||||
batch=settings["batch"],
|
||||
device=settings["device"],
|
||||
workers=settings.get("workers", 8),
|
||||
cache=False,
|
||||
cache="ram",
|
||||
project=os.path.join(out_dir, "runs"),
|
||||
name="train",
|
||||
exist_ok=True,
|
||||
|
||||
Reference in new issue
Block a user