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