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