File size: 1,711 Bytes
1aeffbb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
import argparse
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

import torch
import yaml

from data_loader import SyntheticOceanDataset
from model.glonet import GLONET


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default=str(ROOT / "conf/config.yaml"))
    parser.add_argument("--checkpoint", default=None)
    args = parser.parse_args()
    with open(args.config, encoding="utf-8") as handle:
        config = yaml.safe_load(handle)
    channels = len(config["data"]["channels"])
    model = GLONET(channels * config["data"]["input_steps"], out_channels=channels,
                   hidden_channels=config["model"]["hidden_channels"], modes=config["model"]["modes"],
                   layers=config["model"]["layers"])
    checkpoint = Path(args.checkpoint or ROOT / config["training"]["checkpoint"])
    state = torch.load(checkpoint, map_location="cpu", weights_only=False)
    model.load_state_dict(state["model"])
    model.eval()
    sample, _ = SyntheticOceanDataset(1, channels, config["data"]["grid"],
                                      input_steps=config["data"]["input_steps"],
                                      output_steps=config["data"]["output_steps"],
                                      data_dir=str(ROOT / config["data"]["data_dir"]))[0]
    with torch.no_grad():
        prediction = model(sample.unsqueeze(0))
    output = ROOT / config["project"]["result_dir"] / "data" / "prediction.pt"
    output.parent.mkdir(parents=True, exist_ok=True)
    torch.save(prediction, output)
    print(f"prediction_shape={tuple(prediction.shape)}")


if __name__ == "__main__":
    main()