OneScience-Group/PrecipDD
024
1"""Create spatially coherent daily anomalies with weather and warming signals."""2 3import json4import sys5from pathlib import Path6 7import numpy as np8import torch9import torch.nn.functional as F10 11 12ROOT = Path(__file__).resolve().parents[1]13sys.path.insert(0, str(ROOT))14from model.precipdd import DATA_FORMAT_VERSION, load_config15 16 17def agmt_for_year(year, member, rng):18 forced = -0.45 + 0.0042 * (year - 1850) + 0.000030 * max(year - 1970, 0) ** 219 return forced + 0.08 * np.sin((year - 1850) / 8.0 + member) + rng.normal(0, 0.045)20 21 22def main():23 config = load_config(ROOT / "conf/config.yaml")24 settings = config["data"]25 rng = np.random.default_rng(config["project"]["seed"])26 train_n, val_n = settings["train_samples"], settings["validation_samples"]27 first, last = settings["test_years"]28 test_year = np.repeat(np.arange(first, last + 1), settings["test_days_per_year"])29 train_year = rng.integers(1850, 2101, train_n)30 val_year = rng.integers(1850, 2101, val_n)31 year = np.concatenate((train_year, val_year, test_year)).astype(np.int16)32 split = np.concatenate((np.zeros(train_n, np.int8), np.ones(val_n, np.int8), np.full(len(test_year), 2, np.int8)))33 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)34 member = rng.integers(0, settings["synthetic_members"], len(year), dtype=np.int16)35 agmt = np.array([agmt_for_year(int(y), int(m), rng) for y, m in zip(year, member)], dtype=np.float32)36 37 latitude = np.arange(-60.0, 77.5, 2.5, dtype=np.float32)38 base_longitude = np.arange(0.0, 360.0, 2.5, dtype=np.float32)39 longitude = np.arange(0.0, 400.0, 2.5, dtype=np.float32)40 lat2d, lon2d = np.meshgrid(latitude, base_longitude, indexing="ij")41 east_pacific = np.exp(-((lat2d / 16) ** 2 + ((lon2d - 245) / 35) ** 2))42 storm_tracks = np.exp(-((np.abs(lat2d) - 45) / 12) ** 2) * (0.65 + 0.35 * np.cos(np.deg2rad(lon2d * 2)))43 fingerprint = (east_pacific + storm_tracks).astype(np.float32)44 fingerprint /= fingerprint.max()45 coarse = torch.from_numpy(rng.normal(size=(len(year), 1, 14, 36)).astype(np.float32))46 weather = F.interpolate(coarse, size=(55, 144), mode="bilinear", align_corners=False).numpy()47 synoptic = np.sin(2 * np.pi * day[:, None, None] / 7.0 + np.deg2rad(lon2d)[None])48 noise = rng.normal(0, 0.28, size=weather.shape).astype(np.float32)49 amplitude = 1.0 + 0.34 * np.maximum(agmt, -0.5)[:, None, None] * fingerprint[None]50 fields = weather[:, 0] * amplitude + 0.32 * synoptic * (0.3 + fingerprint[None]) + noise[:, 0]51 fields += 0.22 * agmt[:, None, None] * (fingerprint[None] - 0.35)52 fields = fields.astype(np.float32)53 fields = np.concatenate((fields, fields[:, :, :16]), axis=2)54 fingerprint = np.concatenate((fingerprint, fingerprint[:, :16]), axis=1)55 fields -= fields.mean(axis=2, keepdims=True)56 zonal_std = fields.std(axis=2, keepdims=True).mean(axis=0, keepdims=True)57 fields /= np.maximum(zonal_std, 1e-5)58 59 output = ROOT / config["paths"]["data"]60 output.parent.mkdir(parents=True, exist_ok=True)61 np.savez_compressed(output, format_version=np.array(DATA_FORMAT_VERSION), precipitation=fields[:, None], agmt=agmt,62 split=split, year=year, day_of_year=day, member=member, latitude=latitude,63 longitude=longitude, synthetic=np.array(True), fingerprint=fingerprint)64 metadata = {"synthetic": True, "samples": len(year), "splits": {"train": train_n, "validation": val_n, "test": len(test_year)},65 "precipitation_shape": [len(year), 1, 55, 160], "target_shape": [len(year)],66 "science_note": "Structured engineering data with daily weather variability and an AGMT-dependent spatial variance signal; not CESM2 LE."}67 (output.parent / "metadata.json").write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8")68 print(f"data={output.relative_to(ROOT)} shape={fields[:, None].shape} target={agmt.shape}")69 70 71if __name__ == "__main__":72 main()73 