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