PrithviEO / scripts /inference.py
zhangrenchao's picture
Add engineering reproduction package
4c4d99c verified
Raw
History Blame Contribute Delete
2.25 kB
"""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()