File size: 2,518 Bytes
f4a39ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""Visualize official NeuralGCM pressure-level predictions."""
from __future__ import annotations
import argparse
import matplotlib.pyplot as plt
import numpy as np
import xarray as xr
try:
    from common import load_config, resolve_path
except ModuleNotFoundError:  # supports ``python -m scripts.result``
    from scripts.common import load_config, resolve_path


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default="conf/config.yaml")
    parser.add_argument("--input")
    parser.add_argument("--variable", default="temperature_500")
    parser.add_argument("--level", type=int, default=500)
    parser.add_argument("--lead", type=int, default=0)
    parser.add_argument("--output")
    args = parser.parse_args()
    config = load_config(args.config)
    ds = xr.open_dataset(resolve_path(args.input or config["inference"].get("output", "results/predictions.nc"), args.config))
    # Official PressureLevelModel output stores named variables with an
    # explicit pressure ``level`` dimension; retain compatibility with the
    # earlier channel-packed smoke output when present.
    variable = args.variable
    if variable not in ds.data_vars and variable.rsplit("_", 1)[-1].isdigit():
        base, suffix = variable.rsplit("_", 1)
        if base in ds.data_vars:
            variable = base
            args.level = int(suffix)
    if variable in ds.data_vars:
        field = ds[variable]
        if "level" in field.dims:
            level = int(args.level)
            if level not in ds.level.values:
                raise ValueError(f"Unknown pressure level {level}; available levels={ds.level.values.tolist()}")
            field = field.sel(level=level)
    elif "channel" in ds.coords and args.variable in [str(x) for x in ds.channel.values]:
        field = ds.prediction.sel(channel=args.variable)
    else:
        raise ValueError(f"Unknown channel/variable {args.variable!r}")
    if not bool(np.isfinite(field.values).all()):
        raise ValueError(
            f"Variable {variable!r} contains NaN/Inf; refusing to create "
            "a misleading visualization"
        )
    field.isel(time=args.lead).plot(figsize=(12, 4), cmap="coolwarm")
    output = resolve_path(args.output or "results/forecast.png", args.config)
    output.parent.mkdir(parents=True, exist_ok=True)
    plt.tight_layout(); plt.savefig(output, dpi=150)
    print(f"Saved visualization to {output}")


if __name__ == "__main__":
    main()