Download scripts/velocity.py from OneScience-Group/ClimODE: direct link, hf CLI and curl.
- Browser
- Download file 5.3 kB
-
https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/velocity.py
- Command line
-
hf download hf://OneScience-Group/ClimODE/scripts/velocity.py
-
curl -L -o velocity.py https://huggingface.co/OneScience-Group/ClimODE/resolve/main/scripts/velocity.py
5.3 kB
| """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() | |