File size: 4,147 Bytes
1558db5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2d7eead
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1558db5
 
 
2d7eead
 
 
 
1558db5
 
 
 
 
 
 
 
 
 
 
 
 
 
2d7eead
 
 
1558db5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""Generate deterministic SkySense-format data for connectivity tests."""

import argparse
from pathlib import Path

import numpy as np
import yaml


ROOT = Path(__file__).resolve().parents[1]


def load_config():
    with (ROOT / "conf" / "config.yaml").open(encoding="utf-8") as handle:
        return yaml.safe_load(handle)


def make_split(path, count, config, seed):
    rng = np.random.default_rng(seed)
    data = config["data"]
    hr_size = data["hr_size"]
    low_size = data["s2_size"]
    if low_size != data["s1_size"]:
        raise ValueError("The fake generator expects equal S1/S2 grid sizes")
    y, x = np.mgrid[0:hr_size, 0:hr_size].astype(np.float32) / (hr_size - 1)
    low_y, low_x = np.mgrid[0:low_size, 0:low_size].astype(np.float32) / (low_size - 1)
    land = np.stack([x, y, np.sin(4 * np.pi * x) * np.cos(3 * np.pi * y)], axis=0)
    hr = np.empty((count, data["hr_timesteps"], data["hr_channels"], hr_size, hr_size), np.float32)
    s2 = np.empty((count, data["s2_timesteps"], data["s2_channels"], low_size, low_size), np.float32)
    s1 = np.empty((count, data["s1_timesteps"], data["s1_channels"], low_size, low_size), np.float32)
    for sample in range(count):
        phase = rng.uniform(0, 2 * np.pi)
        hr[sample, 0] = land + rng.normal(0, 0.04, land.shape)
        for time in range(data["s2_timesteps"]):
            seasonal = np.sin(2 * np.pi * time / data["s2_timesteps"] + phase)
            s2[sample, time] = np.stack([
                0.5 + 0.2 * np.sin((band + 1) * low_x + phase) + 0.1 * seasonal * low_y
                for band in range(data["s2_channels"])
            ]) + rng.normal(0, 0.025, (data["s2_channels"], low_size, low_size))
        for time in range(data["s1_timesteps"]):
            seasonal = np.cos(2 * np.pi * time / data["s1_timesteps"] + phase)
            s1[sample, time, 0] = -0.6 + 0.3 * np.sin(4 * np.pi * low_x) * np.cos(3 * np.pi * low_y) + 0.1 * seasonal
            s1[sample, time, 1] = -0.8 + 0.2 * low_y - 0.1 * seasonal
        hr[sample] = np.clip(hr[sample], -1, 1)
        s2[sample] = np.clip(s2[sample], 0, 1)
    dates_hr = rng.integers(0, 365, size=(count, data["hr_timesteps"]), dtype=np.int64)
    dates_s2 = rng.integers(0, 365, size=(count, data["s2_timesteps"]), dtype=np.int64)
    dates_s1 = rng.integers(0, 365, size=(count, data["s1_timesteps"]), dtype=np.int64)
    region = rng.integers(0, config["model"]["num_regions"], size=count, dtype=np.int64)
    indices = np.linspace(0, low_size - 1, hr_size).round().astype(int)
    s2_hr = s2[:, :, 3].mean(axis=1)[:, indices][:, :, indices]
    s1_hr = s1[:, :, 0].mean(axis=1)[:, indices][:, :, indices]
    signal = hr[:, 0, 0] + 0.35 * s2_hr - 0.2 * s1_hr
    bins = np.quantile(signal, np.linspace(0, 1, data["num_classes"] + 1)[1:-1])
    labels = np.digitize(signal, bins).astype(np.int64)
    np.savez_compressed(
        path,
        hr=hr,
        s2=s2,
        s1=s1,
        dates_hr=dates_hr,
        dates_s2=dates_s2,
        dates_s1=dates_s1,
        region=region,
        labels=labels,
        data_source=np.asarray("synthetic"),
        protocol=np.asarray(config["data"]["protocol"]),
        band_order_hr=np.asarray(["R", "G", "B"]),
        band_order_s2=np.asarray(["B2", "B3", "B4", "B5", "B6", "B7", "B8", "B8A", "B11", "B12"]),
        band_order_s1=np.asarray(["VV", "VH"]),
    )


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--force", action="store_true")
    args = parser.parse_args()
    config = load_config()
    data_dir = ROOT / config["data"]["root"]
    data_dir.mkdir(parents=True, exist_ok=True)
    outputs = [("train.npz", config["data"]["train_samples"]), ("test.npz", config["data"]["test_samples"])]
    for offset, (name, count) in enumerate(outputs):
        target = data_dir / name
        if args.force or not target.exists():
            make_split(target, count, config, config["seed"] + offset)
        print(
            f"generated={target.relative_to(ROOT)} samples={count} "
            f"data_source=synthetic protocol={config['data']['protocol']}"
        )


if __name__ == "__main__":
    main()