File size: 2,851 Bytes
659cfca
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
"""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)])
    pairs = []
    for j in range(len(nrev)):
        x = to_rgb8(imgs[start[j]])
        if (x.max(2) > 0).mean() > 0.95 and x.std() > 8:
            out.append(x)
        R = nrev[j]
        for _ in range(min(2, R - 1)):                   # real same-place / different-date pairs
            a, b = rng.choice(R, 2, replace=False)
            u, v = to_rgb8(imgs[start[j] + a]), to_rgb8(imgs[start[j] + b])
            if (u.max(2) > 0).mean() > 0.95 and (v.max(2) > 0).mean() > 0.95:
                pairs.append(np.stack([u, v]))
    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, pairs


if __name__ == "__main__":
    idx = np.random.default_rng(1).choice(7863, N, replace=False)
    with ThreadPoolExecutor(8) as ex:
        res = list(ex.map(shard, idx))
    singles = np.stack([x for o, _ in res for x in o if x.shape == (384, 384, 3)])
    pairs = [p for _, ps in res for p in ps if p.shape == (2, 384, 384, 3)]
    np.save("/workspace/prep/ev_imgs.npy", singles)
    np.save("/workspace/prep/ev_pairs.npy", np.stack(pairs) if pairs else np.zeros((0, 2, 384, 384, 3), np.uint8))
    print("EV_DONE singles", singles.shape, "pairs", len(pairs), flush=True)