"""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()