File size: 2,249 Bytes
4c4d99c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 | """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()
|