PrecipDD / scripts /fake_data.py
zhangrenchao's picture
Publish PrecipDD reproduction
950fc23 verified
Raw
History Blame Contribute Delete
4.1 kB
"""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()