cuibinge's picture
Sync YOLO training and evaluation utilities (part 2)
c2b1b26 verified
Raw
History Blame Contribute Delete
4.48 kB
"""Build a train/val dataset from large images and prediction rasters.
This is a bootstrap utility. Masks created from previous predictions are
pseudo-labels, not human-verified ground truth.
"""
from __future__ import annotations
import argparse
import random
from pathlib import Path
import numpy as np
import rasterio
from rasterio.windows import Window
def find_pairs(source_dir: Path):
images = [p for p in source_dir.glob("*.tif") if "_prediction" not in p.stem.lower()]
pairs = []
for image in images:
pred = None
for candidate in source_dir.glob(f"{image.stem}_*/{image.stem}_prediction.tif"):
pred = candidate
break
if pred:
pairs.append((image, pred))
return pairs
def ensure_layout(output_dir: Path):
for split in ("train", "val"):
(output_dir / split / "images").mkdir(parents=True, exist_ok=True)
(output_dir / split / "masks").mkdir(parents=True, exist_ok=True)
def write_tile(src, mask_src, window: Window, image_path: Path, mask_path: Path, foreground_threshold: int):
image = src.read(window=window)
mask = mask_src.read(1, window=window)
if image.shape[1] != window.height or image.shape[2] != window.width:
return False
if mask.shape[0] != window.height or mask.shape[1] != window.width:
return False
image_meta = src.meta.copy()
image_meta.update(
{
"height": int(window.height),
"width": int(window.width),
"transform": src.window_transform(window),
"compress": "lzw",
}
)
mask_meta = mask_src.meta.copy()
mask_meta.update(
{
"count": 1,
"dtype": "uint8",
"height": int(window.height),
"width": int(window.width),
"transform": mask_src.window_transform(window),
"compress": "lzw",
}
)
binary_mask = (mask >= foreground_threshold).astype(np.uint8) * 255
with rasterio.open(image_path, "w", **image_meta) as dst:
dst.write(image)
with rasterio.open(mask_path, "w", **mask_meta) as dst:
dst.write(binary_mask, 1)
return True
def build_dataset(source_dir: Path, output_dir: Path, tile_size: int, stride: int, val_ratio: float, foreground_threshold: int):
pairs = find_pairs(source_dir)
if not pairs:
raise RuntimeError(f"No image/prediction pairs found under {source_dir}")
ensure_layout(output_dir)
rng = random.Random(42)
written = {"train": 0, "val": 0}
for pair_index, (image_path, pred_path) in enumerate(pairs):
with rasterio.open(image_path) as src, rasterio.open(pred_path) as mask_src:
windows = []
for y in range(0, src.height - tile_size + 1, stride):
for x in range(0, src.width - tile_size + 1, stride):
windows.append(Window(x, y, tile_size, tile_size))
rng.shuffle(windows)
for tile_index, window in enumerate(windows):
split = "val" if rng.random() < val_ratio else "train"
name = f"pair{pair_index:02d}_{tile_index:06d}.tif"
ok = write_tile(
src,
mask_src,
window,
output_dir / split / "images" / name,
output_dir / split / "masks" / name,
foreground_threshold,
)
if ok:
written[split] += 1
print(f"Wrote pseudo dataset to {output_dir}")
print(f"train tiles: {written['train']}")
print(f"val tiles: {written['val']}")
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--source-dir", default="data")
parser.add_argument("--output-dir", default="data_pseudo")
parser.add_argument("--tile-size", type=int, default=256)
parser.add_argument("--stride", type=int, default=256)
parser.add_argument("--val-ratio", type=float, default=0.15)
parser.add_argument("--foreground-threshold", type=int, default=1)
args = parser.parse_args()
build_dataset(
Path(args.source_dir),
Path(args.output_dir),
args.tile_size,
args.stride,
args.val_ratio,
args.foreground_threshold,
)
if __name__ == "__main__":
main()