| import argparse |
| from pathlib import Path |
| import numpy as np |
| import torch |
|
|
| import sys |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) |
| from model.corrdiff import CorrDiff |
|
|
|
|
| def main(): |
| parser = argparse.ArgumentParser(description="Run CorrDiff inference") |
| parser.add_argument("--data", default="data/era5_corrdiff.npz") |
| parser.add_argument("--checkpoint", default="data/checkpoints/model_bak.pth") |
| parser.add_argument("--ensemble-size", type=int, default=1) |
| parser.add_argument("--output", default="result/output/predictions.npz") |
| args = parser.parse_args() |
| data = np.load(args.data) |
| coarse = torch.from_numpy(data["input"]) |
| model = CorrDiff() |
| checkpoint = Path(args.checkpoint) |
| if checkpoint.exists(): |
| model.load_state_dict(torch.load(checkpoint, map_location="cpu")["model"]) |
| model.eval() |
| samples = [] |
| with torch.no_grad(): |
| for _ in range(args.ensemble_size): |
| samples.append(model(coarse).numpy().astype("float32")) |
| output = Path(args.output) |
| output.parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed(output, prediction=np.stack(samples), input=data["input"]) |
| print(f"prediction: {np.stack(samples).shape}") |
| print(f"saved: {output}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|