File size: 2,443 Bytes
53becf5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Generate Clay embeddings and MAE reconstructions for every configured sensor."""

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.clayfoundation import ClayFoundation
from train import ClayDataset, 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 = ClayFoundation(checkpoint["model_config"]).to(device)
    model.load_state_dict(checkpoint["model"])
    model.eval()
    dataset = ClayDataset(ROOT / config["data"]["root"] / "test.npz", config)
    raw = dataset.data
    payload = {
        "format_version": np.asarray(config["data"]["format_version"]),
        "class_target": raw["class_target"],
        "regression_target": raw["regression_target"],
    }
    time = torch.from_numpy(raw["time"]).to(device)
    latlon = torch.from_numpy(raw["latlon"]).to(device)
    teacher = torch.from_numpy(raw["teacher_target"]).to(device)
    with torch.no_grad():
        for name, spec in config["data"]["sensors"].items():
            pixels = torch.from_numpy(raw[f"pixels_{name}"]).to(device)
            waves = torch.from_numpy(raw[f"wavelengths_{name}"]).to(device)
            outputs = model(pixels, time, latlon, float(spec["gsd"]), waves, teacher, mask_ratio=0.0)
            payload[f"embedding_{name}"] = outputs["embedding"].cpu().numpy()
            payload[f"projected_embedding_{name}"] = outputs["projected_embedding"].cpu().numpy()
            payload[f"reconstruction_{name}"] = outputs["reconstruction"].cpu().numpy()
            payload[f"pixels_{name}"] = raw[f"pixels_{name}"]
            payload[f"reconstruction_loss_{name}"] = np.asarray(float(outputs["reconstruction_loss"]))
            payload[f"representation_loss_{name}"] = np.asarray(float(outputs["representation_loss"]))
    output = ROOT / config["paths"]["inference_dir"] / "predictions.npz"
    output.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output, **payload)
    print(f"predictions={output.relative_to(ROOT)}")


if __name__ == "__main__":
    main()