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