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