fnruha0921's picture
anch30 + earthview runs
2b046b1 verified
Raw History Blame Contribute Delete
2.39 kB
"""EarthView (Satellogic 1 m, CC BY 4.0): build real same-place / different-date RGB pairs as 'no change' negatives.
Downloads N shards, keeps locations with >= 2 revisits, stores up to 3 revisit pairs per location as uint8
(K, 2, H, W, 3) at native 1 m. The trainer crops ~115 px and resizes to 192 (~0.6 m/px).
usage: python ev_prep.py N_SHARDS
"""
import os
import sys
from concurrent.futures import ThreadPoolExecutor
import numpy as np
import pyarrow.parquet as pq
from huggingface_hub import hf_hub_download
N = int(sys.argv[1]) if len(sys.argv) > 1 else 20
OUT = "/workspace/prep/ev_imgs.npy"
def to_rgb8(a):
"""(bands,H,W) raw -> (H,W,3) uint8 with a per-image percentile stretch (display-like)."""
a = np.asarray(a, np.float32)[:3].transpose(1, 2, 0)
lo, hi = np.percentile(a, 1), np.percentile(a, 99)
return np.clip((a - lo) * 255.0 / max(hi - lo, 1e-6), 0, 255).astype(np.uint8)
def shard(i):
p = hf_hub_download("satellogic/EarthView", f"satellogic/train-{i:05d}-of-07863.parquet", repo_type="dataset",
local_dir="/workspace/data/ev")
rng = np.random.default_rng(i)
out = []
col = pq.read_table(p, columns=["rgb"])["rgb"].combine_chunks()
nrev = np.diff(col.offsets.to_numpy())
vals = col.flatten().flatten().flatten().flatten().to_numpy(zero_copy_only=False).astype(np.uint8)
per = vals.size // max(nrev.sum(), 1)
side = int(round((per / 3) ** 0.5))
imgs = vals.reshape(-1, 3, side, side)
start = np.concatenate([[0], np.cumsum(nrev)])
for j in range(len(nrev)): # single-date tiles: keep first revisit
x = to_rgb8(imgs[start[j]])
if (x.max(2) > 0).mean() > 0.95 and x.std() > 8:
out.append(x)
print("shard", i, "rows", len(nrev), "multi-revisit", int((nrev >= 2).sum()), "max rev", int(nrev.max()), flush=True)
os.remove(p)
print("shard", i, len(out), flush=True)
return out
if __name__ == "__main__":
idx = np.random.default_rng(0).choice(7863, N, replace=False)
with ThreadPoolExecutor(6) as ex:
res = [r for rs in ex.map(shard, idx) for r in rs]
shapes = {}
for r in res:
shapes[r.shape] = shapes.get(r.shape, 0) + 1
s = max(shapes, key=shapes.get)
arr = np.stack([r for r in res if r.shape == s])
np.save(OUT, arr)
print("EV_DONE", arr.shape, shapes, flush=True)