Download scripts/inference.py from OneScience-Group/DINCAE: direct link, hf CLI and curl.
- Browser
- Download file 2.61 kB
-
https://huggingface.co/OneScience-Group/DINCAE/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/DINCAE/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/DINCAE/resolve/main/scripts/inference.py
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() | |