MetNet-2 / scripts /inference.py
zhangrenchao's picture
Publish MetNet-2 reproduction
efe4fbe verified
Raw
History Blame Contribute Delete
2.27 kB
#!/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)