"""Independent PyTorch implementation of DINCAE for daily SST gap filling.""" from __future__ import annotations from typing import Dict, Tuple import numpy as np import torch from torch import Tensor, nn from torch.nn import functional as F class ConvBlock(nn.Sequential): def __init__(self, in_channels: int, out_channels: int, kernel_size: int = 3): padding = kernel_size // 2 super().__init__( nn.Conv2d(in_channels, out_channels, kernel_size, padding=padding), nn.LeakyReLU(0.1, inplace=True), nn.Conv2d(out_channels, out_channels, kernel_size, padding=padding), nn.LeakyReLU(0.1, inplace=True), ) class DINCAE(nn.Module): """Four-level convolutional autoencoder with the paper-sized dense bottleneck.""" def __init__(self, in_channels: int = 10, out_channels: int = 2, filters=(16, 24, 36, 54), bottleneck_features: int = 529, dropout: float = 0.3, kernel_size: int = 3, grid_size: int = 112): super().__init__() if len(filters) != 4 or out_channels != 2: raise ValueError("DINCAE requires four encoder levels and two output channels") if grid_size < 16 or grid_size % 16: raise ValueError("grid_size must be at least 16 and divisible by 16") self.input_shape = (in_channels, grid_size, grid_size) self.encoders = nn.ModuleList() previous = in_channels for width in filters: self.encoders.append(ConvBlock(previous, width, kernel_size)) previous = width self.pool = nn.AvgPool2d(2) latent_size = grid_size // 16 flattened = filters[-1] * latent_size * latent_size self.latent_size = latent_size self.dense_encode = nn.Linear(flattened, bottleneck_features) self.dropout = nn.Dropout(dropout) self.dense_decode = nn.Linear(bottleneck_features, flattened) decoder_inputs = (filters[3] + filters[3], filters[3] + filters[2], filters[2] + filters[1], filters[1] + filters[0]) decoder_outputs = (filters[3], filters[2], filters[1], filters[0]) self.decoders = nn.ModuleList([ ConvBlock(cin, cout, kernel_size) for cin, cout in zip(decoder_inputs, decoder_outputs) ]) self.output = nn.Conv2d(filters[0], out_channels, 1) def forward(self, x: Tensor) -> Tensor: if tuple(x.shape[-2:]) != self.input_shape[1:] or x.shape[1] != self.input_shape[0]: raise ValueError(f"expected [B,{self.input_shape[0]},{self.input_shape[1]},{self.input_shape[2]}], got {tuple(x.shape)}") skips = [] for encoder in self.encoders: x = encoder(x) skips.append(x) x = self.pool(x) x = x.flatten(1) x = self.dense_decode(self.dropout(F.leaky_relu(self.dense_encode(x), 0.1))) x = x.view(x.shape[0], -1, self.latent_size, self.latent_size) for decoder, skip in zip(self.decoders, reversed(skips)): x = F.interpolate(x, size=skip.shape[-2:], mode="nearest") x = decoder(torch.cat((x, skip), dim=1)) return self.output(x) def output_distribution(output: Tensor, gamma: float = 10.0, delta: float = 1e-3) -> Tuple[Tensor, Tensor, Tensor]: """Recover precision, variance, and mean from [log precision, temperature*precision].""" log_precision = output[:, 0] precision = torch.maximum(torch.exp(torch.minimum(log_precision, log_precision.new_tensor(gamma))), log_precision.new_tensor(delta)) variance = precision.reciprocal() mean = output[:, 1] / precision return mean, variance, precision def masked_gaussian_nll(output: Tensor, target: Tensor, mask: Tensor, gamma: float = 10.0, delta: float = 1e-3) -> Tensor: mean, _, precision = output_distribution(output, gamma, delta) target = target[:, 0] valid = mask[:, 0].bool() & torch.isfinite(target) if not torch.any(valid): return output.sum() * 0.0 error = target[valid] - mean[valid] return (0.5 * (precision[valid] * error.square() - torch.log(precision[valid]))).mean() def reconstruction_metrics(target: np.ndarray, prediction: np.ndarray, mask: np.ndarray) -> Dict[str, float]: valid = mask.astype(bool) & np.isfinite(target) & np.isfinite(prediction) if not np.any(valid): return {"rmse": float("nan"), "crmse": float("nan"), "bias": float("nan")} error = prediction[valid].astype(np.float64) - target[valid].astype(np.float64) bias = float(error.mean()) return { "rmse": float(np.sqrt(np.mean(error ** 2))), "crmse": float(np.sqrt(np.mean((error - bias) ** 2))), "bias": bias, } def build_input(observed: np.ndarray, precision: np.ndarray, index: int, longitude: np.ndarray, latitude: np.ndarray, day_of_year: int) -> np.ndarray: """Build 10 channels: three anomaly*precision, three precision, and four auxiliaries.""" frames = [] for offset in (0, -1, 1): j = min(max(index + offset, 0), observed.shape[0] - 1) anomaly = np.nan_to_num(observed[j], nan=0.0) frames.extend((anomaly * precision[j], precision[j])) angle = 2.0 * np.pi * day_of_year / 365.25 frames.extend((longitude, latitude, np.full_like(longitude, np.cos(angle)), np.full_like(longitude, np.sin(angle)))) return np.stack(frames).astype(np.float32)