File size: 3,179 Bytes
3549cf5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
56
57
58
59
60
61
"""Generate float and paper-style signed-int8 annual embedding fields."""

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))
sys.path.insert(0, str(ROOT / "scripts"))
from model.alphaearthfoundations import AlphaEarthFoundations, dequantize_embeddings, quantize_embeddings
from train import AEFDataset, device_from_config, unpack


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    torch.manual_seed(config["seed"])
    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 do not match")
    model = AlphaEarthFoundations(checkpoint["input_sources"], checkpoint["target_sources"], checkpoint["model_config"]).to(device)
    model.load_state_dict(checkpoint["model"])
    model.eval()
    dataset = AEFDataset(ROOT / config["data"]["root"] / "test.npz", config)
    embeddings, quantized, restored = [], [], []
    reconstruction = {name: [] for name in config["data"]["target_sources"]}
    selected_targets = {name: [] for name in config["data"]["target_sources"]}
    selected_masks = {name: [] for name in config["data"]["target_sources"]}
    with torch.no_grad():
        for index in range(len(dataset)):
            batch = {key: value.unsqueeze(0) for key, value in dataset[index].items()}
            (sources, timestamps, frame_available, targets, masks, target_times,
             target_periods, geometry) = unpack(batch, config, device)
            output = model(sources, timestamps, batch["valid_period"].to(device), frame_available,
                           target_times, geometry, target_periods)
            q = quantize_embeddings(output["embedding"])
            embeddings.append(output["embedding"].cpu().numpy())
            quantized.append(q.cpu().numpy())
            restored.append(dequantize_embeddings(q).cpu().numpy())
            for name, values in output["reconstructions"].items():
                reconstruction[name].append(values.cpu().numpy())
                selected_targets[name].append(targets[name].cpu().numpy())
                selected_masks[name].append(masks[name].cpu().numpy())
    output_dir = ROOT / config["paths"]["inference_dir"]
    output_dir.mkdir(parents=True, exist_ok=True)
    payload = {"embedding": np.concatenate(embeddings), "embedding_s8_power2": np.concatenate(quantized),
               "embedding_dequantized": np.concatenate(restored)}
    payload.update({f"reconstruction_{name}": np.concatenate(values) for name, values in reconstruction.items()})
    payload.update({f"target_{name}": np.concatenate(values) for name, values in selected_targets.items()})
    payload.update({f"mask_{name}": np.concatenate(values) for name, values in selected_masks.items()})
    np.savez_compressed(output_dir / "predictions.npz", **payload)
    print(f"predictions={(output_dir / 'predictions.npz').relative_to(ROOT)}")


if __name__ == "__main__":
    main()