Download scripts/fake_data.py from OneScience-Group/PPNN: direct link, hf CLI and curl.
- Browser
- Download file 3.74 kB
-
https://huggingface.co/OneScience-Group/PPNN/resolve/main/scripts/fake_data.py
- Command line
-
hf download hf://OneScience-Group/PPNN/scripts/fake_data.py
-
curl -L -o fake_data.py https://huggingface.co/OneScience-Group/PPNN/resolve/main/scripts/fake_data.py
3.74 kB
| """Generate a small, structured PPNN dataset while preserving paper dimensions.""" | |
| from pathlib import Path | |
| import numpy as np | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| data = config["data"] | |
| if (data["station_count"], data["member_count"], data["variable_count"], data["lead_hours"]) != (537, 50, 18, 48): | |
| raise ValueError("paper station/member/variable/lead dimensions cannot be reduced") | |
| rng = np.random.default_rng(int(config["seed"])) | |
| station_index_all = np.arange(537) | |
| latitude = (47.0 + 8.0 * ((station_index_all * 0.61803398875) % 1)).astype(np.float32) | |
| longitude = (5.5 + 9.5 * ((station_index_all * 0.41421356237) % 1)).astype(np.float32) | |
| elevation = (5.0 + 1450.0 * ((station_index_all * 0.2718281828) % 1)).astype(np.float32) | |
| station_id = np.asarray([f"DWD-{i:04d}" for i in station_index_all]) | |
| active = np.linspace(0, 536, int(data["active_station_count"]), dtype=np.int64) | |
| dates = np.asarray(data["dates"], dtype="datetime64[D]") | |
| date_id, station_id_index = np.meshgrid(np.arange(len(dates)), active, indexing="ij") | |
| date_id, station_id_index = date_id.ravel(), station_id_index.ravel() | |
| n = len(date_id) | |
| day_of_year = (dates[date_id] - dates[date_id].astype("datetime64[Y]")).astype(int) + 1 | |
| seasonal = np.sin(2 * np.pi * (day_of_year - 172) / 365.25).astype(np.float32) | |
| lat = latitude[station_id_index] | |
| lon = longitude[station_id_index] | |
| elev = elevation[station_id_index] | |
| spatial = (-0.42 * (lat - 51.0) + 0.13 * (lon - 10.0) - 0.0062 * elev).astype(np.float32) | |
| common_weather = rng.normal(0, 2.2, len(dates)).astype(np.float32)[date_id] | |
| target = (11.0 + 9.0 * seasonal + spatial + common_weather + 0.18 * np.sin(station_id_index * 0.17)).astype(np.float32) | |
| variable_offset = np.linspace(-4.0, 5.0, 18, dtype=np.float32) | |
| variable_scale = np.linspace(0.72, 1.25, 18, dtype=np.float32) | |
| centre = target[:, None] * variable_scale[None] + variable_offset[None] + 0.3 * common_weather[:, None] | |
| spread = (0.45 + 0.025 * np.abs(common_weather) + 0.00045 * elev + 0.08 * (1.0 + seasonal)).astype(np.float32) | |
| member_axis = np.linspace(-2.1, 2.1, 50, dtype=np.float32) | |
| wave = np.sin((station_id_index[:, None, None] + 1) * 0.13 + np.arange(18)[None, None] * 0.31 + member_axis[None, :, None] * 1.7) | |
| ensemble = centre[:, None, :] + spread[:, None, None] * variable_scale[None, None, :] * (member_axis[None, :, None] + 0.22 * wave) | |
| ensemble[:, :, 0] += (0.35 + 0.14 * seasonal)[:, None] | |
| auxiliary = np.stack((latitude[station_id_index], longitude[station_id_index], elevation[station_id_index], day_of_year / 366.0), axis=1).astype(np.float32) | |
| output = ROOT / data["file"] | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed( | |
| output, format_version=np.asarray(data["format_version"]), ensemble=ensemble.astype(np.float32), | |
| auxiliary=auxiliary, target=target, station_index=station_id_index.astype(np.int64), date_index=date_id.astype(np.int64), | |
| dates=dates.astype("U10"), lead_hours=np.asarray(data["lead_hours"], np.int64), variables=np.asarray(["t2m"] + [f"var_{i:02d}" for i in range(1, 18)]), | |
| station_id=station_id, station_latitude=latitude, station_longitude=longitude, station_elevation_m=elevation, | |
| ) | |
| if ensemble.shape != (n, 50, 18) or len(station_id) != 537 or not np.isfinite(ensemble).all(): | |
| raise RuntimeError("generated data failed shape or finite-value validation") | |
| print(f"saved={output.relative_to(ROOT)} samples={n} ensemble={ensemble.shape} metadata_stations=537 lead=48h") | |
| if __name__ == "__main__": | |
| main() | |