File size: 4,100 Bytes
950fc23
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Create spatially coherent daily anomalies with weather and warming signals."""

import json
import sys
from pathlib import Path

import numpy as np
import torch
import torch.nn.functional as F


ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.precipdd import DATA_FORMAT_VERSION, load_config


def agmt_for_year(year, member, rng):
    forced = -0.45 + 0.0042 * (year - 1850) + 0.000030 * max(year - 1970, 0) ** 2
    return forced + 0.08 * np.sin((year - 1850) / 8.0 + member) + rng.normal(0, 0.045)


def main():
    config = load_config(ROOT / "conf/config.yaml")
    settings = config["data"]
    rng = np.random.default_rng(config["project"]["seed"])
    train_n, val_n = settings["train_samples"], settings["validation_samples"]
    first, last = settings["test_years"]
    test_year = np.repeat(np.arange(first, last + 1), settings["test_days_per_year"])
    train_year = rng.integers(1850, 2101, train_n)
    val_year = rng.integers(1850, 2101, val_n)
    year = np.concatenate((train_year, val_year, test_year)).astype(np.int16)
    split = np.concatenate((np.zeros(train_n, np.int8), np.ones(val_n, np.int8), np.full(len(test_year), 2, np.int8)))
    day = np.concatenate((rng.integers(1, 366, train_n + val_n), np.tile(np.linspace(15, 350, settings["test_days_per_year"], dtype=int), last - first + 1))).astype(np.int16)
    member = rng.integers(0, settings["synthetic_members"], len(year), dtype=np.int16)
    agmt = np.array([agmt_for_year(int(y), int(m), rng) for y, m in zip(year, member)], dtype=np.float32)

    latitude = np.arange(-60.0, 77.5, 2.5, dtype=np.float32)
    base_longitude = np.arange(0.0, 360.0, 2.5, dtype=np.float32)
    longitude = np.arange(0.0, 400.0, 2.5, dtype=np.float32)
    lat2d, lon2d = np.meshgrid(latitude, base_longitude, indexing="ij")
    east_pacific = np.exp(-((lat2d / 16) ** 2 + ((lon2d - 245) / 35) ** 2))
    storm_tracks = np.exp(-((np.abs(lat2d) - 45) / 12) ** 2) * (0.65 + 0.35 * np.cos(np.deg2rad(lon2d * 2)))
    fingerprint = (east_pacific + storm_tracks).astype(np.float32)
    fingerprint /= fingerprint.max()
    coarse = torch.from_numpy(rng.normal(size=(len(year), 1, 14, 36)).astype(np.float32))
    weather = F.interpolate(coarse, size=(55, 144), mode="bilinear", align_corners=False).numpy()
    synoptic = np.sin(2 * np.pi * day[:, None, None] / 7.0 + np.deg2rad(lon2d)[None])
    noise = rng.normal(0, 0.28, size=weather.shape).astype(np.float32)
    amplitude = 1.0 + 0.34 * np.maximum(agmt, -0.5)[:, None, None] * fingerprint[None]
    fields = weather[:, 0] * amplitude + 0.32 * synoptic * (0.3 + fingerprint[None]) + noise[:, 0]
    fields += 0.22 * agmt[:, None, None] * (fingerprint[None] - 0.35)
    fields = fields.astype(np.float32)
    fields = np.concatenate((fields, fields[:, :, :16]), axis=2)
    fingerprint = np.concatenate((fingerprint, fingerprint[:, :16]), axis=1)
    fields -= fields.mean(axis=2, keepdims=True)
    zonal_std = fields.std(axis=2, keepdims=True).mean(axis=0, keepdims=True)
    fields /= np.maximum(zonal_std, 1e-5)

    output = ROOT / config["paths"]["data"]
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output, format_version=np.array(DATA_FORMAT_VERSION), precipitation=fields[:, None], agmt=agmt,
                        split=split, year=year, day_of_year=day, member=member, latitude=latitude,
                        longitude=longitude, synthetic=np.array(True), fingerprint=fingerprint)
    metadata = {"synthetic": True, "samples": len(year), "splits": {"train": train_n, "validation": val_n, "test": len(test_year)},
                "precipitation_shape": [len(year), 1, 55, 160], "target_shape": [len(year)],
                "science_note": "Structured engineering data with daily weather variability and an AGMT-dependent spatial variance signal; not CESM2 LE."}
    (output.parent / "metadata.json").write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8")
    print(f"data={output.relative_to(ROOT)} shape={fields[:, None].shape} target={agmt.shape}")


if __name__ == "__main__":
    main()