Download scripts/inference.py from OneScience-Group/ClimODE: direct link, hf CLI and curl.
- Browser
- Download file 8.77 kB
-
https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/ClimODE/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/inference.py
8.77 kB
| """Run ClimODE global forecasts and save machine-readable outputs.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| from torch.utils.data import DataLoader | |
| PROJECT_ROOT = Path(__file__).resolve().parents[1] | |
| if str(PROJECT_ROOT) not in sys.path: | |
| sys.path.insert(0, str(PROJECT_ROOT)) | |
| from model.climode import load_checkpoint | |
| from scripts.data_loader import ClimODESeriesDataset, load_constants | |
| from scripts.metrics import evaluate, save_metrics | |
| from scripts.velocity import fit_velocity_cache, load_velocity_cache | |
| def _load_config(path: Path) -> dict: | |
| with path.open("r", encoding="utf-8") as handle: | |
| return yaml.safe_load(handle) | |
| def _parse_years(value: str | None, fallback: list[int]) -> list[int]: | |
| if value is None: | |
| return list(fallback) | |
| years = [int(item.strip()) for item in value.split(",") if item.strip()] | |
| if not years: | |
| raise ValueError("year override must contain at least one integer") | |
| return years | |
| def _device(value: str | None) -> torch.device: | |
| if value: | |
| return torch.device(value) | |
| return torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| def _resolve(path: str | Path) -> Path: | |
| value = Path(path) | |
| return value if value.is_absolute() else PROJECT_ROOT / value | |
| def main() -> None: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", type=Path, default=PROJECT_ROOT / "conf/config.yaml") | |
| parser.add_argument("--checkpoint", type=Path, default=None) | |
| parser.add_argument("--device", type=str, default=None) | |
| parser.add_argument("--test-years", type=str, default=None) | |
| parser.add_argument("--sequence-length", type=int, default=None) | |
| parser.add_argument("--max-samples", type=int, default=None) | |
| parser.add_argument("--velocity-epochs", type=int, default=None) | |
| parser.add_argument("--velocity-cache", type=Path, default=None) | |
| parser.add_argument("--data-dir", type=Path, default=None, help="Override data.data_dir") | |
| parser.add_argument("--stats-dir", type=Path, default=None, help="Override data.stats_dir") | |
| parser.add_argument("--static-file", type=Path, default=None, help="Override data.static_file") | |
| parser.add_argument("--output-dir", type=Path, default=None) | |
| args = parser.parse_args() | |
| args.config = _resolve(args.config) | |
| args.checkpoint = _resolve(args.checkpoint) if args.checkpoint is not None else None | |
| args.velocity_cache = ( | |
| _resolve(args.velocity_cache) if args.velocity_cache is not None else None | |
| ) | |
| args.data_dir = _resolve(args.data_dir) if args.data_dir is not None else None | |
| args.stats_dir = _resolve(args.stats_dir) if args.stats_dir is not None else None | |
| args.static_file = _resolve(args.static_file) if args.static_file is not None else None | |
| args.output_dir = _resolve(args.output_dir) if args.output_dir is not None else None | |
| config = _load_config(args.config) | |
| data_cfg, model_cfg, vel_cfg = config["data"], config["model"], config["velocity"] | |
| root = _resolve(args.data_dir or data_cfg["data_dir"]) | |
| stats_dir = _resolve(args.stats_dir or data_cfg.get("stats_dir", root / "static")) | |
| test_years = _parse_years(args.test_years, data_cfg["test_years"]) | |
| sequence_length = args.sequence_length or data_cfg.get("sequence_length", 8) | |
| dataset = ClimODESeriesDataset( | |
| root, | |
| test_years, | |
| stats_dir=stats_dir, | |
| model_size=(data_cfg["model_height"], data_cfg["model_width"]), | |
| sequence_length=sequence_length, | |
| normalize=data_cfg.get("normalize", True), | |
| ) | |
| loader = DataLoader(dataset, batch_size=1, shuffle=False, num_workers=0) | |
| static_file = _resolve(args.static_file or data_cfg["static_file"]) | |
| constants, lat, lon = load_constants( | |
| static_file, (data_cfg["model_height"], data_cfg["model_width"]) | |
| ) | |
| device = _device(args.device) | |
| constants = constants.to(device) | |
| lat_device, lon_device = lat.unsqueeze(0).to(device), lon.unsqueeze(0).to(device) | |
| velocity_root = _resolve(vel_cfg["cache_dir"]) | |
| velocity_path = args.velocity_cache or (velocity_root / "test.pt") | |
| if velocity_path.is_file(): | |
| velocity = load_velocity_cache(velocity_path, len(dataset)) | |
| else: | |
| velocity = fit_velocity_cache( | |
| dataset, | |
| constants, | |
| lat, | |
| lon, | |
| velocity_path, | |
| epochs=args.velocity_epochs if args.velocity_epochs is not None else vel_cfg["epochs"], | |
| learning_rate=vel_cfg["learning_rate"], | |
| smoothing_alpha=vel_cfg["smoothing_alpha"], | |
| kernel_sigma=vel_cfg["kernel_sigma"], | |
| ) | |
| checkpoint_path = args.checkpoint | |
| if checkpoint_path is None: | |
| checkpoint_path = _resolve(model_cfg["default_checkpoint"]) | |
| if not checkpoint_path.is_file(): | |
| pretrained = _resolve(model_cfg["pretrained_checkpoint"]) | |
| if pretrained.is_file(): | |
| checkpoint_path = pretrained | |
| if checkpoint_path is None or not checkpoint_path.is_file(): | |
| raise FileNotFoundError( | |
| "No checkpoint found; pass --checkpoint or provide model.default_checkpoint" | |
| ) | |
| model = load_checkpoint(checkpoint_path, map_location="cpu").to(device).eval() | |
| predictions, uncertainties, targets = [], [], [] | |
| with torch.no_grad(): | |
| for sample_index, batch in enumerate(loader): | |
| if args.max_samples is not None and sample_index >= args.max_samples: | |
| break | |
| observations = batch["observations"].squeeze(0).to(device) | |
| time_steps = batch["time_steps"].squeeze(0).to(device) | |
| initial = observations[0].unsqueeze(1) | |
| model.update_param([velocity[sample_index].to(device), constants, lat_device, lon_device]) | |
| mean, std, _ = model( | |
| time_steps, | |
| initial, | |
| atol=model_cfg["atol"], | |
| rtol=model_cfg["rtol"], | |
| ) | |
| # Index 0 is the analysis state used to initialize the ODE. Official | |
| # evaluation starts at index 1, corresponding to a six-hour lead. | |
| if mean.shape[0] > 1: | |
| predictions.append(mean[1:].detach().cpu().numpy()) | |
| uncertainties.append(std[1:].detach().cpu().numpy()) | |
| targets.append(observations[1:].detach().cpu().numpy()) | |
| if not predictions: | |
| raise RuntimeError("No test samples were processed") | |
| valid_lengths = np.asarray([item.shape[0] for item in predictions], dtype=np.int64) | |
| max_lead = int(valid_lengths.max()) | |
| def _pad(items: list[np.ndarray]) -> np.ndarray: | |
| shape = (len(items), max_lead) + tuple(items[0].shape[1:]) | |
| padded = np.full(shape, np.nan, dtype=np.float32) | |
| for index, item in enumerate(items): | |
| padded[index, : item.shape[0]] = item | |
| return padded | |
| pred_array = _pad(predictions) | |
| std_array = _pad(uncertainties) | |
| target_array = _pad(targets) | |
| scale = (dataset.maximum - dataset.minimum).numpy().reshape(1, 1, 1, 5, 1, 1) | |
| offset = dataset.minimum.numpy().reshape(1, 1, 1, 5, 1, 1) | |
| pred_physical = pred_array * scale + offset | |
| target_physical = target_array * scale + offset | |
| std_physical = std_array * scale | |
| output_dir = args.output_dir or _resolve(data_cfg["output_dir"]) | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| np.save(output_dir / "predictions.npy", pred_array) | |
| np.save(output_dir / "std.npy", std_array) | |
| np.save(output_dir / "targets.npy", target_array) | |
| np.save(output_dir / "valid_lengths.npy", valid_lengths) | |
| metrics = evaluate( | |
| pred_physical, | |
| target_physical, | |
| lat.numpy(), | |
| std_physical, | |
| crps_predictions=pred_array, | |
| crps_targets=target_array, | |
| crps_std=std_array, | |
| valid_lengths=valid_lengths, | |
| ) | |
| metrics["checkpoint"] = str(checkpoint_path) | |
| metrics["outputs_normalized"] = True | |
| metrics_path = _resolve(config["output"]["metrics_file"]) | |
| if args.output_dir is not None: | |
| metrics_path = output_dir.parent / "metrics.json" | |
| save_metrics(metrics, metrics_path) | |
| manifest = { | |
| "checkpoint": str(checkpoint_path), | |
| "samples": int(pred_array.shape[0]), | |
| "shape": list(pred_array.shape), | |
| "valid_lengths": valid_lengths.tolist(), | |
| "variables": ["z", "t", "t2m", "u10", "v10"], | |
| "output_dir": str(output_dir), | |
| "metrics": str(metrics_path), | |
| } | |
| (output_dir / "inference_manifest.json").write_text( | |
| json.dumps(manifest, indent=2), encoding="utf-8" | |
| ) | |
| print(json.dumps(manifest)) | |
| if __name__ == "__main__": | |
| main() | |