import argparse import csv import os import shutil from pathlib import Path def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Split a flat image folder by source-video prefixes from CSV." ) parser.add_argument( "--flat-root", type=Path, required=True, help="Flat folder with extracted images.", ) parser.add_argument( "--split-csv", type=Path, required=True, help="CSV with recommended_split and flat_prefix columns.", ) parser.add_argument( "--output-root", type=Path, required=True, help="Output folder that will contain train/valid/test.", ) parser.add_argument( "--image-ext", default="jpg", help="Image extension in the flat folder, default: jpg.", ) parser.add_argument( "--mode", choices=("hardlink", "copy"), default="hardlink", help="Write mode for split files. hardlink avoids doubling disk usage.", ) parser.add_argument( "--overwrite", action="store_true", help="Remove existing output_root before writing.", ) return parser.parse_args() def load_split_rows(csv_path: Path) -> list[dict]: with csv_path.open("r", encoding="utf-8-sig", newline="") as f: rows = list(csv.DictReader(f)) required = {"recommended_split", "flat_prefix"} if not rows: raise ValueError(f"No rows found in {csv_path}") missing = required - set(rows[0].keys()) if missing: raise ValueError(f"Missing required CSV columns: {sorted(missing)}") return rows def ensure_clean_dir(path: Path, overwrite: bool) -> None: if path.exists(): if not overwrite: raise FileExistsError( f"Output path already exists: {path}. Use --overwrite to replace it." ) shutil.rmtree(path) path.mkdir(parents=True, exist_ok=True) def link_or_copy(src: Path, dst: Path, mode: str) -> str: if mode == "hardlink": try: os.link(src, dst) return "hardlink" except OSError: shutil.copy2(src, dst) return "copy_fallback" shutil.copy2(src, dst) return "copy" def main() -> None: args = parse_args() rows = load_split_rows(args.split_csv) ext = args.image_ext.lower().lstrip(".") ensure_clean_dir(args.output_root, args.overwrite) manifest_path = args.output_root / "_split_manifest.csv" split_name_map = { "train": "train", "val": "valid", "valid": "valid", "test": "test", } total_written = 0 total_missing = 0 with manifest_path.open("w", encoding="utf-8-sig", newline="") as mf: writer = csv.DictWriter( mf, fieldnames=[ "recommended_split", "output_split", "group_name", "file_name", "flat_prefix", "written_count", "missing", "write_mode", ], ) writer.writeheader() for row in rows: split_raw = row["recommended_split"].strip().lower() output_split = split_name_map.get(split_raw) if output_split is None: raise ValueError(f"Unsupported split value: {row['recommended_split']}") prefix = row["flat_prefix"].strip() pattern = f"{prefix}__f*.{ext}" matches = sorted(args.flat_root.glob(pattern)) out_dir = args.output_root / output_split out_dir.mkdir(parents=True, exist_ok=True) write_mode = "" for src in matches: dst = out_dir / src.name write_mode = link_or_copy(src, dst, args.mode) missing = 0 if matches else 1 total_missing += missing total_written += len(matches) writer.writerow( { "recommended_split": row["recommended_split"], "output_split": output_split, "group_name": row.get("group_name", ""), "file_name": row.get("file_name", ""), "flat_prefix": prefix, "written_count": len(matches), "missing": missing, "write_mode": write_mode, } ) print(f"output_root={args.output_root}") print(f"manifest={manifest_path}") print(f"total_written={total_written}") print(f"missing_prefixes={total_missing}") if __name__ == "__main__": main()