Files
asus 5c7c122105 feat: add counting bench, triage, and dataset modules
This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
2026-08-14 16:28:52 +07:00

221 lines
8.2 KiB
Python

"""Base datasets: externally-labelled images that ride along with training.
A base dataset is not a batch and never becomes one. It has no frames, no
review, no triage — it is a folder of images and YOLO label files that the user
already trusts, registered against a project and offered as a checkbox next to
the project's own datasets (REQ-120…123).
It stays outside `dataset_items` on purpose. That table keys on `frame_id`, and
inventing frame rows for images this app never extracted would put fake batches
in front of the user for the rest of the project's life. Instead the rows are
handed straight to `dataset._build_selected_tree`, which only ever wanted an
image path and a label path.
"""
import json
import os
import shutil
import time
from typing import List, Optional
from backend import config, db
IMAGE_EXTS = (".jpg", ".jpeg", ".png", ".bmp", ".webp")
class BaseDatasetError(Exception):
pass
def root_dir(project_slug: str, base_id: int) -> str:
return os.path.join(config.project_dir(project_slug), "base_datasets", str(base_id))
def listing(project_id: int) -> List[dict]:
with db.cursor() as cur:
cur.execute(
"""SELECT id, project_id, name, source, image_count, box_count, classes, created_at
FROM base_datasets WHERE project_id = ? ORDER BY created_at DESC""",
(project_id,),
)
return [_row(r) for r in cur.fetchall()]
def get(base_id: int) -> Optional[dict]:
with db.cursor() as cur:
cur.execute(
"""SELECT id, project_id, name, source, image_count, box_count, classes, created_at
FROM base_datasets WHERE id = ?""",
(base_id,),
)
row = cur.fetchone()
return _row(row) if row else None
def _row(row) -> dict:
return {
"id": row[0], "project_id": row[1], "name": row[2], "source": row[3],
"image_count": row[4], "box_count": row[5],
"classes": json.loads(row[6] or "[]"), "created_at": row[7],
}
def delete(base_id: int, project_slug: str) -> bool:
with db.cursor() as cur:
cur.execute("DELETE FROM base_datasets WHERE id = ?", (base_id,))
removed = cur.rowcount > 0
if removed:
shutil.rmtree(root_dir(project_slug, base_id), ignore_errors=True)
return removed
def rows(project_id: int, base_ids: List[int], project_slug: str) -> List[dict]:
"""Image/label pairs for the run's `selected/` tree.
Everything is `train`. A base dataset must not contribute validation images:
the base-vs-new comparison is only meaningful measured on this project's own
val split, and REQ-052 keeps that split stable (REQ-122).
"""
if not base_ids:
return []
out = []
for base_id in base_ids:
record = get(base_id)
if record is None or record["project_id"] != project_id:
continue
root = root_dir(project_slug, base_id)
images_dir = os.path.join(root, "images")
labels_dir = os.path.join(root, "labels")
if not os.path.isdir(images_dir):
continue
for name in sorted(os.listdir(images_dir)):
if not name.lower().endswith(IMAGE_EXTS):
continue
out.append({
"frame_id": None,
"split": "train",
"source_image": os.path.join(images_dir, name),
"source_label": os.path.join(labels_dir, os.path.splitext(name)[0] + ".txt"),
})
return out
def import_tree(project_id: int, project_slug: str, source_dir: str, name: str,
keep_class_ids: List[int], on_progress=None) -> dict:
"""Adopt an unpacked YOLO export, keeping only `keep_class_ids`.
Class ids are kept as they are — the caller has already checked that the
export numbers its classes the same way the project does. A label line for a
class we are not keeping is dropped; an image left with no lines at all is
dropped with it, because an empty label is a claim that the image contains
none of the kept classes, and here it only means "the box was a truck".
"""
pairs = _collect(source_dir)
if not pairs:
raise BaseDatasetError(f"No image/label pairs found under {source_dir}")
keep = set(keep_class_ids)
with db.cursor() as cur:
cur.execute(
"""INSERT INTO base_datasets (project_id, name, source, classes, created_at)
VALUES (?, ?, ?, ?, ?)""",
(project_id, name, os.path.basename(source_dir.rstrip("/")),
json.dumps(sorted(keep)), time.time()),
)
base_id = cur.lastrowid
root = root_dir(project_slug, base_id)
images_dir = os.path.join(root, "images")
labels_dir = os.path.join(root, "labels")
os.makedirs(images_dir, exist_ok=True)
os.makedirs(labels_dir, exist_ok=True)
images = boxes = skipped = 0
for index, (image_path, label_path) in enumerate(pairs):
lines = []
if os.path.isfile(label_path):
with open(label_path, encoding="utf-8") as handle:
for line in handle.read().splitlines():
if not line.strip():
continue
parts = line.split()
try:
class_id = int(parts[0])
except (ValueError, IndexError):
continue
if class_id not in keep:
continue
normalised = _to_bbox(parts)
if normalised is not None:
lines.append(normalised)
if not lines:
skipped += 1
continue
stem = os.path.basename(image_path)
shutil.copyfile(image_path, os.path.join(images_dir, stem))
with open(os.path.join(labels_dir, os.path.splitext(stem)[0] + ".txt"),
"w", encoding="utf-8") as handle:
handle.write("\n".join(lines) + "\n")
images += 1
boxes += len(lines)
if on_progress is not None and index % 25 == 0:
on_progress(index + 1, len(pairs))
with db.cursor() as cur:
cur.execute("UPDATE base_datasets SET image_count = ?, box_count = ? WHERE id = ?",
(images, boxes, base_id))
return {**get(base_id), "skipped": skipped, "candidates": len(pairs)}
def _to_bbox(parts: List[str]) -> Optional[str]:
"""Normalise one YOLO label line to `class cx cy w h`.
Roboflow exports segmentation polygons when the source project was drawn
that way, and a detect model reads the first four numbers of such a line as
a box — which lands somewhere near the first two polygon vertices and is
nowhere near the object. Polygons are collapsed to their bounding box, which
is the honest projection of a mask onto a bbox dataset.
"""
class_id, coords = parts[0], parts[1:]
if len(coords) == 4:
return " ".join([class_id] + coords)
if len(coords) < 6 or len(coords) % 2 != 0:
return None
try:
values = [float(v) for v in coords]
except ValueError:
return None
xs, ys = values[0::2], values[1::2]
x0, x1 = min(xs), max(xs)
y0, y1 = min(ys), max(ys)
width, height = x1 - x0, y1 - y0
if width <= 0 or height <= 0:
return None
return (f"{class_id} {(x0 + x1) / 2:.6f} {(y0 + y1) / 2:.6f} "
f"{width:.6f} {height:.6f}")
def _collect(source_dir: str) -> List[tuple]:
"""Every image under the tree, paired with its sibling label file.
Handles both a flat `images/`+`labels/` pair and the split layout Roboflow
exports (`train/images`, `valid/labels`, …). Deduplicated by file name, so
re-importing an export that overlaps an earlier one cannot double-weight the
same picture.
"""
seen = {}
for current, _dirs, files in os.walk(source_dir):
if os.path.basename(current) != "images":
continue
labels = os.path.join(os.path.dirname(current), "labels")
for name in sorted(files):
if not name.lower().endswith(IMAGE_EXTS):
continue
if name in seen:
continue
seen[name] = (os.path.join(current, name),
os.path.join(labels, os.path.splitext(name)[0] + ".txt"))
return [seen[key] for key in sorted(seen)]