| """Run multi-temporal reconstruction and embedding inference.""" |
|
|
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| import yaml |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.prithvi_eo import PrithviEO2 |
| from train import PrithviDataset, device_from_config |
|
|
|
|
| def main(): |
| config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) |
| device = device_from_config(config) |
| checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=True) |
| if checkpoint["format_version"] != config["data"]["format_version"]: |
| raise ValueError("checkpoint and data formats are incompatible") |
| model = PrithviEO2(checkpoint["model_config"]).to(device) |
| model.load_state_dict(checkpoint["model"]) |
| model.eval() |
| dataset = PrithviDataset(ROOT / config["data"]["root"] / "test.npz", config) |
| pixels = torch.stack([dataset[index]["pixels"] for index in range(len(dataset))]).to(device) |
| temporal = torch.stack([dataset[index]["temporal"] for index in range(len(dataset))]).to(device) |
| location = torch.stack([dataset[index]["location"] for index in range(len(dataset))]).to(device) |
| torch.manual_seed(int(config["seed"])) |
| with torch.no_grad(): |
| output = model(pixels, temporal, location) |
| cls_embedding, patch_embeddings = model.encode(pixels, temporal, location) |
| target = ROOT / config["paths"]["inference_dir"] / "predictions.npz" |
| target.parent.mkdir(parents=True, exist_ok=True) |
| np.savez_compressed( |
| target, |
| format_version=np.asarray(config["data"]["format_version"]), |
| pixels=pixels.cpu().numpy(), |
| reconstruction=output["reconstruction"].cpu().numpy(), |
| mask=output["mask"].cpu().numpy(), |
| embedding=cls_embedding.cpu().numpy(), |
| patch_embeddings=patch_embeddings.cpu().numpy(), |
| temporal_coords=temporal.cpu().numpy(), |
| location_coords=location.cpu().numpy(), |
| class_target=dataset.data["class_target"], |
| regression_target=dataset.data["regression_target"], |
| masked_patch_mse=np.asarray(float(output["loss"])), |
| ) |
| print(f"predictions={target.relative_to(ROOT)}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|