cuibinge's picture
Sync YOLO training and evaluation utilities (part 2)
c2b1b26 verified
Raw
History Blame Contribute Delete
10.8 kB
"""Build a normalized manifest for marine ecological feature datasets.
The script scans one or more roots, classifies likely feature types from path
keywords, pairs images and masks when possible, and writes a manifest without
copying large raster files.
"""
from __future__ import annotations
import argparse
import csv
import json
import re
from dataclasses import asdict, dataclass
from datetime import datetime
from pathlib import Path
from typing import Iterable
IMAGE_EXTS = {".tif", ".tiff", ".png", ".jpg", ".jpeg"}
MASK_HINTS = ("mask", "masks", "label", "labels", "gt", "annotation", "annotations", "seg")
IMAGE_HINTS = ("image", "images", "img", "imgs", "tif", "tile", "tiles")
ELEMENT_KEYWORDS = {
"green_tide": ("浒苔", "绿潮", "green_tide", "greentide", "entgreentide", "enteromorpha", "seaweed"),
"red_tide": ("赤潮", "red_tide", "redtide", "harmful_algal", "hab"),
"golden_tide": ("马尾藻", "金潮", "sarg", "sargassum", "golden_tide", "goldentide"),
"aquaculture": ("养殖", "aquaculture", "raft", "cage", "pond"),
}
SATELLITE_PATTERN = re.compile(r"\b(GF\d+|HY\d+|Sentinel-?2|Landsat-?\d*)\b", re.IGNORECASE)
DATE_PATTERN = re.compile(r"(20\d{6}|19\d{6})")
PATCH_SIZE_PATTERN = re.compile(r"(?:^|[_\\/\-])(?:size)?(128|256|512|1024)(?:[_\\/\-]|$)")
@dataclass
class AssetRecord:
asset_id: str
path: str
filename: str
suffix: str
role: str
element: str
satellite: str | None
sensor: str | None
acquired_at: str | None
patch_size: int | None
source_project: str
source_dataset: str
size_bytes: int
modified_at: str
quality_flags: list[str]
@dataclass
class SampleRecord:
sample_id: str
element: str
task_type: str
image_path: str
mask_path: str | None
label_encoding: dict[str, str] | None
satellite: str | None
sensor: str | None
resolution_m: float | None
patch_size: int | None
bands: list[str] | None
band_count: int | None
dtype: str | None
fusion: dict
acquired_at: str | None
source_project: str
source_dataset: str
split: str | None
quality_flags: list[str]
notes: str
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--roots", nargs="+", required=True, help="Dataset roots to scan.")
parser.add_argument("--output-root", required=True, help="Output normalized dataset root.")
parser.add_argument("--max-files", type=int, default=0, help="Optional scan limit for debugging.")
return parser.parse_args()
def norm_text(path: Path) -> str:
return str(path).replace("\\", "/").lower()
def infer_element(path: Path) -> str:
text = norm_text(path)
for element, keywords in ELEMENT_KEYWORDS.items():
if any(keyword.lower() in text for keyword in keywords):
return element
return "unknown"
def infer_role(path: Path) -> str:
parts = [part.lower() for part in path.parts]
stem = path.stem.lower()
if any(hint in parts or hint in stem for hint in MASK_HINTS):
return "mask"
if any(hint in parts for hint in IMAGE_HINTS):
return "image"
if path.suffix.lower() in {".png", ".jpg", ".jpeg"} and any(hint in stem for hint in MASK_HINTS):
return "mask"
return "image"
def infer_satellite(path: Path) -> str | None:
match = SATELLITE_PATTERN.search(str(path))
return match.group(1).upper().replace("-", "") if match else None
def infer_sensor(path: Path) -> str | None:
upper = path.name.upper()
for sensor in ("PMS", "MUX", "MSS", "PAN", "WFV"):
if sensor in upper:
return sensor
return None
def infer_date(path: Path) -> str | None:
match = DATE_PATTERN.search(path.name)
if not match:
return None
raw = match.group(1)
try:
return datetime.strptime(raw, "%Y%m%d").date().isoformat()
except ValueError:
return None
def infer_patch_size(path: Path) -> int | None:
match = PATCH_SIZE_PATTERN.search(str(path))
return int(match.group(1)) if match else None
def infer_source_project(path: Path, roots: list[Path]) -> str:
for root in roots:
try:
rel = path.relative_to(root)
except ValueError:
continue
return rel.parts[0] if len(rel.parts) > 1 else root.name
return path.parent.name
def infer_split(path: Path) -> str | None:
parts = {part.lower() for part in path.parts}
for split in ("train", "val", "test"):
if split in parts:
return split
return None
def source_dataset(path: Path) -> str:
for part in reversed(path.parts):
lower = part.lower()
if any(token in lower for token in ("gf", "sentinel", "landsat", "浒苔", "赤潮", "马尾藻", "养殖")):
return part
return path.parent.name
def is_fused(path: Path) -> bool | None:
lower = path.name.lower()
if "fuse" in lower or "fusion" in lower or "pan" not in lower and "mux" in lower:
return True
if "pan" in lower or "mss" in lower:
return False
return None
def infer_fusion(path: Path) -> dict:
fused = is_fused(path)
lower = path.name.lower()
if fused is True:
state = "fused_product"
method = "unknown_vendor_product"
persisted = True
elif fused is False and ("pan" in lower or "mss" in lower):
state = "none"
method = "none"
persisted = False
else:
state = "unknown"
method = "unknown"
persisted = False
return {
"state": state,
"method": method,
"sources": [{"role": "source", "path": str(path), "resolution_m": None}],
"target_resolution_m": None,
"native_multispectral_resolution_m": None,
"persisted": persisted,
"reproducible": state != "unknown",
"spectral_preservation": "unknown",
"notes": "Auto-inferred from local filename; verify before training.",
}
def asset_id(path: Path) -> str:
safe = re.sub(r"[^A-Za-z0-9]+", "_", str(path.stem)).strip("_").lower()
return safe[:180]
def iter_files(roots: Iterable[Path], max_files: int) -> Iterable[Path]:
count = 0
for root in roots:
if not root.exists():
continue
for path in root.rglob("*"):
if not path.is_file() or path.suffix.lower() not in IMAGE_EXTS:
continue
yield path
count += 1
if max_files and count >= max_files:
return
def write_jsonl(path: Path, rows: Iterable[dict]) -> None:
with path.open("w", encoding="utf-8") as f:
for row in rows:
f.write(json.dumps(row, ensure_ascii=False) + "\n")
def main() -> None:
args = parse_args()
roots = [Path(root) for root in args.roots]
output_root = Path(args.output_root)
manifest_dir = output_root / "manifests"
report_dir = output_root / "reports"
manifest_dir.mkdir(parents=True, exist_ok=True)
report_dir.mkdir(parents=True, exist_ok=True)
assets: list[AssetRecord] = []
for path in iter_files(roots, args.max_files):
stat = path.stat()
role = infer_role(path)
flags = []
if infer_element(path) == "unknown":
flags.append("unknown_element")
if role == "image" and "black" in norm_text(path):
flags.append("possibly_invalid")
assets.append(
AssetRecord(
asset_id=asset_id(path),
path=str(path),
filename=path.name,
suffix=path.suffix.lower(),
role=role,
element=infer_element(path),
satellite=infer_satellite(path),
sensor=infer_sensor(path),
acquired_at=infer_date(path),
patch_size=infer_patch_size(path),
source_project=infer_source_project(path, roots),
source_dataset=source_dataset(path),
size_bytes=stat.st_size,
modified_at=datetime.fromtimestamp(stat.st_mtime).isoformat(timespec="seconds"),
quality_flags=flags,
)
)
masks_by_stem = {Path(asset.path).stem.lower(): asset for asset in assets if asset.role == "mask"}
samples: list[SampleRecord] = []
for asset in assets:
if asset.role != "image":
continue
path = Path(asset.path)
mask_asset = masks_by_stem.get(path.stem.lower())
element = asset.element if asset.element != "unknown" else (mask_asset.element if mask_asset else "unknown")
flags = list(asset.quality_flags)
if mask_asset is None:
flags.append("unpaired_image")
sample_id = f"{element}_{asset.asset_id}"
samples.append(
SampleRecord(
sample_id=sample_id,
element=element,
task_type="semantic_segmentation",
image_path=asset.path,
mask_path=mask_asset.path if mask_asset else None,
label_encoding={"0": "background", "1": element} if mask_asset else None,
satellite=asset.satellite,
sensor=asset.sensor,
resolution_m=None,
patch_size=asset.patch_size,
bands=None,
band_count=None,
dtype=None,
fusion=infer_fusion(path),
acquired_at=asset.acquired_at,
source_project=asset.source_project,
source_dataset=asset.source_dataset,
split=infer_split(path),
quality_flags=flags,
notes="auto-generated; verify ambiguous labels before training",
)
)
write_jsonl(manifest_dir / "assets_raw.jsonl", (asdict(asset) for asset in assets))
write_jsonl(manifest_dir / "samples.jsonl", (asdict(sample) for sample in samples))
with (report_dir / "asset_inventory.csv").open("w", newline="", encoding="utf-8-sig") as f:
writer = csv.DictWriter(f, fieldnames=list(asdict(assets[0]).keys()) if assets else ["asset_id"])
writer.writeheader()
for asset in assets:
row = asdict(asset)
row["quality_flags"] = ";".join(row["quality_flags"])
writer.writerow(row)
summary = {
"roots": [str(root) for root in roots],
"assets": len(assets),
"samples": len(samples),
"by_element": {},
"output_root": str(output_root),
}
for sample in samples:
summary["by_element"][sample.element] = summary["by_element"].get(sample.element, 0) + 1
print(json.dumps(summary, indent=2, ensure_ascii=False))
if __name__ == "__main__":
main()