fnruha0921's picture
add code
eea5f0e verified
Raw History Blame Contribute Delete
3.08 kB
"""Resample HF single-date datasets to ~0.55-0.65 m/px and pack into npy arrays.
FLAIR (0.2 m, 512^2) -> 192^2 (~0.53 m/px). classes: 1 building, 18 greenhouse -> building;
6 conifer, 7 deciduous, 16 mixed, 17 ligneous -> tree; 2,4,10,11,12,15 -> ground donors.
OAM-TCD (0.1 m, 2048^2) -> 320^2 (0.64 m/px). annotation > 0 -> tree.
"""
import glob
import io
import os
from concurrent.futures import ProcessPoolExecutor
import cv2
import numpy as np
import pyarrow.parquet as pq
import tifffile
from PIL import Image
OUT = "/workspace/prep"
os.makedirs(OUT, exist_ok=True)
BUILD = [1, 18]
TREE = [6, 7, 16, 17]
GROUND = [2, 4, 10, 11, 12, 15]
def flair_one(img_path):
msk_path = img_path.replace("/aerial/", "/labels/").replace("IMG_", "MSK_")
if not os.path.exists(msk_path):
cand = glob.glob(os.path.join(os.path.dirname(img_path).replace("img", "msk"), "MSK_" + img_path.split("IMG_")[-1]))
if not cand:
return None
msk_path = cand[0]
im = tifffile.imread(img_path)
if im.shape[0] in (4, 5):
im = np.transpose(im, (1, 2, 0))
rgb = np.ascontiguousarray(im[..., :3]).astype(np.uint8)
m = tifffile.imread(msk_path)
if m.ndim == 3:
m = m[0] if m.shape[0] < m.shape[-1] else m[..., 0]
rgb = cv2.resize(rgb, (192, 192), interpolation=cv2.INTER_AREA)
lab = np.zeros(m.shape, np.uint8)
lab[np.isin(m, TREE)] = 2
lab[np.isin(m, GROUND)] = 3
lab[np.isin(m, BUILD)] = 1
lab = cv2.resize(lab, (192, 192), interpolation=cv2.INTER_NEAREST)
return rgb, lab
def main():
imgs = sorted(glob.glob("/workspace/data/flair/unz/**/IMG_*.tif", recursive=True))
print("flair imgs", len(imgs), flush=True)
with ProcessPoolExecutor(48) as ex:
res = [r for r in ex.map(flair_one, imgs, chunksize=32) if r is not None]
fx = np.stack([r[0] for r in res]); fy = np.stack([r[1] for r in res])
np.save(f"{OUT}/flair_x.npy", fx); np.save(f"{OUT}/flair_y.npy", fy)
print("flair", fx.shape, "building frac", (fy == 1).mean(), "tree frac", (fy == 2).mean(), flush=True)
xs, ys = [], []
for f in sorted(glob.glob("/workspace/data/tcd/data/train-*.parquet")):
t = pq.read_table(f, columns=["image", "annotation"]).to_pylist()
for row in t:
im = np.asarray(Image.open(io.BytesIO(row["image"]["bytes"])).convert("RGB"))
an = np.asarray(Image.open(io.BytesIO(row["annotation"]["bytes"])))
if an.ndim == 3:
an = an[..., 0]
s = 320 * max(im.shape[:2]) // 2048
xs.append(cv2.resize(im, (s, s), interpolation=cv2.INTER_AREA) if s == 320 else cv2.resize(im, (320, 320), interpolation=cv2.INTER_AREA))
ys.append(cv2.resize(((an > 0) * 2).astype(np.uint8), (320, 320), interpolation=cv2.INTER_NEAREST))
print(f, len(xs), flush=True)
tx = np.stack(xs); ty = np.stack(ys)
np.save(f"{OUT}/tcd_x.npy", tx); np.save(f"{OUT}/tcd_y.npy", ty)
print("tcd", tx.shape, "tree frac", (ty == 2).mean(), flush=True)
if __name__ == "__main__":
main()