cuibinge/typical-marine-ecological-feature-recognition-code
0
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 