Download scripts/fake_data.py from OneScience-Group/IMPROVER-AIFS: direct link, hf CLI and curl.
- Browser
- Download file 3.85 kB
-
https://huggingface.co/OneScience-Group/IMPROVER-AIFS/resolve/main/scripts/fake_data.py
- Command line
-
hf download hf://OneScience-Group/IMPROVER-AIFS/scripts/fake_data.py
-
curl -L -o fake_data.py https://huggingface.co/OneScience-Group/IMPROVER-AIFS/resolve/main/scripts/fake_data.py
3.85 kB
| """Write a lazy 569-station protocol manifest; fields are generated per chunk.""" | |
| from pathlib import Path | |
| import numpy as np | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def station_metadata(count: int): | |
| index = np.arange(count, dtype=np.float32) | |
| latitude = -43.5 + 33.0 * ((index * 0.61803398875) % 1) | |
| longitude = 113.0 + 40.0 * ((index * 0.41421356237) % 1) | |
| elevation = 20 + 1450 * ((index * 0.2718281828) % 1) | |
| return latitude, longitude, elevation.astype(np.float32) | |
| def make_split(path: Path, config: dict): | |
| data = config["data"] | |
| dates = np.asarray(data["valid_dates"], dtype="U10") | |
| lat, lon, elevation = station_metadata(data["station_count"]) | |
| history_dates = np.empty((len(dates), data["history_days"]), dtype="U10") | |
| for i, date in enumerate(dates.astype("datetime64[D]")): | |
| history_dates[i] = (date - np.arange(data["history_days"], 0, -1)).astype("U10") | |
| np.savez(path, format_version=np.asarray(data["format_version"]), valid_dates=dates, | |
| history_dates=history_dates, initialization_utc=np.asarray("1200"), lead_hours=np.arange(241), | |
| station_id=np.asarray([f"JIVE-{i:04d}" for i in range(data["station_count"])]), | |
| station_latitude=lat, station_longitude=lon, station_elevation_m=elevation, | |
| variables=np.asarray(data["variables"]), units=np.asarray(data["units"]), sources=np.asarray(data["sources"]), | |
| neighborhood_shape=np.asarray([3, 3]), target_grid_projection=np.asarray("Albers"), | |
| target_grid_nx=np.asarray(None), target_grid_ny=np.asarray(None), | |
| representation=np.asarray(data["target_representation"])) | |
| def generate_chunk(meta, date_index: int, station_selector, history: bool, seed: int, include_patch: bool = True): | |
| """Emulate lazy reads after authoritative station patches have been extracted.""" | |
| station = np.arange(len(meta["station_id"]))[station_selector] | |
| days = 30 if history else 1 | |
| lead = np.arange(241, dtype=np.float32)[None, :, None, None] | |
| day = np.arange(days, dtype=np.float32)[:, None, None, None] | |
| lat = meta["station_latitude"][station][None, None, None, :] | |
| lon = meta["station_longitude"][station][None, None, None, :] | |
| phase = date_index * 0.7 + day * 0.11 | |
| temp = 20 - 0.35 * (lat + 25) + 4 * np.sin(2 * np.pi * (lead + 12) / 24 + phase) | |
| dew = temp - 5 - 1.5 * np.cos(np.deg2rad(lon) + lead / 48) | |
| wind = 5 + 1.2 * np.abs(np.sin(np.deg2rad(lon) + lead / 18 + phase)) | |
| truth = np.concatenate((temp, dew, wind), axis=2).astype(np.float32) | |
| source = np.arange(3, dtype=np.float32)[None, None, None, :, None] | |
| forecast = truth[:, :, :, None] + (source - 0.6) * np.asarray([1.1, 0.8, 0.5], np.float32)[None, None, :, None, None] | |
| forecast += (lead[..., None] / 240) * np.asarray([0.8, -0.5, 0.6], np.float32)[None, None, :, None, None] | |
| rng = np.random.default_rng(seed + date_index * 1000 + int(station[0])) | |
| forecast += rng.normal(0, 0.08, forecast.shape).astype(np.float32) | |
| offsets = np.asarray([[-0.18, -0.10, -0.04], [-0.08, 0.0, 0.09], [0.03, 0.12, 0.20]], np.float32) | |
| elevation_delta = meta["station_elevation_m"][station] - np.mean(meta["station_elevation_m"]) | |
| if not include_patch: | |
| return forecast.astype(np.float32), truth, elevation_delta.astype(np.float32) | |
| patches = forecast[..., None, None] + offsets | |
| return patches.astype(np.float32), truth, elevation_delta.astype(np.float32) | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| output = ROOT / config["data"]["root"] | |
| output.mkdir(parents=True, exist_ok=True) | |
| make_split(output / "protocol.npz", config) | |
| print("generated=data/protocol.npz dates=2 history_days=30 leads=241 stations=569 patch=3x3 nx=unknown ny=unknown") | |
| if __name__ == "__main__": | |
| main() | |