314 lines
12 KiB
Python
314 lines
12 KiB
Python
"""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}")
|