Team Ai
Modelpublic

cuibinge/typical-marine-ecological-feature-recognition-code

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
prepare_polygon_dataset.py277 linesDownload Raw Back to scripts
1"""Build the normalized polygon dataset used by SAMPoly-style training.2 3The importer is intentionally strict: bbox-only annotations are rejected because4they cannot supervise true polygon boundaries or vertices.5"""6 7from __future__ import annotations8 9import argparse10import json11import random12import shutil13from dataclasses import asdict, dataclass14from pathlib import Path15from typing import Any16 17from PIL import Image18 19 20IMAGE_SUFFIXES = {".jpg", ".jpeg", ".png", ".tif", ".tiff"}21MASK_SUFFIXES = {".png", ".tif", ".tiff", ".jpg", ".jpeg"}22POLYGON_FORMATS = {"coco_polygon", "coco_segmentation", "geojson", "shp", "mask", "binary_mask", "semantic_mask"}23BBOX_FORMATS = {"bbox", "box_txt", "coco_bbox", "voc_bbox"}24 25 26@dataclass27class ImportStats:28    scanned: int = 029    accepted: int = 030    rejected: int = 031    accepted_masks: int = 032    accepted_polygons: int = 033    rejected_bbox_only: int = 034    rejected_missing_image: int = 035    rejected_missing_label: int = 036    rejected_unknown_format: int = 037 38 39def parse_args() -> argparse.Namespace:40    parser = argparse.ArgumentParser(description=__doc__)41    parser.add_argument("--manifest", required=True, help="JSONL manifest with standardized sample records.")42    parser.add_argument("--bbox-source-root", default=None, help="Optional local bbox dataset mirror for rejection auditing.")43    parser.add_argument("--extra-source-root", action="append", default=[], help="Local source roots to scan for mask/polygon datasets.")44    parser.add_argument("--output-root", required=True)45    parser.add_argument("--train-ratio", type=float, default=0.8)46    parser.add_argument("--val-ratio", type=float, default=0.1)47    parser.add_argument("--seed", type=int, default=0)48    parser.add_argument("--min-quality-score", type=float, default=0.9)49    parser.add_argument("--element", default=None)50    return parser.parse_args()51 52 53def read_jsonl(path: Path) -> list[dict[str, Any]]:54    rows = []55    if not path.exists():56        return rows57    for line in path.read_text(encoding="utf-8").splitlines():58        if line.strip():59            rows.append(json.loads(line))60    return rows61 62 63def write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None:64    path.parent.mkdir(parents=True, exist_ok=True)65    path.write_text("\n".join(json.dumps(row, ensure_ascii=False) for row in rows) + ("\n" if rows else ""), encoding="utf-8")66 67 68def safe_name(sample_id: str, fallback: str) -> str:69    raw = sample_id or Path(fallback).stem70    return "".join(ch if ch.isalnum() or ch in "._-" else "_" for ch in raw)71 72 73def local_path_from_record(record: dict[str, Any], key: str) -> Path | None:74    value = record.get(key)75    if not value or not isinstance(value, str):76        return None77    if value.startswith("hf://"):78        return None79    path = Path(value)80    return path if path.exists() else None81 82 83def find_local_bbox_image(record: dict[str, Any], bbox_root: Path | None) -> Path | None:84    if bbox_root is None:85        return None86    source = str(record.get("image_path") or "")87    stem = Path(source).stem.lower()88    for split in ("train", "val", "test"):89        image_dir = bbox_root / "images" / split90        if not image_dir.exists():91            continue92        for path in image_dir.iterdir():93            if path.suffix.lower() in IMAGE_SUFFIXES and path.stem.lower().endswith(stem):94                return path95    return None96 97 98def mask_has_foreground(path: Path) -> bool:99    try:100        img = Image.open(path).convert("L")101        extrema = img.getextrema()102        return bool(extrema and extrema[1] > 0)103    except Exception:104        return False105 106 107def find_extra_samples(root: Path, min_quality: float, element: str | None) -> list[dict[str, Any]]:108    rows: list[dict[str, Any]] = []109    for image_dir in root.rglob("images"):110        if not image_dir.is_dir():111            continue112        split = image_dir.parent.name if image_dir.parent.name in {"train", "val", "test"} else None113        mask_dir_candidates = [114            image_dir.parent / "masks",115            image_dir.parent.parent / "masks" / image_dir.name,116            image_dir.parent.parent / "masks" / (split or ""),117        ]118        for image_path in image_dir.iterdir():119            if image_path.suffix.lower() not in IMAGE_SUFFIXES:120                continue121            mask_path = None122            for mask_dir in mask_dir_candidates:123                if not mask_dir.exists():124                    continue125                for suffix in MASK_SUFFIXES:126                    candidate = mask_dir / f"{image_path.stem}{suffix}"127                    if candidate.exists():128                        mask_path = candidate129                        break130                if mask_path:131                    break132            if not mask_path or not mask_has_foreground(mask_path):133                continue134            rows.append(135                {136                    "sample_id": f"local_{safe_name(image_path.stem, image_path.name)}",137                    "element": element or "unknown",138                    "task_type": "polygon_extraction",139                    "image_path": str(image_path),140                    "mask_path": str(mask_path),141                    "annotation_path": str(mask_path),142                    "annotation_format": "binary_mask",143                    "quality_score": max(min_quality, 0.95),144                    "quality_flags": ["accepted", "local_mask_pair", "polygon_trainable"],145                    "split": split,146                }147            )148    return rows149 150 151def split_rows(rows: list[dict[str, Any]], train_ratio: float, val_ratio: float, seed: int) -> dict[str, list[dict[str, Any]]]:152    grouped = {"train": [], "val": [], "test": []}153    presplit = [row for row in rows if row.get("split") in grouped]154    unsplit = [row for row in rows if row.get("split") not in grouped]155    for row in presplit:156        grouped[str(row["split"])].append(row)157    random.Random(seed).shuffle(unsplit)158    n = len(unsplit)159    n_train = int(n * train_ratio)160    n_val = int(n * val_ratio)161    grouped["train"].extend(unsplit[:n_train])162    grouped["val"].extend(unsplit[n_train : n_train + n_val])163    grouped["test"].extend(unsplit[n_train + n_val :])164    return grouped165 166 167def copy_sample(row: dict[str, Any], split: str, output_root: Path) -> dict[str, Any]:168    image_path = Path(str(row["image_path"]))169    mask_path = Path(str(row.get("mask_path") or row.get("annotation_path")))170    name = safe_name(str(row.get("sample_id") or image_path.stem), image_path.name)171    image_out = output_root / "images" / split / f"{name}{image_path.suffix.lower()}"172    mask_out = output_root / "masks" / split / f"{name}.png"173    image_out.parent.mkdir(parents=True, exist_ok=True)174    mask_out.parent.mkdir(parents=True, exist_ok=True)175    shutil.copy2(image_path, image_out)176    Image.open(mask_path).convert("L").save(mask_out)177    copied = dict(row)178    copied.update(179        {180            "sample_id": name,181            "split": split,182            "image_path": str(image_out),183            "mask_path": str(mask_out),184            "annotation_path": str(mask_out),185            "annotation_format": "binary_mask",186            "task_type": "polygon_extraction",187            "quality_flags": sorted(set(row.get("quality_flags", []) + ["accepted_for_polygon_training"])),188        }189    )190    return copied191 192 193def main() -> None:194    args = parse_args()195    manifest = Path(args.manifest)196    output_root = Path(args.output_root)197    output_root.mkdir(parents=True, exist_ok=True)198    bbox_root = Path(args.bbox_source_root) if args.bbox_source_root else None199    stats = ImportStats()200    accepted: list[dict[str, Any]] = []201    rejected: list[dict[str, Any]] = []202 203    records = read_jsonl(manifest)204    for root in args.extra_source_root:205        records.extend(find_extra_samples(Path(root), args.min_quality_score, args.element))206 207    for record in records:208        stats.scanned += 1209        if args.element and record.get("element") != args.element:210            continue211        quality = float(record.get("quality_score") or 0.0)212        fmt = str(record.get("annotation_format") or "").lower()213        image_path = local_path_from_record(record, "image_path") or find_local_bbox_image(record, bbox_root)214        label_path = local_path_from_record(record, "mask_path") or local_path_from_record(record, "annotation_path")215 216        reject_reason = None217        if quality < args.min_quality_score:218            reject_reason = "quality_below_threshold"219        elif fmt in BBOX_FORMATS:220            reject_reason = "bbox_only_not_polygon_trainable"221            stats.rejected_bbox_only += 1222        elif fmt not in POLYGON_FORMATS:223            reject_reason = "unknown_or_unsupported_annotation_format"224            stats.rejected_unknown_format += 1225        elif image_path is None:226            reject_reason = "missing_local_image"227            stats.rejected_missing_image += 1228        elif label_path is None or not label_path.exists():229            reject_reason = "missing_local_mask_or_polygon"230            stats.rejected_missing_label += 1231        elif fmt in {"mask", "binary_mask", "semantic_mask"} and not mask_has_foreground(label_path):232            reject_reason = "empty_or_invalid_mask"233 234        if reject_reason:235            item = dict(record)236            item["polygon_import_status"] = "rejected"237            item["reject_reason"] = reject_reason238            if image_path:239                item["local_image_path"] = str(image_path)240            rejected.append(item)241            stats.rejected += 1242            continue243 244        item = dict(record)245        item["image_path"] = str(image_path)246        item["mask_path"] = str(label_path)247        item["annotation_path"] = str(label_path)248        item["polygon_import_status"] = "accepted"249        accepted.append(item)250        stats.accepted += 1251        if fmt in {"mask", "binary_mask", "semantic_mask"}:252            stats.accepted_masks += 1253        else:254            stats.accepted_polygons += 1255 256    grouped = split_rows(accepted, args.train_ratio, args.val_ratio, args.seed)257    copied_rows = []258    for split, rows in grouped.items():259        for row in rows:260            copied_rows.append(copy_sample(row, split, output_root))261 262    write_jsonl(output_root / "manifests" / "accepted_polygon_samples.jsonl", copied_rows)263    write_jsonl(output_root / "manifests" / "rejected_polygon_samples.jsonl", rejected)264    summary = {265        **asdict(stats),266        "output_root": str(output_root),267        "splits": {split: len(rows) for split, rows in grouped.items()},268        "quality_policy": "Only mask or polygon annotations are accepted for SAMPoly-style polygon training; bbox-only samples are rejected.",269        "source_manifest": str(manifest),270    }271    (output_root / "dataset_card.json").write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")272    print(json.dumps(summary, indent=2, ensure_ascii=False), flush=True)273 274 275if __name__ == "__main__":276    main()277