Team Ai
Modelpublic

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

sourceHugging Faceupdated 3mo agoView on Hugging Face
0likes
import_patch_polygon_dataset.py391 linesDownload Raw Back to scripts
1"""Import patch folders into the project's normalized polygon dataset layout."""2 3from __future__ import annotations4 5import argparse6import csv7import json8import random9import re10import shutil11from collections import Counter, defaultdict12from dataclasses import asdict, dataclass13from pathlib import Path14from typing import Any15 16from PIL import Image, ImageDraw17 18 19ELEMENT_MAP = {20    "海上养殖区": {21        "id": "aquaculture",22        "name": "海上养殖区",23        "task_type": "polygon_extraction",24    },25    "港口岸线": {26        "id": "port_shoreline",27        "name": "港口岸线",28        "task_type": "polygon_extraction",29    },30    "粉砂淤泥质岸线": {31        "id": "silty_muddy_shoreline",32        "name": "粉砂淤泥质岸线",33        "task_type": "polygon_extraction",34    },35}36 37IMAGE_SUFFIXES = {".jpg", ".jpeg", ".png", ".tif", ".tiff"}38KEY_SUFFIX = re.compile(r"_(Orig|True|False|TrueColor|FalseColor|Binary|Label)_[^_]+$")39 40 41@dataclass42class ImportStats:43    scanned: int = 044    accepted: int = 045    rejected: int = 046    inference_assets: int = 047    missing_train_image: int = 048    missing_mask: int = 049    missing_geojson: int = 050    empty_mask: int = 051    invalid_geojson: int = 052    size_mismatch: int = 053 54 55def parse_args() -> argparse.Namespace:56    parser = argparse.ArgumentParser(description=__doc__)57    parser.add_argument("--source-root", required=True)58    parser.add_argument("--output-root", required=True)59    parser.add_argument("--train-ratio", type=float, default=0.8)60    parser.add_argument("--val-ratio", type=float, default=0.1)61    parser.add_argument("--seed", type=int, default=20260704)62    parser.add_argument("--preview-limit", type=int, default=36)63    return parser.parse_args()64 65 66def read_geojson(path: Path) -> dict[str, Any] | None:67    try:68        data = json.loads(path.read_text(encoding="utf-8"))69    except Exception:70        return None71    if data.get("type") != "FeatureCollection":72        return None73    features = data.get("features")74    if not isinstance(features, list) or len(features) == 0:75        return None76    for feature in features:77        geom = feature.get("geometry") if isinstance(feature, dict) else None78        if isinstance(geom, dict) and geom.get("type") in {"Polygon", "MultiPolygon", "LineString", "MultiLineString"}:79            return data80    return None81 82 83def file_key(path: Path) -> str:84    return KEY_SUFFIX.sub("", path.stem)85 86 87def index_files(directory: Path) -> dict[str, Path]:88    if not directory.exists():89        return {}90    return {file_key(path): path for path in directory.iterdir() if path.is_file()}91 92 93def safe_name(text: str) -> str:94    return "".join(ch if ch.isalnum() or ch in "._-" else "_" for ch in text)95 96 97def mask_stats(mask_path: Path) -> tuple[tuple[int, int] | None, int, float]:98    try:99        mask = Image.open(mask_path).convert("L")100    except Exception:101        return None, 0, 0.0102    extrema = mask.getextrema()103    if extrema is None or extrema[1] == 0:104        return mask.size, 0, 0.0105    binary = mask.point(lambda v: 255 if v > 0 else 0)106    foreground = 0107    hist = binary.histogram()108    if len(hist) > 255:109        foreground = hist[255]110    ratio = foreground / max(binary.size[0] * binary.size[1], 1)111    return binary.size, foreground, ratio112 113 114def image_size(path: Path) -> tuple[int, int] | None:115    try:116        return Image.open(path).size117    except Exception:118        return None119 120 121def discover_groups(source_root: Path) -> list[tuple[str, str, int, Path]]:122    groups = []123    for element_dir in sorted(source_root.iterdir()):124        if not element_dir.is_dir():125            continue126        for satellite_dir in sorted(element_dir.iterdir()):127            if not satellite_dir.is_dir():128                continue129            for size_dir in sorted(satellite_dir.glob("Size_*")):130                if not size_dir.is_dir():131                    continue132                try:133                    patch_size = int(size_dir.name.split("_", 1)[1])134                except Exception:135                    continue136                groups.append((element_dir.name, satellite_dir.name, patch_size, size_dir))137    return groups138 139 140def build_records(source_root: Path) -> tuple[list[dict[str, Any]], list[dict[str, Any]], ImportStats]:141    accepted: list[dict[str, Any]] = []142    rejected: list[dict[str, Any]] = []143    stats = ImportStats()144    for element_name, satellite, patch_size, size_dir in discover_groups(source_root):145        element = ELEMENT_MAP.get(element_name, {"id": safe_name(element_name), "name": element_name, "task_type": "polygon_extraction"})146        true_images = index_files(size_dir / "Image_TrueColor")147        false_images = index_files(size_dir / "Image_FalseColor")148        orig_images = index_files(size_dir / "Image_Orig")149        masks = index_files(size_dir / "Label_Binary")150        geojsons = index_files(size_dir / "Label_GeoJSON")151        keys = sorted(set(true_images) | set(false_images) | set(orig_images) | set(masks) | set(geojsons))152        for key in keys:153            stats.scanned += 1154            true_path = true_images.get(key)155            false_path = false_images.get(key)156            orig_path = orig_images.get(key)157            image_path = true_path or false_path158            mask_path = masks.get(key)159            geojson_path = geojsons.get(key)160            reject_reason = None161            if image_path is None:162                reject_reason = "missing_true_or_false_color_image"163                stats.missing_train_image += 1164            elif mask_path is None:165                reject_reason = "missing_binary_mask"166                stats.missing_mask += 1167            elif geojson_path is None:168                reject_reason = "missing_geojson_polygon"169                stats.missing_geojson += 1170            else:171                img_size = image_size(image_path)172                m_size, foreground, fg_ratio = mask_stats(mask_path)173                geojson = read_geojson(geojson_path)174                if m_size is None:175                    reject_reason = "unreadable_binary_mask"176                    stats.missing_mask += 1177                elif foreground == 0:178                    reject_reason = "empty_binary_mask"179                    stats.empty_mask += 1180                elif img_size is not None and img_size != m_size:181                    reject_reason = "image_mask_size_mismatch"182                    stats.size_mismatch += 1183                elif geojson is None:184                    reject_reason = "invalid_or_empty_geojson"185                    stats.invalid_geojson += 1186                else:187                    sample_id = safe_name(f"{element['id']}_{satellite}_{patch_size}_{key}")188                    accepted.append(189                        {190                            "sample_id": sample_id,191                            "source_key": key,192                            "element": element["id"],193                            "element_name": element["name"],194                            "task_type": element["task_type"],195                            "image_path": str(image_path),196                            "mask_path": str(mask_path),197                            "annotation_path": str(geojson_path),198                            "annotation_format": "geojson_polygon",199                            "label_encoding": {"0": "background", "1": element["id"]},200                            "satellite": satellite,201                            "sensor": infer_sensor(key),202                            "resolution_m": infer_resolution_m(satellite),203                            "patch_size": patch_size,204                            "bands": ["red", "green", "blue"],205                            "band_count": 3,206                            "dtype": "uint8",207                            "fusion": {208                                "state": infer_fusion_state(satellite),209                                "method": "unknown_patch_product",210                                "sources": [211                                    {"role": "true_color", "path": str(true_path) if true_path else None, "resolution_m": infer_resolution_m(satellite)},212                                    {"role": "false_color", "path": str(false_path) if false_path else None, "resolution_m": infer_resolution_m(satellite)},213                                    {"role": "original_patch", "path": str(orig_path) if orig_path else None, "resolution_m": infer_resolution_m(satellite)},214                                ],215                                "target_resolution_m": infer_resolution_m(satellite),216                                "native_multispectral_resolution_m": None,217                                "persisted": True,218                                "reproducible": False,219                                "spectral_preservation": "unknown",220                            },221                            "source_project": "BaiduNetdisk_Patches",222                            "source_dataset": str(source_root),223                            "quality_flags": [224                                "accepted_for_polygon_training",225                                "paired_image_mask_geojson",226                                f"foreground_ratio:{fg_ratio:.6f}",227                            ],228                            "quality_score": 0.96,229                        }230                    )231                    stats.accepted += 1232                    continue233            item = {234                "source_key": key,235                "element": element["id"],236                "element_name": element["name"],237                "satellite": satellite,238                "patch_size": patch_size,239                "image_path": str(image_path) if image_path else None,240                "orig_image_path": str(orig_path) if orig_path else None,241                "mask_path": str(mask_path) if mask_path else None,242                "annotation_path": str(geojson_path) if geojson_path else None,243                "polygon_import_status": "rejected",244                "reject_reason": reject_reason,245            }246            if image_path and not mask_path and not geojson_path:247                item["polygon_import_status"] = "inference_asset"248                stats.inference_assets += 1249            else:250                stats.rejected += 1251            rejected.append(item)252    return accepted, rejected, stats253 254 255def infer_sensor(key: str) -> str | None:256    for token in ("PMS1", "PMS2", "PMS", "MUX"):257        if token in key:258            return token259    return None260 261 262def infer_resolution_m(satellite: str) -> float | None:263    return {"GF1": 2.0, "GF2": 1.0, "GF6": 2.0}.get(satellite.upper())264 265 266def infer_fusion_state(satellite: str) -> str:267    return "fused_product" if satellite.upper() in {"GF1", "GF2", "GF6"} else "unknown"268 269 270def split_records(records: list[dict[str, Any]], train_ratio: float, val_ratio: float, seed: int) -> dict[str, list[dict[str, Any]]]:271    grouped: dict[tuple[str, str, int], list[dict[str, Any]]] = defaultdict(list)272    for record in records:273        grouped[(record["element"], record["satellite"], int(record["patch_size"]))].append(record)274    rng = random.Random(seed)275    splits = {"train": [], "val": [], "test": []}276    for rows in grouped.values():277        rng.shuffle(rows)278        n = len(rows)279        n_train = int(n * train_ratio)280        n_val = int(n * val_ratio)281        if n >= 3:282            n_train = max(1, min(n_train, n - 2))283            n_val = max(1, min(n_val, n - n_train - 1))284        splits["train"].extend(rows[:n_train])285        splits["val"].extend(rows[n_train : n_train + n_val])286        splits["test"].extend(rows[n_train + n_val :])287    return splits288 289 290def copy_record(record: dict[str, Any], split: str, output_root: Path) -> dict[str, Any]:291    image_src = Path(record["image_path"])292    mask_src = Path(record["mask_path"])293    geojson_src = Path(record["annotation_path"])294    sample_id = record["sample_id"]295    image_dst = output_root / "images" / split / f"{sample_id}.jpg"296    mask_dst = output_root / "masks" / split / f"{sample_id}.png"297    annotation_dst = output_root / "annotations" / split / f"{sample_id}.geojson"298    image_dst.parent.mkdir(parents=True, exist_ok=True)299    mask_dst.parent.mkdir(parents=True, exist_ok=True)300    annotation_dst.parent.mkdir(parents=True, exist_ok=True)301    Image.open(image_src).convert("RGB").save(image_dst, quality=95)302    mask = Image.open(mask_src).convert("L").point(lambda v: 255 if v > 0 else 0)303    mask.save(mask_dst)304    shutil.copy2(geojson_src, annotation_dst)305    copied = dict(record)306    copied.update({"split": split, "image_path": str(image_dst), "mask_path": str(mask_dst), "annotation_path": str(annotation_dst)})307    return copied308 309 310def write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None:311    path.parent.mkdir(parents=True, exist_ok=True)312    path.write_text("\n".join(json.dumps(row, ensure_ascii=False) for row in rows) + ("\n" if rows else ""), encoding="utf-8")313 314 315def write_csv(path: Path, rows: list[dict[str, Any]]) -> None:316    path.parent.mkdir(parents=True, exist_ok=True)317    if not rows:318        path.write_text("", encoding="utf-8")319        return320    keys = sorted({key for row in rows for key in row.keys()})321    with path.open("w", newline="", encoding="utf-8-sig") as fp:322        writer = csv.DictWriter(fp, fieldnames=keys)323        writer.writeheader()324        writer.writerows(rows)325 326 327def make_preview(rows: list[dict[str, Any]], output_path: Path, limit: int) -> None:328    rows = rows[:limit]329    if not rows:330        return331    tile = 256332    cols = 6333    rows_n = (len(rows) + cols - 1) // cols334    canvas = Image.new("RGB", (cols * tile, rows_n * tile), (30, 30, 30))335    for idx, row in enumerate(rows):336        image = Image.open(row["image_path"]).convert("RGB").resize((tile, tile))337        mask = Image.open(row["mask_path"]).convert("L").resize((tile, tile))338        overlay = Image.new("RGBA", (tile, tile), (255, 60, 40, 0))339        overlay.putalpha(mask.point(lambda v: 100 if v > 0 else 0))340        image = Image.alpha_composite(image.convert("RGBA"), overlay).convert("RGB")341        draw = ImageDraw.Draw(image)342        draw.text((6, 6), f"{row['element']} {row['satellite']} S{row['patch_size']}", fill=(255, 255, 255))343        x = (idx % cols) * tile344        y = (idx // cols) * tile345        canvas.paste(image, (x, y))346    output_path.parent.mkdir(parents=True, exist_ok=True)347    canvas.save(output_path, quality=92)348 349 350def main() -> None:351    args = parse_args()352    source_root = Path(args.source_root)353    output_root = Path(args.output_root)354    if output_root.exists():355        shutil.rmtree(output_root)356    output_root.mkdir(parents=True, exist_ok=True)357    accepted, rejected, stats = build_records(source_root)358    splits = split_records(accepted, args.train_ratio, args.val_ratio, args.seed)359    copied: list[dict[str, Any]] = []360    for split, rows in splits.items():361        for row in rows:362            copied.append(copy_record(row, split, output_root))363 364    manifests = output_root / "manifests"365    write_jsonl(manifests / "accepted_polygon_samples.jsonl", copied)366    write_jsonl(manifests / "rejected_samples.jsonl", rejected)367    write_csv(output_root / "reports" / "rejected_samples.csv", rejected)368    write_csv(output_root / "reports" / "accepted_samples.csv", copied)369    make_preview(copied, output_root / "previews" / "mask_overlay_contact_sheet.jpg", args.preview_limit)370 371    by_element = Counter(row["element"] for row in copied)372    by_satellite = Counter(row["satellite"] for row in copied)373    by_patch_size = Counter(str(row["patch_size"]) for row in copied)374    summary = {375        **asdict(stats),376        "output_root": str(output_root),377        "splits": {split: len(rows) for split, rows in splits.items()},378        "accepted_by_element": dict(by_element),379        "accepted_by_satellite": dict(by_satellite),380        "accepted_by_patch_size": dict(by_patch_size),381        "dataset_layout": "images/{split}, masks/{split}, annotations/{split}, manifests, reports, previews",382        "quality_policy": "Only paired image + non-empty binary mask + valid GeoJSON polygon samples are accepted.",383        "source_root": str(source_root),384    }385    (output_root / "dataset_card.json").write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")386    print(json.dumps(summary, indent=2, ensure_ascii=False), flush=True)387 388 389if __name__ == "__main__":390    main()391