File size: 3,933 Bytes
53becf5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Generate deterministic multi-sensor chips for the Clay engineering workflow."""

import argparse
from pathlib import Path

import numpy as np
import yaml


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


def load_config():
    return yaml.safe_load((ROOT / "conf/config.yaml").read_text())


def make_split(path, count, config, seed):
    rng = np.random.default_rng(seed)
    size = int(config["data"]["image_size"])
    y, x = np.mgrid[-1:1:complex(size), -1:1:complex(size)].astype(np.float32)
    payload = {
        "format_version": np.asarray(config["data"]["format_version"]),
        "data_source": np.asarray("synthetic"),
        "time": np.empty((count, 2), np.float32),
        "latlon": np.empty((count, 2), np.float32),
        "class_target": np.empty(count, np.int64),
        "regression_target": np.empty(count, np.float32),
        "teacher_target": np.empty((count, config["model"]["teacher_dim"]), np.float32),
    }
    projection = rng.normal(size=(6, config["model"]["teacher_dim"])).astype(np.float32)
    for sample in range(count):
        phase = rng.uniform(0, 2 * np.pi)
        week = rng.uniform(0, 2 * np.pi)
        hour = rng.uniform(0, 2 * np.pi)
        latitude = rng.uniform(-math_pi_over_two(), math_pi_over_two())
        longitude = rng.uniform(-np.pi, np.pi)
        payload["time"][sample] = (week, hour)
        payload["latlon"][sample] = (latitude, longitude)
        payload["class_target"][sample] = int(np.sin(phase) > 0)
        payload["regression_target"][sample] = np.cos(phase) + 0.2 * np.sin(latitude)
        descriptor = np.asarray([
            np.sin(phase), np.cos(phase), np.sin(week), np.cos(week),
            np.sin(latitude), np.cos(longitude),
        ], np.float32)
        target = descriptor @ projection
        payload["teacher_target"][sample] = target / np.linalg.norm(target).clip(1e-6)

    for sensor_index, (name, spec) in enumerate(config["data"]["sensors"].items()):
        channels = int(spec["channels"])
        pixels = np.empty((count, channels, size, size), np.float32)
        valid = np.ones((count, 1, size, size), np.float32)
        for sample in range(count):
            phase = np.arctan2(payload["teacher_target"][sample, 0], payload["teacher_target"][sample, 1])
            base = np.sin((2.0 + sensor_index) * np.pi * x + phase) * np.cos(2 * np.pi * y - phase)
            base += 0.35 * x + 0.15 * y
            for channel in range(channels):
                spectral = 0.12 * channel + 0.2 * np.sin(phase + channel / max(channels, 1))
                pixels[sample, channel] = base + spectral + rng.normal(0, 0.03, (size, size))
            valid[sample, :, :2] = 0
        pixels -= pixels.mean(axis=(0, 2, 3), keepdims=True)
        pixels /= pixels.std(axis=(0, 2, 3), keepdims=True).clip(1e-6)
        payload[f"pixels_{name}"] = np.clip(pixels, -6, 6).astype(np.float32)
        payload[f"valid_{name}"] = valid
        payload[f"wavelengths_{name}"] = np.tile(np.asarray(spec["wavelengths_nm"], np.float32), (count, 1))
        payload[f"gsd_{name}"] = np.full(count, float(spec["gsd"]), np.float32)
    np.savez_compressed(path, **payload)


def math_pi_over_two():
    return np.pi / 2


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)
    for offset, (filename, count) in enumerate((
        ("train.npz", config["data"]["train_samples"]),
        ("test.npz", config["data"]["test_samples"]),
    )):
        path = data_dir / filename
        if args.force or not path.exists():
            make_split(path, int(count), config, int(config["seed"]) + offset)
        print(f"generated={path.relative_to(ROOT)} samples={count} format={config['data']['format_version']}")


if __name__ == "__main__":
    main()