| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
|
|
| if __package__ in (None, ""): |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1])) |
|
|
| from model.earthformer import Earthformer |
| from script.data_loader import make_loader |
| from script.utils import clean_state_dict, load_checkpoint_payload, load_config, resolve_cli_path, resolve_device |
|
|
|
|
| def parse_args() -> argparse.Namespace: |
| parser = argparse.ArgumentParser(description="Run Earthformer inference on the first batch") |
| parser.add_argument("--config", help="Optional data/config override; checkpoint config is used by default") |
| parser.add_argument("--checkpoint", default="data/checkpoint/earthformer.pt") |
| parser.add_argument("--split", choices=("train", "val", "test"), default="test") |
| parser.add_argument("--output", default="output/predictions.npz") |
| parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") |
| return parser.parse_args() |
|
|
|
|
| def main() -> None: |
| args = parse_args() |
| device = resolve_device(args.device) |
| payload = load_checkpoint_payload(resolve_cli_path(args.checkpoint), device) |
| config = load_config(args.config) if args.config else payload["config"] |
| model = Earthformer(config).to(device) |
| model.load_state_dict(clean_state_dict(payload["model"])) |
| model.eval() |
| loader, _ = make_loader(config, args.split, shuffle=False) |
| inputs, targets = next(iter(loader)) |
| with torch.no_grad(): |
| predictions = model(inputs.to(device)).clamp(0.0, 1.0).cpu().numpy() |
| output = Path(resolve_cli_path(args.output)) |
| output.parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed(output, inputs=inputs.numpy(), targets=targets.numpy(), predictions=predictions) |
| print(json.dumps({"output": str(output), "shape": list(predictions.shape)}, indent=2)) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|