"""Initial velocity fitting and cache management for ClimODE training.""" from __future__ import annotations import sys from pathlib import Path import numpy as np import torch import torch.nn as nn import torch.optim as optim 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 OptimVelocity try: from torchcubicspline import NaturalCubicSpline, natural_cubic_spline_coeffs except ImportError: # pragma: no cover - dependency is optional for tiny smoke tests NaturalCubicSpline = None natural_cubic_spline_coeffs = None def _time_derivative(history: torch.Tensor, interval_hours: float = 6.0) -> torch.Tensor: """Estimate the derivative at the final history point. The cubic-spline path is identical to the official implementation. The finite-difference fallback is only for environments without the optional package and is explicitly reported to the caller. """ if history.ndim != 5: raise ValueError(f"Expected history [N,3,K,H,W], got {tuple(history.shape)}") if natural_cubic_spline_coeffs is not None: times = torch.arange(3, device=history.device, dtype=history.dtype) * interval_hours values = history.permute(1, 0, 2, 3, 4) coeffs = natural_cubic_spline_coeffs(times, values) spline = NaturalCubicSpline(coeffs) return spline.derivative(times[-1]) return (3.0 * history[:, 2] - 4.0 * history[:, 1] + history[:, 0]) / (2.0 * interval_hours) def build_rbf_kernel( lat2d: torch.Tensor, lon2d: torch.Tensor, sigma: float = 1.0, ) -> torch.Tensor: coords = torch.stack([lat2d.reshape(-1), lon2d.reshape(-1)], dim=1).float() distances = torch.cdist(coords, coords).square() kernel = torch.exp(-distances / (2.0 * sigma * sigma)) return torch.linalg.inv(kernel) def optimize_velocity( history: torch.Tensor, current: torch.Tensor, kernel_inv: torch.Tensor, epochs: int = 200, learning_rate: float = 2.0, smoothing_alpha: float = 1.0e-7, ) -> torch.Tensor: """Fit [N,2K,H,W] velocities using the official penalized objective.""" if current.ndim != 4: raise ValueError(f"Expected current [N,K,H,W], got {tuple(current.shape)}") num_years, channels, height, width = current.shape model = OptimVelocity(num_years, height, width, channels).to(current.device) optimizer = optim.Adam(model.parameters(), lr=learning_rate) delta_u = _time_derivative(history) best_loss = float("inf") best_velocity = None for _ in range(max(int(epochs), 1)): optimizer.zero_grad(set_to_none=True) out, vx, vy = model(current.unsqueeze(1)) vx_flat = vx.view(num_years, channels, -1, 1) vy_flat = vy.view(num_years, channels, -1, 1) kernel = kernel_inv.to(current.device).expand(num_years, channels, -1, -1) smooth_x = torch.matmul(torch.matmul(vx_flat.transpose(2, 3), kernel), vx_flat).mean() smooth_y = torch.matmul(torch.matmul(vy_flat.transpose(2, 3), kernel), vy_flat).mean() loss = nn.functional.mse_loss(out.squeeze(1), delta_u) + smoothing_alpha * (smooth_x + smooth_y) loss.backward() optimizer.step() if float(loss.detach()) < best_loss: best_loss = float(loss.detach()) best_velocity = torch.cat([vx.detach(), vy.detach()], dim=2).squeeze(1).clone() if best_velocity is None: raise RuntimeError("Velocity optimization produced no result") return best_velocity def fit_velocity_cache( dataset, constants: torch.Tensor, lat2d: torch.Tensor, lon2d: torch.Tensor, output_path: str | Path, epochs: int = 200, learning_rate: float = 2.0, smoothing_alpha: float = 1.0e-7, kernel_sigma: float = 1.0, ) -> torch.Tensor: del constants # Kept in the signature to make the training handoff explicit. kernel_inv = build_rbf_kernel(lat2d, lon2d, kernel_sigma) velocities = [] for index in range(len(dataset)): item = dataset[index] history = item["history"].float() current = item["observations"][0].float() velocities.append( optimize_velocity( history, current, kernel_inv, epochs=epochs, learning_rate=learning_rate, smoothing_alpha=smoothing_alpha, ) ) result = torch.stack(velocities) path = Path(output_path) path.parent.mkdir(parents=True, exist_ok=True) torch.save({"velocity": result, "starts": dataset.starts, "years": dataset.years}, path) return result def load_velocity_cache(path: str | Path, expected_length: int | None = None) -> torch.Tensor: try: checkpoint = torch.load(path, map_location="cpu", weights_only=True) except TypeError: # PyTorch before weights_only support. checkpoint = torch.load(path, map_location="cpu") velocity = checkpoint["velocity"] if isinstance(checkpoint, dict) else checkpoint if expected_length is not None and len(velocity) != expected_length: raise ValueError(f"Velocity cache length {len(velocity)} != dataset length {expected_length}") return velocity.float()