"""The master dataset: approved frames merged in, batch after batch (REQ-050…054). The one rule that matters here is the stable val split. A frame's membership is recorded once in `dataset_items` and never revised, so an image that was in `val` for the last comparison is still in `val` for the next one. Without that, a rising mAP could just mean an easier val set. Label files are plain YOLO: detect class_id cx cy w h (normalized) segment class_id x1 y1 x2 y2 … (normalized polygon) """ import os import shutil import time from typing import List, Optional from backend import batches, config, db, jobs, projects, review class DatasetError(Exception): pass def dataset_dir(project_slug: str) -> str: return os.path.join(config.project_dir(project_slug), "dataset") def approve(batch_id: int) -> dict: """Sign a batch off and queue its merge (REQ-045, REQ-050).""" batch = batches.get(batch_id) if batch is None: raise DatasetError("No such batch") if batch["status"] == "merged": raise DatasetError("This batch is already in the master dataset") if batch["review"]["pending"] > 0: raise DatasetError( f"{batch['review']['pending']} frame(s) still need a decision before this " "batch can be approved" ) if batch["review"]["approved"] == 0: raise DatasetError("Every frame was rejected — there is nothing to merge") batches.set_status(batch_id, "approved") job = jobs.create( "merge", params={"batch_id": batch_id}, project_id=batch["project_id"], batch_id=batch_id, message=f"{batch['date_label']}/{batch['batch_label']}", ) return job.to_dict() def _label_line(class_id: int, geometry: dict, label_type: str) -> str: if label_type == "bbox": x0, y0, x1, y1 = review.to_box(geometry) return (f"{class_id} {(x0 + x1) / 2:.6f} {(y0 + y1) / 2:.6f} " f"{x1 - x0:.6f} {y1 - y0:.6f}") points = geometry["points"] if geometry["type"] == "bbox": x0, y0, x1, y1 = geometry["points"] points = [[x0, y0], [x1, y0], [x1, y1], [x0, y1]] coords = " ".join(f"{value:.6f}" for point in points for value in point) return f"{class_id} {coords}" def _next_split(cur, project_id: int, val_every: int) -> str: """Continue the every-Nth pattern from wherever the last merge left off.""" if val_every <= 0: return "train" cur.execute("SELECT COUNT(*) FROM dataset_items WHERE project_id = ?", (project_id,)) position = cur.fetchone()[0] return "val" if position % val_every == val_every - 1 else "train" 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"] 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: placeholders = ",".join("?" for _ in batch_ids) cur.execute( f"""SELECT d.image_rel, d.split FROM dataset_items d JOIN frames f ON f.id = d.frame_id WHERE d.project_id = ? AND f.batch_id IN ({placeholders})""", [project["id"]] + list(batch_ids), ) rows = cur.fetchall() train_files = [row[0] for row in rows if row[1] == "train"] val_files = [row[0] for row in rows if row[1] == "val"] or train_files train_txt = os.path.join(root, "selected_train.txt") val_txt = os.path.join(root, "selected_val.txt") with open(train_txt, "w", encoding="utf-8") as handle: handle.write("\n".join(os.path.join(root, rel) for rel in train_files) + "\n") with open(val_txt, "w", encoding="utf-8") as handle: handle.write("\n".join(os.path.join(root, rel) for rel in val_files) + "\n") path = os.path.join(root, "selected_data.yaml") with open(path, "w", encoding="utf-8") as handle: handle.write(f"path: {root}\n") handle.write(f"train: {train_txt}\n") handle.write(f"val: {val_txt}\n\n") handle.write(f"nc: {len(project['classes'])}\n") handle.write(f"names: [{names}]\n") return path path = os.path.join(root, "data.yaml") with open(path, "w", encoding="utf-8") as handle: handle.write(f"path: {root}\n") handle.write("train: images/train\n") handle.write(f"val: images/{'val' if counts['val'] > 0 else 'train'}\n\n") handle.write(f"nc: {len(project['classes'])}\n") handle.write(f"names: [{names}]\n") return path def summary(project_id: int) -> dict: with db.cursor() as cur: cur.execute( "SELECT split, COUNT(*) FROM dataset_items WHERE project_id = ? GROUP BY split", (project_id,), ) splits = {"train": 0, "val": 0} for split, count in cur.fetchall(): splits[split] = count cur.execute( """SELECT b.id, b.date_label, b.batch_label, b.merged_at, COUNT(d.id) AS images FROM batches b LEFT JOIN frames f ON f.batch_id = b.id LEFT JOIN dataset_items d ON d.frame_id = f.id WHERE b.project_id = ? AND b.status = 'merged' GROUP BY b.id ORDER BY b.merged_at""", (project_id,), ) merged = [dict(row) for row in cur.fetchall()] return {"splits": splits, "total": splits["train"] + splits["val"], "batches": merged} def drop_class_from_labels(project: dict, class_id: int) -> dict: """Rewrite every label file on disk after a class is deleted (REQ-007). Two edits per file: lines of the deleted class go, and every id above it comes down by one. Skipping this would leave `2` in old files meaning a class that is now `1` — labels that quietly name the wrong thing are worse than labels that are missing. """ root = dataset_dir(project["slug"]) with db.cursor() as cur: cur.execute("SELECT label_rel FROM dataset_items WHERE project_id = ?", (project["id"],)) label_files = [row[0] for row in cur.fetchall()] rewritten = 0 dropped = 0 for rel in label_files: path = os.path.join(root, rel) if not os.path.isfile(path): continue with open(path, encoding="utf-8") as handle: lines = handle.read().splitlines() kept, touched = [], False for line in lines: if not line.strip(): continue head, _, rest = line.partition(" ") try: current = int(head) except ValueError: kept.append(line) continue if current == class_id: dropped += 1 touched = True continue if current > class_id: current -= 1 touched = True kept.append(f"{current} {rest}") if touched: # An emptied file stays as an empty file: the image is still a valid # negative sample (REQ-033), it just has nothing on it any more. with open(path, "w", encoding="utf-8") as handle: handle.write("\n".join(kept) + ("\n" if kept else "")) rewritten += 1 return {"label_files_rewritten": rewritten, "dataset_lines_removed": dropped} def zip_path(project: dict) -> str: """Zip the master dataset for download (REQ-054).""" root = dataset_dir(project["slug"]) if not os.path.isdir(os.path.join(root, "images")): raise DatasetError("This project's dataset is still empty") archive = os.path.join(config.project_dir(project["slug"]), "dataset") return shutil.make_archive(archive, "zip", root) @jobs.handler("merge") def _run_merge(job) -> None: batch = batches.get(job.params["batch_id"]) if batch is None: raise DatasetError("The batch disappeared before the merge started") project = projects.get(batch["project_id"]) root = dataset_dir(project["slug"]) for split in ("train", "val"): os.makedirs(os.path.join(root, "images", split), exist_ok=True) os.makedirs(os.path.join(root, "labels", split), exist_ok=True) frames = [f for f in batches.frames(batch["id"]) if f["review_status"] == "approved"] source_dir = batches.frames_dir(project["slug"], batch["id"]) job.progress(0, len(frames)) job.log(f"Merging {len(frames)} approved frame(s) into the master dataset") added = {"train": 0, "val": 0} skipped = 0 for index, frame in enumerate(frames): if job.cancelled: job.log(f"Cancelled after {index} frame(s)") break with db.cursor() as cur: cur.execute("SELECT 1 FROM dataset_items WHERE frame_id = ?", (frame["id"],)) if cur.fetchone() is not None: skipped += 1 job.progress(index + 1, len(frames)) continue split = _next_split(cur, project["id"], project["val_every"]) stem = f"{batch['id']}__{os.path.splitext(frame['filename'])[0]}" image_rel = f"images/{split}/{stem}.jpg" label_rel = f"labels/{split}/{stem}.txt" shutil.copyfile(os.path.join(source_dir, frame["filename"]), os.path.join(root, image_rel)) lines = [_label_line(item["class_id"], item["geometry"], project["label_type"]) for item in review.listing(frame["id"])] # An approved frame with nothing on it is a negative sample, and an # empty .txt is how YOLO spells that (REQ-033). with open(os.path.join(root, label_rel), "w", encoding="utf-8") as handle: handle.write("\n".join(lines) + ("\n" if lines else "")) cur.execute( """INSERT INTO dataset_items (project_id, frame_id, split, image_rel, label_rel, added_at) VALUES (?, ?, ?, ?, ?, ?)""", (project["id"], frame["id"], split, image_rel, label_rel, time.time()), ) added[split] += 1 job.progress(index + 1, len(frames)) with db.cursor() as cur: cur.execute("UPDATE batches SET status = 'merged', merged_at = ? WHERE id = ?", (time.time(), batch["id"])) path = write_data_yaml(projects.get(project["id"])) totals = summary(project["id"])["splits"] job.log(f"Added {added['train']} train / {added['val']} val" + (f", skipped {skipped} already merged" if skipped else "")) job.log(f"Master dataset now {totals['train']} train / {totals['val']} val — {path}")