Team Ai
Modelpublic

OneScience-Group/MetNet-2

sourceHugging Faceapache-2.0updated 27d agoView on Hugging Face
0likes26downloads
inference.py42 linesDownload Raw Back to scripts
1#!/usr/bin/env python32import argparse3from pathlib import Path4 5import numpy as np6import torch7from model.metnet_2 import CLASS_RATES, ProceduralField, build_model, load_checkpoint, load_config8 9parser = argparse.ArgumentParser(description="Run selected-window or streamed full-domain inference")10parser.add_argument("--config", default="conf/config.yaml")11parser.add_argument("--lead", type=int, default=None)12parser.add_argument("--full", action="store_true")13parser.add_argument("--cdf", action="store_true")14args = parser.parse_args()15config = load_config(args.config)16torch.set_num_threads(config["runtime"]["num_threads"])17device = torch.device("cuda" if config["runtime"]["device"] == "auto" and torch.cuda.is_available()18                      else "cpu" if config["runtime"]["device"] == "auto" else config["runtime"]["device"])19model = build_model(config).to(device)20load_checkpoint(config["paths"]["checkpoint"], model)21field, lead = ProceduralField(2001), args.lead or config["inference"]["lead_minutes"]22if args.full:23    output = Path(config["paths"]["predictions"]).with_suffix(".npy")24    print(model.assemble_full(field, lead, output, config["data"]["window"], config["data"]["halo"],25                              config["training"]["class_chunk"], "cdf" if args.cdf else "probability", device))26else:27    window = config["data"]["window"]28    model.eval()29    with torch.no_grad():30        logits = model(field.window(0, 0, window, config["data"]["halo"]).unsqueeze(0).to(device),31                       torch.tensor([lead], device=device), window)[0]32        probabilities = logits.softmax(0).cpu().numpy().astype(np.float32)33    if not np.isfinite(probabilities).all():34        raise FloatingPointError("inference probabilities are not finite")35    output = Path(config["paths"]["predictions"])36    output.parent.mkdir(parents=True, exist_ok=True)37    np.savez_compressed(output, probabilities=probabilities, cdf=np.cumsum(probabilities, axis=0),38                        target=field.target_window(0, 0, window, lead).numpy(), rates=CLASS_RATES,39                        lead_minutes=np.int32(lead), coverage=np.array(config["inference"]["coverage"]),40                        is_complete=np.bool_(config["inference"]["is_complete"]))41    print(output)42