GLONET / scripts /inference.py
yzt15806542928's picture
Upload folder using huggingface_hub
1aeffbb verified
Raw
History Blame Contribute Delete
1.71 kB
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()