File size: 2,273 Bytes
efe4fbe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
import argparse
from pathlib import Path

import numpy as np
import torch
from model.metnet_2 import CLASS_RATES, ProceduralField, build_model, load_checkpoint, load_config

parser = argparse.ArgumentParser(description="Run selected-window or streamed full-domain inference")
parser.add_argument("--config", default="conf/config.yaml")
parser.add_argument("--lead", type=int, default=None)
parser.add_argument("--full", action="store_true")
parser.add_argument("--cdf", action="store_true")
args = parser.parse_args()
config = load_config(args.config)
torch.set_num_threads(config["runtime"]["num_threads"])
device = torch.device("cuda" if config["runtime"]["device"] == "auto" and torch.cuda.is_available()
                      else "cpu" if config["runtime"]["device"] == "auto" else config["runtime"]["device"])
model = build_model(config).to(device)
load_checkpoint(config["paths"]["checkpoint"], model)
field, lead = ProceduralField(2001), args.lead or config["inference"]["lead_minutes"]
if args.full:
    output = Path(config["paths"]["predictions"]).with_suffix(".npy")
    print(model.assemble_full(field, lead, output, config["data"]["window"], config["data"]["halo"],
                              config["training"]["class_chunk"], "cdf" if args.cdf else "probability", device))
else:
    window = config["data"]["window"]
    model.eval()
    with torch.no_grad():
        logits = model(field.window(0, 0, window, config["data"]["halo"]).unsqueeze(0).to(device),
                       torch.tensor([lead], device=device), window)[0]
        probabilities = logits.softmax(0).cpu().numpy().astype(np.float32)
    if not np.isfinite(probabilities).all():
        raise FloatingPointError("inference probabilities are not finite")
    output = Path(config["paths"]["predictions"])
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output, probabilities=probabilities, cdf=np.cumsum(probabilities, axis=0),
                        target=field.target_window(0, 0, window, lead).numpy(), rates=CLASS_RATES,
                        lead_minutes=np.int32(lead), coverage=np.array(config["inference"]["coverage"]),
                        is_complete=np.bool_(config["inference"]["is_complete"]))
    print(output)