CorrDiff / scripts /inference.py
zhangrenchao's picture
Upload folder using huggingface_hub
986404c verified
Raw
History Blame
1.3 kB
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()