DINCAE / scripts /inference.py
zhangrenchao's picture
Publish DINCAE reproduction
43dce84 verified
Raw
History Blame Contribute Delete
2.61 kB
"""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()