"""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)]