alami-vision-api / ml /scripts /prepare_taco.py
alami-ci
Deploy from alami-eco/alami-trash-ai@aee69796b70947e95efdb9c7483fa52f8d3b4520
76838d6
Raw
History Blame Contribute Delete
33.5 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
prepare_taco.py
Functionality:
- Mode A (default): Local TACO images + a single COCO annotations file
-> Map classes to Alami super-classes
-> Stratified split (train/val/test)
-> Write YOLOv8-seg labels + dataset.yaml + stats
- Mode B (pre-split JSONs already exist): --annotations_train/val/test
-> No stratification; use provided splits as-is
-> Write YOLOv8-seg labels + dataset.yaml + stats
- Optional in both modes:
1) If images are missing locally, download them directly from COCO image fields
(flickr_640_url, flickr_url, coco_url, url) and save under image['file_name'] in --images_dir.
2) If still missing or explicitly desired: CSV fallback (--images_urls_csv),
matching by basename (URL pathname) to image['file_name'].
Notes:
- Dedupe via MD5 available (default on).
- RLE masks are skipped (polygons only).
NEW (backwards compatible, default keeps behavior):
- --min_poly_area_px: discard polygons smaller than this area (pixel^2). default 0 = off.
"""
import argparse
import json
import os
import sys
import shutil
import random
from collections import defaultdict, Counter
from pathlib import Path
from typing import Dict, List, Tuple, Any, Optional
import re
# ---- Required deps ----
try:
from PIL import Image, ImageOps
except ImportError:
print("Please `pip install pillow`", file=sys.stderr); raise
try:
import numpy as np
except ImportError:
print("Please `pip install numpy`", file=sys.stderr); raise
# ---------------------------
# Helpers: IO / JSON / FS
# ---------------------------
def read_json(path: Path) -> Any:
with path.open("r", encoding="utf-8") as f:
return json.load(f)
def write_json(path: Path, obj: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as f:
json.dump(obj, f, ensure_ascii=False, indent=2)
def write_text(path: Path, s: str) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
with path.open("w", encoding="utf-8") as f:
f.write(s)
def safe_symlink_or_copy(src: Path, dst: Path) -> None:
dst.parent.mkdir(parents=True, exist_ok=True)
if os.name == "nt": # Windows: Symlink pain -> copy
if not dst.exists():
shutil.copy2(src, dst)
else:
shutil.copy2(src, dst)
return
try:
if dst.exists() or dst.is_symlink():
dst.unlink()
os.symlink(src, dst)
except Exception:
if not dst.exists():
shutil.copy2(src, dst)
else:
shutil.copy2(src, dst)
# ---------------------------
# Label map (rule-based)
# ---------------------------
def load_label_map(label_map_path: Path) -> Dict[str, Any]:
lm = read_json(label_map_path)
required = ["target_classes", "default", "rules"]
for k in required:
if k not in lm:
raise ValueError(f"label_map.json missing key: {k}")
# normalize for case-insensitive matching
rules = []
for r in lm["rules"]:
nr = {
"to": r["to"],
"match_any_substring": [s.lower() for s in r.get("match_any_substring", [])],
"match_any_regex": r.get("match_any_regex", []),
}
rules.append(nr)
lm["rules"] = rules
return lm
def map_class(name: str, label_map: Dict[str, Any]) -> str:
n = name.lower().strip()
for r in label_map["rules"]:
for sub in r["match_any_substring"]:
if sub in n:
return r["to"]
for pattern in r["match_any_regex"]:
if re.search(pattern, n):
return r["to"]
return label_map["default"]
# ---------------------------
# COCO / TACO reading
# ---------------------------
def index_coco(coco: Dict[str, Any]) -> Tuple[Dict[int, dict], Dict[int, dict], Dict[int, List[dict]]]:
images_by_id = {im["id"]: im for im in coco.get("images", [])}
cats_by_id = {c["id"]: c for c in coco.get("categories", [])}
anns_by_image = defaultdict(list)
for ann in coco.get("annotations", []):
anns_by_image[ann["image_id"]].append(ann)
return images_by_id, cats_by_id, anns_by_image
# ---------------------------
# Geometry helpers (YOLOv8 seg expects normalized polygon points)
# ---------------------------
def coco_segmentation_to_yolo_polys(seg, img_w: int, img_h: int) -> List[List[float]]:
"""
COCO 'segmentation' can be list of polygons (each as flat list of x,y) or RLE.
We only support polygons here; RLE is skipped.
Returns list of [x1_norm,y1_norm, x2_norm,y2_norm, ...] per polygon.
"""
polys = []
if isinstance(seg, list):
for poly in seg:
if not isinstance(poly, list) or len(poly) < 6:
continue
norm = []
for i, v in enumerate(poly):
if i % 2 == 0:
x = max(0.0, min(float(v) / img_w, 1.0))
norm.append(x)
else:
y = max(0.0, min(float(v) / img_h, 1.0))
norm.append(y)
polys.append(norm)
# else: RLE not supported here
return polys
def _polygon_area_px(norm_poly: List[float], img_w: int, img_h: int) -> float:
"""
Shoelace area in pixels from normalized coordinates.
norm_poly: [x1n,y1n,x2n,y2n,...]
"""
if len(norm_poly) < 6:
return 0.0
xs = [norm_poly[i] * img_w for i in range(0, len(norm_poly), 2)]
ys = [norm_poly[i] * img_h for i in range(1, len(norm_poly), 2)]
area = 0.0
n = len(xs)
for i in range(n):
j = (i + 1) % n
area += xs[i] * ys[j] - xs[j] * ys[i]
return abs(area) * 0.5
# ---------------------------
# Hashing / Dedupe (optional)
# ---------------------------
def image_md5(path: Path) -> Optional[str]:
try:
import hashlib
h = hashlib.md5()
with path.open("rb") as f:
for chunk in iter(lambda: f.read(8192), b""):
h.update(chunk)
return h.hexdigest()
except Exception:
return None
def image_phash(path: Path) -> Optional[str]:
try:
import scipy.fftpack
except Exception:
return None
try:
img = Image.open(path).convert("L")
img = ImageOps.fit(img, (32, 32))
arr = np.asarray(img, dtype=np.float32)
dct = scipy.fftpack.dct(scipy.fftpack.dct(arr.T, norm="ortho").T, norm="ortho")
dct_low = dct[:8, :8]
med = np.median(dct_low)
bits = (dct_low > med).flatten()
return "".join("1" if b else "0" for b in bits)
except Exception:
return None
# ---------------------------
# Split (stratify by primary class)
# ---------------------------
def stratified_split(items: List[dict], y: List[str], train_ratio=0.8, val_ratio=0.1, seed=42):
rnd = random.Random(seed)
by_class = defaultdict(list)
for i, c in enumerate(y):
by_class[c].append(i)
train, val, test = [], [], []
for c, idx_list in by_class.items():
rnd.shuffle(idx_list)
n = len(idx_list)
n_train = int(round(n * train_ratio))
n_val = int(round(n * val_ratio))
n_test = n - n_train - n_val
train += idx_list[:n_train]
val += idx_list[n_train:n_train+n_val]
test += idx_list[n_train+n_val:]
for arr in (train, val, test):
rnd.shuffle(arr)
return train, val, test
# ---------------------------
# URL download support (CSV)
# ---------------------------
def read_urls_csv(csv_path: Path) -> List[str]:
urls = []
with csv_path.open("r", encoding="utf-8") as f:
for line in f:
u = line.strip()
if not u or u.lower().startswith("#"):
continue
urls.append(u)
return urls
def basename_from_url(u: str) -> str:
try:
from urllib.parse import urlparse, unquote
p = urlparse(u)
b = Path(unquote(p.path)).name
return b
except Exception:
return Path(u).name
def ensure_images_from_urls(
urls_csv: Path,
target_images_dir: Path,
coco_images_by_id: Dict[int, dict]
) -> None:
"""
Download images from URLs into target_images_dir if they match COCO file_name by basename.
- Heuristic: basename(URL) == basename(image['file_name'])
- Non-matching URLs are ignored.
- Existing files are not re-downloaded.
"""
urls = read_urls_csv(urls_csv)
if not urls:
print(f"[WARN] No URLs found in {urls_csv}", file=sys.stderr)
return
from collections import defaultdict
by_base = defaultdict(list)
for u in urls:
by_base[basename_from_url(u).lower()].append(u)
target_images_dir.mkdir(parents=True, exist_ok=True)
q: List[Tuple[str, Path]] = []
for im in coco_images_by_id.values():
base = Path(im["file_name"]).name.lower()
cand_urls = by_base.get(base)
out = target_images_dir / Path(im["file_name"]).name
if out.exists():
continue
if cand_urls:
q.append((cand_urls[0], out))
if not q:
print("[INFO] No downloads required (all images present or no matches).")
return
print(f"[INFO] Downloading {len(q)} images from CSV into {target_images_dir} ...")
import requests
from concurrent.futures import ThreadPoolExecutor, as_completed
def _fetch(u_dst: Tuple[str, Path]) -> Tuple[Path, bool, str]:
u, dst = u_dst
try:
last = ""
for attempt in range(3):
try:
r = requests.get(u, timeout=20)
r.raise_for_status()
dst.parent.mkdir(parents=True, exist_ok=True)
with dst.open("wb") as f:
f.write(r.content)
return dst, True, ""
except Exception as e:
last = str(e)
return dst, False, last
except Exception as e:
return dst, False, str(e)
ok, fail = 0, 0
with ThreadPoolExecutor(max_workers=min(16, os.cpu_count() or 4)) as ex:
futs = [ex.submit(_fetch, job) for job in q]
for fu in as_completed(futs):
dst, success, msg = fu.result()
if success:
ok += 1
else:
fail += 1
print(f"[WARN] Download failed for {dst.name}: {msg}", file=sys.stderr)
print(f"[INFO] CSV downloads done. ok={ok}, failed={fail}")
# ---------------------------
# NEW: URL download from COCO annotations (flickr_* / coco_url / url)
# ---------------------------
def _candidate_url_from_image(im: dict) -> Optional[str]:
# Priority order
for k in ("flickr_640_url", "flickr_url", "coco_url", "url"):
u = im.get(k)
if isinstance(u, str) and u.strip():
return u.strip()
return None
def ensure_images_from_ann_urls(
coco_images_by_id: Dict[int, dict],
target_images_dir: Path
) -> None:
"""
Download images based on URLs directly from COCO image objects.
Target filename is ALWAYS image['file_name'].
"""
target_images_dir.mkdir(parents=True, exist_ok=True)
jobs: List[Tuple[str, Path]] = []
for im in coco_images_by_id.values():
url = _candidate_url_from_image(im)
if not url:
continue
dst = target_images_dir / Path(im["file_name"]).name
if dst.exists():
continue
jobs.append((url, dst))
if not jobs:
print("[INFO] No annotation-based downloads required (all images present or no URLs).")
return
print(f"[INFO] Downloading {len(jobs)} images from annotation URLs into {target_images_dir} ...")
import time as _time
import requests
from concurrent.futures import ThreadPoolExecutor, as_completed
def _fetch(u_dst: Tuple[str, Path]) -> Tuple[Path, bool, str]:
u, dst = u_dst
try:
last = ""
for attempt in range(3):
try:
r = requests.get(u, timeout=25)
r.raise_for_status()
dst.parent.mkdir(parents=True, exist_ok=True)
with dst.open("wb") as f:
f.write(r.content)
return dst, True, ""
except Exception as e:
last = str(e)
_time.sleep(0.5 * (attempt + 1))
return dst, False, last
except Exception as e:
return dst, False, str(e)
ok, fail = 0, 0
with ThreadPoolExecutor(max_workers=min(16, os.cpu_count() or 4)) as ex:
futs = [ex.submit(_fetch, job) for job in jobs]
for fu in as_completed(futs):
dst, success, msg = fu.result()
if success:
ok += 1
else:
fail += 1
print(f"[WARN] Annotation URL download failed for {dst.name}: {msg}", file=sys.stderr)
print(f"[INFO] Annotation-based downloads done. ok={ok}, failed={fail}")
# ---------------------------
# Core helpers
# ---------------------------
def build_primary_label(anns: List[dict], cats_by_id: Dict[int, dict], label_map: Dict[str, Any]) -> Optional[str]:
if not anns:
return None
mapped = []
for a in anns:
cat = cats_by_id.get(a["category_id"])
if not cat:
continue
mapped.append(map_class(cat["name"], label_map))
if not mapped:
return None
cnt = Counter(mapped)
return cnt.most_common(1)[0][0]
def write_yolo_seg_label_file(
label_path: Path,
ann_list: List[dict],
cats_by_id: Dict[int, dict],
label_map: Dict[str, Any],
img_w: int,
img_h: int,
class_to_index: Dict[str, int],
min_poly_area_px: float = 0.0
) -> int:
"""
Write YOLOv8-seg label file.
NEW: min_poly_area_px > 0 filters small polygons (shoelace area in pixels).
"""
lines = []
for ann in ann_list:
cat = cats_by_id.get(ann["category_id"])
if not cat:
continue
mapped = map_class(cat["name"], label_map)
cls_id = class_to_index[mapped]
polys = coco_segmentation_to_yolo_polys(ann.get("segmentation"), img_w, img_h)
for poly in polys:
if len(poly) < 6:
continue
if min_poly_area_px > 0.0:
area = _polygon_area_px(poly, img_w, img_h)
if area < min_poly_area_px:
continue
lines.append(" ".join([str(cls_id)] + [f"{p:.6f}" for p in poly]))
if lines:
label_path.parent.mkdir(parents=True, exist_ok=True)
with label_path.open("w", encoding="utf-8") as f:
f.write("\n".join(lines))
return len(lines)
def items_from_coco(
coco: Dict[str, Any],
images_dir: Path,
cats_by_id: Dict[int, dict],
anns_by_image: Dict[int, List[dict]],
label_map: Dict[str, Any]
) -> Tuple[List[dict], Counter, Counter]:
items = []
raw_class_stats = Counter()
unmapped = Counter()
for img_id, img in {im["id"]: im for im in coco.get("images", [])}.items():
file_name = img["file_name"]
width = img.get("width")
height = img.get("height")
# get dims if missing
if width is None or height is None:
try:
with Image.open(images_dir / file_name) as im:
width, height = im.size
except Exception:
# try with basename
try:
with Image.open(images_dir / Path(file_name).name) as im:
width, height = im.size
file_name = Path(file_name).name
except Exception:
continue
ann_list = anns_by_image.get(img_id, [])
for a in ann_list:
cat = cats_by_id.get(a["category_id"])
if cat:
raw_class_stats[cat["name"]] += 1
mapped = map_class(cat["name"], label_map)
if mapped not in label_map["target_classes"]:
unmapped[mapped] += 1
primary = build_primary_label(ann_list, cats_by_id, label_map)
items.append({
"id": img_id,
"file_name": file_name,
"width": width,
"height": height,
"primary_label": primary,
"anns": ann_list
})
return items, raw_class_stats, unmapped
def write_split(
split_name: str,
split_items: List[dict],
out_root: Path,
images_dir: Path,
cats_by_id: Dict[int, dict],
label_map: Dict[str, Any],
class_to_index: Dict[str, int],
min_poly_area_px: float = 0.0
) -> Tuple[int, Counter, List[dict]]:
(out_root / "images" / split_name).mkdir(parents=True, exist_ok=True)
(out_root / "labels" / split_name).mkdir(parents=True, exist_ok=True)
count_used = 0
class_presence = Counter()
used_items: List[dict] = []
for it in split_items:
# copy/symlink image
src_img = images_dir / it["file_name"]
if not src_img.exists():
# try basename fallback
src_img = images_dir / Path(it["file_name"]).name
if not src_img.exists():
# skip if missing even after downloads
continue
dst_img = out_root / "images" / split_name / Path(it["file_name"]).name
safe_symlink_or_copy(src_img, dst_img)
# write label file
lbl_path = out_root / "labels" / split_name / (Path(it["file_name"]).stem + ".txt")
n_written = write_yolo_seg_label_file(
lbl_path, it["anns"], cats_by_id, label_map, it["width"], it["height"], class_to_index,
min_poly_area_px=min_poly_area_px
)
if n_written == 0:
if lbl_path.exists():
lbl_path.unlink()
if dst_img.exists():
try:
dst_img.unlink()
except Exception:
pass
continue
mapped_classes = []
for a in it["anns"]:
cat = cats_by_id.get(a["category_id"])
if cat:
mapped_classes.append(map_class(cat["name"], label_map))
for mc in set(mapped_classes):
class_presence[mc] += 1
count_used += 1
used_items.append(it)
return count_used, class_presence, used_items
# ---------------------------
# Manifests (per split)
# ---------------------------
def write_split_manifest(
manifests_root: Path,
split_name: str,
used_items: List[dict],
images_by_id: Dict[int, dict]
) -> int:
"""
Write a manifest file per split:
<file_name>\t<best_url>
Preferred URL: flickr_640_url, then flickr_url, coco_url, url.
Entries without URL are skipped.
Returns: number of written lines.
"""
path = manifests_root / f"{split_name}.txt"
path.parent.mkdir(parents=True, exist_ok=True)
lines: List[str] = []
for it in used_items:
im = images_by_id.get(it["id"]) if images_by_id is not None else None
url = _candidate_url_from_image(im) if im else None
if not url:
continue
lines.append(f"{Path(it['file_name']).name}\t{url}")
if lines:
with path.open("w", encoding="utf-8") as f:
f.write("\n".join(lines))
return len(lines)
else:
with path.open("w", encoding="utf-8") as f:
f.write("")
return 0
# ---------------------------
# Pipeline: single annotations (stratified split)
# ---------------------------
def prepare_single_annotations(
images_dir: Path,
annotations_path: Path,
label_map_path: Path,
out_root: Path,
train_ratio: float,
val_ratio: float,
seed: int,
dedupe: bool,
images_urls_csv: Optional[Path],
min_poly_area_px: float = 0.0
):
label_map = load_label_map(label_map_path)
target_classes: List[str] = label_map["target_classes"]
class_to_index = {c: i for i, c in enumerate(target_classes)}
coco = read_json(annotations_path)
images_by_id, cats_by_id, anns_by_image = index_coco(coco)
# 1) Annotation-based download (if directory is empty/not present)
if not images_dir.exists() or not any(images_dir.glob("*")):
print("[INFO] images_dir empty/not found -> downloading from annotation URLs (flickr_640_url/flickr_url/etc.)")
ensure_images_from_ann_urls(images_by_id, images_dir)
# 2) CSV fallback (optional) if still empty
if images_urls_csv and (not any(images_dir.glob("*"))):
print(f"[INFO] images_dir still empty -> downloading from CSV {images_urls_csv}")
ensure_images_from_urls(images_urls_csv, images_dir, images_by_id)
# Build items
items, raw_stats, unmapped = items_from_coco(coco, images_dir, cats_by_id, anns_by_image, label_map)
# Filter images without annotations
items = [x for x in items if x["anns"]]
# Dedupe
if dedupe:
seen = {}
deduped = []
for it in items:
cand = images_dir / it["file_name"]
if not cand.exists():
cand = images_dir / Path(it["file_name"]).name
ph = image_md5(cand) if cand.exists() else None
key = ph or it["file_name"]
if key in seen:
continue
seen[key] = True
deduped.append(it)
items = deduped
# Stratified split
y = [it["primary_label"] or "other" for it in items]
train_idx, val_idx, test_idx = stratified_split(items, y, train_ratio, val_ratio, seed)
idx_set = {"train": set(train_idx), "val": set(val_idx), "test": set(test_idx)}
# Write splits (one item per call -> logic unchanged)
per_split_counts = {"train": 0, "val": 0, "test": 0}
per_split_class = {s: Counter() for s in ("train", "val", "test")}
per_split_used: Dict[str, List[dict]] = {"train": [], "val": [], "test": []}
mapping_report = []
for i, it in enumerate(items):
split = "train" if i in idx_set["train"] else "val" if i in idx_set["val"] else "test"
n_used_before = per_split_counts[split]
used, class_presence, used_items = write_split(
split, [it], out_root, images_dir, cats_by_id, label_map, class_to_index,
min_poly_area_px=min_poly_area_px
)
per_split_counts[split] += used
for k, v in class_presence.items():
per_split_class[split][k] += v
if used_items:
per_split_used[split].extend(used_items)
if per_split_counts[split] > n_used_before:
mapped_classes = []
for a in it["anns"]:
cat = cats_by_id.get(a["category_id"])
if cat:
mapped_classes.append(map_class(cat["name"], label_map))
mapping_report.append({
"file_name": it["file_name"],
"primary": it["primary_label"] or "other",
"classes_in_image": list(sorted(set(mapped_classes)))
})
# dataset.yaml
names = target_classes
dataset_yaml = {
"path": str(out_root.resolve()),
"train": "images/train",
"val": "images/val",
"test": "images/test",
"names": names
}
write_json(out_root / "dataset.yolov8.json", dataset_yaml)
write_text(out_root / "dataset.yaml",
"path: {}\ntrain: {}\nval: {}\ntest: {}\nnames:\n".format(
dataset_yaml["path"], dataset_yaml["train"], dataset_yaml["val"], dataset_yaml["test"]
) + "".join([f" {i}: {n}\n" for i, n in enumerate(names)]))
# Stats
stats = {
"total_images_after_filter": sum(per_split_counts.values()),
"per_split_counts": per_split_counts,
"per_split_class_presence": {k: dict(v) for k, v in per_split_class.items()},
"raw_class_counts": dict(raw_stats),
"unmapped_buckets_seen": dict(unmapped)
}
write_json(out_root / "class_stats.json", stats)
write_json(out_root / "mapping_report.jsonl", mapping_report)
# Write manifests per split
manif_root = out_root / "manifests"
for split in ("train", "val", "test"):
write_split_manifest(manif_root, split, per_split_used[split], images_by_id)
print("=== DONE (single annotations) ===")
print(json.dumps(stats, indent=2, ensure_ascii=False))
print(f"dataset.yaml -> {out_root / 'dataset.yaml'}")
# ---------------------------
# Pipeline: pre-split annotations (train/val/test provided)
# ---------------------------
def prepare_presplit_annotations(
images_dir: Path,
annotations_train: Path,
annotations_val: Path,
annotations_test: Path,
label_map_path: Path,
out_root: Path,
dedupe: bool,
images_urls_csv: Optional[Path],
min_poly_area_px: float = 0.0
):
label_map = load_label_map(label_map_path)
target_classes: List[str] = label_map["target_classes"]
class_to_index = {c: i for i, c in enumerate(target_classes)}
# load three COCOs
coco_train = read_json(annotations_train)
coco_val = read_json(annotations_val)
coco_test = read_json(annotations_test)
# index
imgs_tr, cats_tr, anns_tr = index_coco(coco_train)
imgs_vl, cats_vl, anns_vl = index_coco(coco_val)
imgs_te, cats_te, anns_te = index_coco(coco_test)
# Ensure images (first: from annotation URLs)
if not images_dir.exists() or not any(images_dir.glob("*")):
print("[INFO] images_dir empty/not found -> downloading from annotation URLs (train/val/test merged)")
merged = {**imgs_tr, **imgs_vl, **imgs_te}
ensure_images_from_ann_urls(merged, images_dir)
# CSV fallback
if images_urls_csv and (not any(images_dir.glob("*"))):
print(f"[INFO] images_dir still empty -> downloading from CSV {images_urls_csv}")
merged = {**imgs_tr, **imgs_vl, **imgs_te}
ensure_images_from_urls(images_urls_csv, images_dir, merged)
# build items per split
items_tr, raw_tr, unm_tr = items_from_coco(coco_train, images_dir, cats_tr, anns_tr, label_map)
items_vl, raw_vl, unm_vl = items_from_coco(coco_val, images_dir, cats_vl, anns_vl, label_map)
items_te, raw_te, unm_te = items_from_coco(coco_test, images_dir, cats_te, anns_te, label_map)
# Filter empties
items_tr = [x for x in items_tr if x["anns"]]
items_vl = [x for x in items_vl if x["anns"]]
items_te = [x for x in items_te if x["anns"]]
# Dedupe (within each split)
if dedupe:
def _dedupe(items: List[dict]) -> List[dict]:
seen = {}
out = []
for it in items:
cand = images_dir / it["file_name"]
if not cand.exists():
cand = images_dir / Path(it["file_name"]).name
ph = image_md5(cand) if cand.exists() else None
key = ph or it["file_name"]
if key in seen:
continue
seen[key] = True
out.append(it)
return out
items_tr = _dedupe(items_tr)
items_vl = _dedupe(items_vl)
items_te = _dedupe(items_te)
# Write each split
per_split_counts = {"train": 0, "val": 0, "test": 0}
per_split_class = {s: Counter() for s in ("train", "val", "test")}
used, cpres, used_items_tr = write_split("train", items_tr, out_root, images_dir, cats_tr, label_map, class_to_index,
min_poly_area_px=min_poly_area_px)
per_split_counts["train"] += used
per_split_class["train"].update(cpres)
used, cpres, used_items_vl = write_split("val", items_vl, out_root, images_dir, cats_vl, label_map, class_to_index,
min_poly_area_px=min_poly_area_px)
per_split_counts["val"] += used
per_split_class["val"].update(cpres)
used, cpres, used_items_te = write_split("test", items_te, out_root, images_dir, cats_te, label_map, class_to_index,
min_poly_area_px=min_poly_area_px)
per_split_counts["test"] += used
per_split_class["test"].update(cpres)
# dataset.yaml
names = target_classes
dataset_yaml = {
"path": str(out_root.resolve()),
"train": "images/train",
"val": "images/val",
"test": "images/test",
"names": names
}
write_json(out_root / "dataset.yolov8.json", dataset_yaml)
write_text(out_root / "dataset.yaml",
"path: {}\ntrain: {}\nval: {}\ntest: {}\nnames:\n".format(
dataset_yaml["path"], dataset_yaml["train"], dataset_yaml["val"], dataset_yaml["test"]
) + "".join([f" {i}: {n}\n" for i, n in enumerate(names)]))
# Stats
raw_total = Counter()
raw_total.update(raw_tr); raw_total.update(raw_vl); raw_total.update(raw_te)
unm_total = Counter()
unm_total.update(unm_tr); unm_total.update(unm_vl); unm_total.update(unm_te)
stats = {
"total_images_after_filter": sum(per_split_counts.values()),
"per_split_counts": per_split_counts,
"per_split_class_presence": {k: dict(v) for k, v in per_split_class.items()},
"raw_class_counts": dict(raw_total),
"unmapped_buckets_seen": dict(unm_total)
}
write_json(out_root / "class_stats.json", stats)
# Manifests per split
manif_root = out_root / "manifests"
write_split_manifest(manif_root, "train", used_items_tr, imgs_tr)
write_split_manifest(manif_root, "val", used_items_vl, imgs_vl)
write_split_manifest(manif_root, "test", used_items_te, imgs_te)
print("=== DONE (pre-split annotations) ===")
print(json.dumps(stats, indent=2, ensure_ascii=False))
print(f"dataset.yaml -> {out_root / 'dataset.yaml'}")
# ---------------------------
# CLI
# ---------------------------
def parse_args():
p = argparse.ArgumentParser(description="Prepare TACO for YOLOv8-seg with Alami label mapping.")
# Image source
p.add_argument("--images_dir", required=True, type=Path,
help="Path to target image directory (created if missing). "
"If empty -> download from annotation URLs; optional CSV fallback.")
p.add_argument("--images_urls_csv", type=Path, default=None,
help="CSV with image URLs (one per line). Optional fallback.")
# Mode A (single annotations file; we produce the split)
p.add_argument("--annotations", type=Path,
help="Path to COCO annotations.json (single-file mode with internal split).")
# Mode B (pre-split annotations)
p.add_argument("--annotations_train", type=Path, help="COCO annotations train.json")
p.add_argument("--annotations_val", type=Path, help="COCO annotations val.json")
p.add_argument("--annotations_test", type=Path, help="COCO annotations test.json")
# General
p.add_argument("--label_map", required=True, type=Path, help="Path to ml/configs/label_map.json")
p.add_argument("--out", required=True, type=Path, help="Output root, e.g., ml/datasets/taco")
# Only for Mode A (we create the split)
p.add_argument("--train_ratio", type=float, default=0.8)
p.add_argument("--val_ratio", type=float, default=0.1)
p.add_argument("--seed", type=int, default=42)
# Optional
p.add_argument("--no_dedupe", action="store_true", help="Disable duplicate filtering")
# NEW
p.add_argument("--min_poly_area_px", type=float, default=0.0,
help="Minimum polygon area in pixels; polygons below are discarded (default 0 = off)")
args = p.parse_args()
have_single = args.annotations is not None
have_presplit = all([args.annotations_train, args.annotations_val, args.annotations_test])
if not have_single and not have_presplit:
p.error("Specify either --annotations OR (--annotations_train, --annotations_val, --annotations_test).")
if have_single:
if not (0 < args.train_ratio < 1):
p.error("--train_ratio must be in (0,1)")
if not (0 <= args.val_ratio < 1):
p.error("--val_ratio must be in [0,1)")
if args.train_ratio + args.val_ratio >= 1:
p.error("train_ratio + val_ratio must be < 1 (remainder is test)")
return args
def main():
args = parse_args()
out_root = args.out
out_root.mkdir(parents=True, exist_ok=True)
if args.annotations:
prepare_single_annotations(
images_dir=args.images_dir,
annotations_path=args.annotations,
label_map_path=args.label_map,
out_root=out_root,
train_ratio=args.train_ratio,
val_ratio=args.val_ratio,
seed=args.seed,
dedupe=(not args.no_dedupe),
images_urls_csv=args.images_urls_csv,
min_poly_area_px=float(args.min_poly_area_px)
)
else:
prepare_presplit_annotations(
images_dir=args.images_dir,
annotations_train=args.annotations_train,
annotations_val=args.annotations_val,
annotations_test=args.annotations_test,
label_map_path=args.label_map,
out_root=out_root,
dedupe=(not args.no_dedupe),
images_urls_csv=args.images_urls_csv,
min_poly_area_px=float(args.min_poly_area_px)
)
if __name__ == "__main__":
main()