Download model/dincae.py from OneScience-Group/DINCAE: direct link, hf CLI and curl.
- Browser
- Download file 5.59 kB
-
https://huggingface.co/OneScience-Group/DINCAE/resolve/main/model/dincae.py
- Command line
-
hf download hf://OneScience-Group/DINCAE/model/dincae.py
-
curl -L -o dincae.py https://huggingface.co/OneScience-Group/DINCAE/resolve/main/model/dincae.py
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) | |