| |
| 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) |
|
|