| """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() |
|
|