DINCAE / model /dincae.py
zhangrenchao's picture
Publish DINCAE reproduction
43dce84 verified
Raw History Blame Contribute Delete
5.59 kB
"""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)