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