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