Team Ai
Modelpublic

OneScience-Group/PrecipDD

sourceHugging Faceapache-2.0updated 24d agoView on Hugging Face
0likes24downloads
fake_data.py73 linesDownload Raw Back to scripts
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