alami-vision-api / ml /scripts /augment_external_datasets.py
alami-ci
Deploy from alami-eco/alami-trash-ai@aee69796b70947e95efdb9c7483fa52f8d3b4520
76838d6
Raw
History Blame Contribute Delete
12.7 kB
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
augment_external_datasets.py
Integrate external COCO datasets (e.g., Keremberke / Roboflow / Kaggle) into your TACO→YOLOv8-seg workflow:
- Class mapping via label_map.json
- Polygons from COCO 'segmentation' (list). If only bounding boxes are available: generate rectangle polygons.
- Write YOLOv8-seg labels, link/copy images into target structure
- Extend manifests {train,val,test}.txt (informational; not needed for local prefetching)
Example:
python ml/scripts/augment_external_datasets.py \
--coco external/keremberke/annotations.json \
--images_dir external/keremberke/images \
--out_root ml/datasets/taco \
--manifest_root ml/datasets/taco/manifests \
--label_map ml/configs/label_map.json \
--split_policy train_only
Split policies:
- train_only (default): put all external material into 'train'
- stratify: 80/10/10 (reproducible via --seed)
- keep_existing: use 'images', 'annotations' and any existing 'train/val/test' keys like in COCO (if present)
"""
from __future__ import annotations
import argparse, json, os, random, sys, shutil
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
from collections import defaultdict, Counter
# --- Utils ---
def read_json(p: Path) -> Any:
return json.loads(p.read_text(encoding="utf-8"))
def write_text(p: Path, s: str) -> None:
p.parent.mkdir(parents=True, exist_ok=True)
p.write_text(s, encoding="utf-8")
def safe_symlink_or_copy(src: Path, dst: Path) -> None:
dst.parent.mkdir(parents=True, exist_ok=True)
if os.name == "nt":
if not dst.exists():
shutil.copy2(src, dst)
else:
try:
if dst.exists():
dst.unlink()
dst.symlink_to(src.resolve())
except Exception:
if not dst.exists():
shutil.copy2(src, dst)
# --- Label map ---
def load_label_map(label_map_path: Path) -> Dict[str, Any]:
lm = read_json(label_map_path)
rules = []
for r in lm.get("rules", []):
rules.append({
"to": r["to"],
"match_any_substring": [s.lower() for s in r.get("match_any_substring", [])],
"match_any_regex": r.get("match_any_regex", []),
})
lm["rules"] = rules
return lm
import re
def map_class(name: str, label_map: Dict[str, Any]) -> str:
n = name.lower().strip()
for r in label_map["rules"]:
if any(sub in n for sub in r["match_any_substring"]):
return r["to"]
for pat in r["match_any_regex"]:
if re.search(pat, n):
return r["to"]
return label_map.get("default", "other")
# --- COCO helpers ---
def index_coco(coco: Dict[str, Any]):
imgs = {im["id"]: im for im in coco.get("images", [])}
cats = {c["id"]: c for c in coco.get("categories", [])}
by_img = defaultdict(list)
for a in coco.get("annotations", []):
by_img[a["image_id"]].append(a)
return imgs, cats, by_img
def coco_seg_to_yolo_polys(seg, img_w: int, img_h: int) -> List[List[float]]:
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)
return polys
def bbox_to_rect_poly_xyxy(x1,y1,x2,y2, img_w:int, img_h:int) -> List[float]:
# clamp + normalize to [0,1], clockwise rectangle
x1n = max(0.0, min(x1 / img_w, 1.0)); y1n = max(0.0, min(y1 / img_h, 1.0))
x2n = max(0.0, min(x2 / img_w, 1.0)); y2n = max(0.0, min(y2 / img_h, 1.0))
return [x1n,y1n, x1n,y2n, x2n,y2n, x2n,y1n]
def write_yolo_seg_label(lbl_path: Path, anns: 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]) -> int:
lines = []
for a in anns:
cat = cats_by_id.get(a.get("category_id"))
if not cat:
continue
mapped = map_class(cat.get("name",""), label_map)
if mapped not in class_to_index:
# falls Mapping außerhalb eurer Zielklassen liegt -> skip
continue
cls_id = class_to_index[mapped]
wrote = 0
# 1) bevorzugt echte Polygone
seg = a.get("segmentation")
polys = coco_seg_to_yolo_polys(seg, img_w, img_h)
for poly in polys:
if len(poly) >= 6:
lines.append(" ".join([str(cls_id)] + [f"{p:.6f}" for p in poly]))
wrote += 1
# 2) falls keine Polygone und BBox existiert -> Rechteck-Polygon
if wrote == 0 and "bbox" in a and isinstance(a["bbox"], (list,tuple)) and len(a["bbox"]) >= 4:
x,y,w,h = a["bbox"][:4]
rect = bbox_to_rect_poly_xyxy(x, y, x+w, y+h, img_w, img_h)
lines.append(" ".join([str(cls_id)] + [f"{p:.6f}" for p in rect]))
wrote += 1
if lines:
lbl_path.parent.mkdir(parents=True, exist_ok=True)
lbl_path.write_text("\n".join(lines) + "\n", encoding="utf-8")
return len(lines)
# --- Split policies ---
def stratified_split(items: List[dict], y: List[str], train_ratio=0.8, val_ratio=0.1, seed=42):
rnd = random.Random(seed)
by = defaultdict(list)
for i,c in enumerate(y):
by[c].append(i)
train,val,test = [],[],[]
for c, idxs in by.items():
rnd.shuffle(idxs)
n=len(idxs); ntr=int(round(n*train_ratio)); nv=int(round(n*val_ratio))
train += idxs[:ntr]
val += idxs[ntr:ntr+nv]
test += idxs[ntr+nv:]
for arr in (train,val,test): rnd.shuffle(arr)
return train,val,test
# --- Main augmentation ---
def augment(
coco_path: Path,
images_dir: Path,
out_root: Path,
manifest_root: Path,
label_map_path: Path,
split_policy: str = "train_only",
train_ratio: float = 0.8,
val_ratio: float = 0.1,
seed: int = 42,
min_bbox_area_frac: float = 0.0,
target_split: str = "auto"
):
label_map = load_label_map(label_map_path)
target_classes = label_map["target_classes"]
class_to_index = {c:i for i,c in enumerate(target_classes)}
coco = read_json(coco_path)
images_by_id, cats_by_id, anns_by_image = index_coco(coco)
# Build items (only those with at least one annotation)
items = []
for img_id, im in images_by_id.items():
file_name = im["file_name"]
w = im.get("width") or 0
h = im.get("height") or 0
anns = anns_by_image.get(img_id, [])
if not anns:
continue
# optional: filter very small bounding boxes
if min_bbox_area_frac > 0 and w>0 and h>0:
keep = []
for a in anns:
if "bbox" in a:
bx,by,bw,bh = a["bbox"][:4]
if (bw*bh)/(w*h + 1e-9) >= min_bbox_area_frac:
keep.append(a)
else:
keep.append(a)
anns = keep
if not anns:
continue
# primary label
mapped = []
for a in anns:
cat = cats_by_id.get(a.get("category_id"))
if cat: mapped.append(map_class(cat.get("name",""), label_map))
primary = None
if mapped:
c = Counter(mapped)
primary = c.most_common(1)[0][0]
items.append({"id": img_id, "file_name": file_name, "width": w, "height": h, "anns": anns, "primary": primary})
# Split
splits = {"train": [], "val": [], "test": []}
def _infer_split_from_paths() -> str:
name = (str(coco_path).lower() + " " + str(images_dir).lower())
if "val" in name or "valid" in name or "validation" in name:
return "val"
if "test" in name:
return "test"
return "train"
if target_split in ("train", "val", "test"):
splits[target_split] = items
else:
if split_policy == "train_only":
splits["train"] = items
elif split_policy == "stratify":
y = [it["primary"] or "other" for it in items]
ti,vi,si = stratified_split(items, y, train_ratio, val_ratio, seed)
splits["train"] = [items[i] for i in ti]
splits["val"] = [items[i] for i in vi]
splits["test"] = [items[i] for i in si]
elif split_policy == "keep_existing":
# Try to infer target split from file names/paths
inferred = _infer_split_from_paths()
splits[inferred] = items
else:
raise ValueError(f"unknown split_policy: {split_policy}")
# Write
total_written = 0
per_split_written = {}
for split, arr in splits.items():
img_out = out_root / "images" / split
lbl_out = out_root / "labels" / split
mani = manifest_root / f"{split}.txt"
mani.parent.mkdir(parents=True, exist_ok=True)
# load existing manifest lines (keep)
existing = []
if mani.exists():
for ln in mani.read_text(encoding="utf-8").splitlines():
if not ln.strip(): continue
parts = ln.split("\t")
if len(parts)>=2: existing.append((parts[0], parts[1]))
new_lines = []
used = 0
for it in arr:
src = images_dir / Path(it["file_name"]).name
if not src.exists():
# exports often have subfolders -> fallback: try relative path from COCO
alt = images_dir / it["file_name"]
if alt.exists(): src = alt
if not src.exists():
# image missing -> skip
continue
# link/copy image into target
dst_img = img_out / Path(src).name
safe_symlink_or_copy(src, dst_img)
# write label
lbl_path = lbl_out / (dst_img.stem + ".txt")
n = write_yolo_seg_label(lbl_path, it["anns"], cats_by_id, label_map, it["width"], it["height"], class_to_index)
if n == 0:
# no usable segment -> optionally remove image again to keep clean
try:
if dst_img.exists(): dst_img.unlink()
except Exception:
pass
continue
# extend manifest (informational; local file)
new_lines.append(f"{dst_img.name}\t{src.resolve().as_posix()}")
used += 1
# append to manifest
with mani.open("a", encoding="utf-8") as f:
for ln in new_lines:
f.write(ln + "\n")
total_written += used
per_split_written[split] = used
print(json.dumps({"total_images_added": total_written, "per_split_added": per_split_written}, indent=2))
def parse_args():
ap = argparse.ArgumentParser(description="Augment an existing TACO-prepared dataset with an external COCO dataset.")
ap.add_argument("--coco", required=True, type=Path, help="Path to the external COCO annotations file")
ap.add_argument("--images_dir", required=True, type=Path, help="Folder with the corresponding images")
ap.add_argument("--out_root", required=True, type=Path, help="Target dataset root, e.g., ml/datasets/taco")
ap.add_argument("--manifest_root", required=True, type=Path, help="Manifest folder, e.g., ml/datasets/taco/manifests")
ap.add_argument("--label_map", required=True, type=Path, help="Your ml/configs/label_map.json")
ap.add_argument("--split_policy", choices=["train_only","stratify","keep_existing"], default="train_only")
ap.add_argument("--train_ratio", type=float, default=0.8)
ap.add_argument("--val_ratio", type=float, default=0.1)
ap.add_argument("--seed", type=int, default=42)
ap.add_argument("--min_bbox_area_frac", type=float, default=0.0, help="BBox area minimum relative to the image (0..1) to filter tiny objects")
ap.add_argument("--target_split", choices=["auto","train","val","test"], default="auto", help="Force write items into a specific split; auto will infer for keep_existing or use split_policy")
return ap.parse_args()
if __name__ == "__main__":
args = parse_args()
augment(args.coco, args.images_dir, args.out_root, args.manifest_root, args.label_map,
split_policy=args.split_policy, train_ratio=args.train_ratio, val_ratio=args.val_ratio,
seed=args.seed, min_bbox_area_frac=args.min_bbox_area_frac, target_split=args.target_split)