| """Run SatMAE masked reconstruction inference.""" |
|
|
| import argparse |
| import importlib.util |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
|
|
|
|
| def parse_args(): |
| parser = argparse.ArgumentParser(description=__doc__) |
| parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml") |
| parser.add_argument("--data", type=Path, default=None) |
| parser.add_argument("--checkpoint", type=Path, default=None) |
| parser.add_argument("--output-dir", type=Path, default=None) |
| parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto") |
| parser.add_argument("--mask-ratio", type=float, default=None) |
| return parser.parse_args() |
|
|
|
|
| def main(): |
| args = parse_args() |
| config = yaml.safe_load(args.config.read_text()) |
| spec = importlib.util.spec_from_file_location("satmae", ROOT / "model/satmae.py") |
| module = importlib.util.module_from_spec(spec) |
| spec.loader.exec_module(module) |
| model_args = { |
| key: value for key, value in config["model"].items() |
| if key not in {"architecture", "runtime_profile"} |
| } |
| model = module.SatMAE(**model_args) |
| checkpoint_path = args.checkpoint or ROOT / config["paths"]["checkpoint"] |
| if not checkpoint_path.exists(): |
| raise FileNotFoundError(f"checkpoint not found: {checkpoint_path}") |
| checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False) |
| model.load_state_dict(checkpoint["model"]) |
| use_cuda = torch.cuda.is_available() and args.device != "cpu" |
| if args.device == "cuda" and not torch.cuda.is_available(): |
| raise RuntimeError("CUDA was requested but is unavailable") |
| device = torch.device("cuda" if use_cuda else "cpu") |
| model.to(device).eval() |
|
|
| data_path = args.data or ROOT / config["data"]["root"] / "test.npz" |
| archive = np.load(data_path) |
| images = torch.from_numpy(archive["images"]).to(device) |
| timestamps = None |
| if "timestamps" in archive: |
| timestamps = torch.from_numpy(archive["timestamps"]).to(device) |
| with torch.inference_mode(): |
| output = model(images, timestamps=timestamps, mask_ratio=args.mask_ratio) |
| output_dir = args.output_dir or ROOT / config["paths"]["inference_dir"] |
| output_dir.mkdir(parents=True, exist_ok=True) |
| payload = { |
| "target": output["target"].cpu().numpy(), |
| "prediction": output["prediction"].cpu().numpy(), |
| "mask": output["mask"].cpu().numpy(), |
| "labels": archive["labels"], |
| } |
| if timestamps is not None: |
| payload["timestamps"] = timestamps.cpu().numpy() |
| for index, (prediction, target) in enumerate(zip( |
| output["group_predictions"], output["group_targets"] |
| )): |
| payload[f"prediction_group_{index}"] = prediction.cpu().numpy() |
| payload[f"target_group_{index}"] = target.cpu().numpy() |
| np.savez_compressed(output_dir / "reconstruction.npz", **payload) |
| print("inference=", output_dir / "reconstruction.npz") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|