ClimODE / scripts /inference.py
yzt15806542928's picture
Upload folder using huggingface_hub
807a08b verified
Raw
History Blame Contribute Delete
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()