Download scripts/inference.py from OneScience-Group/Pangu-ICON-DKE: direct link, hf CLI and curl.
- Browser
- Download file 3.55 kB
-
https://huggingface.co/OneScience-Group/Pangu-ICON-DKE/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/Pangu-ICON-DKE/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/Pangu-ICON-DKE/resolve/main/scripts/inference.py
3.55 kB
| """Compute complete-grid DKE products in latitude chunks and compact T719 spectra.""" | |
| import os | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| sys.path.insert(0, str(ROOT)) | |
| from model.pangu_icon_dke import PanguIconDKEDiagnostics | |
| def main(): | |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) | |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location="cpu", weights_only=False) | |
| if checkpoint["format_version"] != config["data"]["format_version"]: | |
| raise ValueError("checkpoint protocol mismatch") | |
| model = PanguIconDKEDiagnostics(checkpoint["protocol"], checkpoint["coefficients"]) | |
| world = int(os.environ.get("WORLD_SIZE", "1")) | |
| distributed = world > 1 | |
| if distributed: | |
| torch.distributed.init_process_group("gloo") | |
| rank = torch.distributed.get_rank() if distributed else 0 | |
| experiments, times, _, nlat, nlon = model.fields.field_shape | |
| chunk = int(config["runtime"]["latitude_chunk"]) | |
| local = {} | |
| for experiment in range(rank, experiments, world): | |
| global_dke = np.zeros(times, dtype=np.float64) | |
| dke_maps = np.empty((times, nlat, nlon), dtype=np.float32) | |
| spectrum = np.empty((times, int(checkpoint["protocol"]["triangular_truncation"]) + 1), dtype=np.float64) | |
| weights_total = 0.0 | |
| weighted_sum = np.zeros(times, dtype=np.float64) | |
| for start in range(0, nlat, chunk): | |
| stop = min(start + chunk, nlat) | |
| dke = model.fields.dke_chunk_all_times(experiment, slice(start, stop)) | |
| weights = np.cos(np.deg2rad(model.fields.latitudes[start:stop]))[None, :, None] | |
| weighted_sum += np.sum(dke * weights, axis=(1, 2)) | |
| dke_maps[:, start:stop] = dke | |
| weights_total += float(nlon * weights.sum()) | |
| global_dke[:] = weighted_sum / weights_total | |
| for time in range(times): | |
| spectrum[time] = model.fields.spectrum(experiment, time) | |
| local[experiment] = (global_dke, dke_maps, spectrum) | |
| if distributed: | |
| gathered = [None] * world if rank == 0 else None | |
| torch.distributed.gather_object(local, gathered, dst=0) | |
| if rank == 0: | |
| local = {key: value for shard in gathered for key, value in shard.items()} | |
| if rank == 0: | |
| ordered = [local[index] for index in range(experiments)] | |
| global_dke = np.stack([item[0] for item in ordered]) | |
| maps = np.stack([item[1] for item in ordered]) | |
| spectra = np.stack([item[2] for item in ordered]) | |
| output = ROOT / config["paths"]["inference"] | |
| output.parent.mkdir(parents=True, exist_ok=True) | |
| np.savez_compressed(output, global_dke=global_dke, dke_maps=maps, spectra=spectra, experiments=np.asarray(checkpoint["protocol"]["experiments"]), times_hours=np.asarray(checkpoint["protocol"]["times_hours"]), latitudes_degrees=np.asarray(checkpoint["protocol"]["latitudes_degrees"]), longitudes_degrees=np.asarray(checkpoint["protocol"]["longitudes_degrees"]), total_wavenumber=np.arange(spectra.shape[-1]), format_version=np.asarray(checkpoint["format_version"]), logical_field_shape=np.asarray(checkpoint["protocol"]["field_shape"]), logical_spectral_shape=np.asarray(checkpoint["protocol"]["spectral_shape"])) | |
| print(f"output={output.relative_to(ROOT)} global={global_dke.shape} maps={maps.shape} spectra={spectra.shape}") | |
| if distributed: | |
| torch.distributed.destroy_process_group() | |
| if __name__ == "__main__": | |
| main() | |