File size: 2,614 Bytes
43dce84
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Run probabilistic DINCAE reconstruction on held-out daily SST fields."""

from datetime import date
from pathlib import Path
import sys

import numpy as np
import torch
import yaml

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

from model.dincae import DINCAE, build_input, output_distribution


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    device = torch.device("cuda" if torch.cuda.is_available() and config["runtime"]["device"] != "cpu" else "cpu")
    data_path = ROOT / config["data"]["root"] / config["data"]["file"]
    with np.load(data_path) as loaded:
        data = {key: loaded[key] for key in loaded.files}
    checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location=device, weights_only=False)
    model = DINCAE(**checkpoint["model_config"]).to(device)
    model.load_state_dict(checkpoint["model"])
    model.eval()
    start = min(int(config["data"]["train_samples"]), len(data["timestamps"]) - 1)
    means, variances, targets, missing_masks, inputs_observed = [], [], [], [], []
    with torch.no_grad():
        for index in range(start, len(data["timestamps"])):
            timestamp = date.fromisoformat(str(data["timestamps"][index]))
            model_input = build_input(data["observed_anomaly"], data["precision"], index,
                                      data["longitude"], data["latitude"], timestamp.timetuple().tm_yday)
            output = model(torch.from_numpy(model_input[None]).to(device))
            mean, variance, _ = output_distribution(output, checkpoint["gamma"], checkpoint["delta"])
            means.append(mean[0].cpu().numpy() + data["climatology"])
            variances.append(variance[0].cpu().numpy())
            targets.append(data["sst_anomaly"][index] + data["climatology"])
            missing_masks.append((data["precision"][index] == 0) & data["ocean_mask"])
            inputs_observed.append(data["observed_anomaly"][index] + data["climatology"])
    output_path = ROOT / config["paths"]["inference"]
    output_path.parent.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output_path, prediction=np.asarray(means), variance=np.asarray(variances),
                        target=np.asarray(targets), missing_mask=np.asarray(missing_masks),
                        observed=np.asarray(inputs_observed), ocean_mask=data["ocean_mask"],
                        timestamps=data["timestamps"][start:], units=data["units"])
    print(f"predictions={len(means)} shape={np.asarray(means).shape} output={output_path}")


if __name__ == "__main__":
    main()