feat: setup dataset enrichment app codebase and scripts
This commit is contained in:
1 parent
b5c28cc98a
commit
d07578462e
72 files changed
+11370
No files matched your search
@@ -0,0 +1,269 @@
|
||||
"""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 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 write_data_yaml(project: dict, batch_ids: list = None) -> str:
|
||||
"""Rebuild data.yaml from the project's classes (REQ-051)."""
|
||||
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"])
|
||||
|
||||
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}")
|
||||
Reference in new issue
Block a user