Download scripts/inference.py from OneScience-Group/SatMAE: direct link, hf CLI and curl.
- Browser
- Download file 3.05 kB
-
https://huggingface.co/OneScience-Group/SatMAE/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/SatMAE/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/SatMAE/resolve/main/scripts/inference.py
3.05 kB
| """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() | |