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