File size: 3,549 Bytes
e66b66f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
"""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()