File size: 2,660 Bytes
1aeffbb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 | """GLONET reference architecture based on the public paper description."""
import torch
from torch import nn
class SpectralConv2d(nn.Module):
def __init__(self, channels, modes):
super().__init__()
self.modes_y, self.modes_x = modes
self.weight = nn.Parameter(torch.randn(channels, channels, self.modes_y, self.modes_x, 2) * 0.02)
def forward(self, x):
height, width = x.shape[-2:]
spectrum = torch.fft.rfft2(x, norm="ortho")
out = torch.zeros_like(spectrum)
modes_y = min(self.modes_y, height)
modes_x = min(self.modes_x, spectrum.shape[-1])
weight = torch.view_as_complex(self.weight[:, :, :modes_y, :modes_x].contiguous())
out[:, :, :modes_y, :modes_x] = torch.einsum(
"bixy,ioxy->boxy", spectrum[:, :, :modes_y, :modes_x], weight
)
return torch.fft.irfft2(out, s=(height, width), norm="ortho")
class SpectralBlock(nn.Module):
def __init__(self, channels, modes):
super().__init__()
self.spectral = SpectralConv2d(channels, modes)
self.pointwise = nn.Conv2d(channels, channels, 1)
self.activation = nn.GELU()
def forward(self, x):
return self.activation(self.spectral(x) + self.pointwise(x))
class CNNBranch(nn.Module):
def __init__(self, channels):
super().__init__()
self.net = nn.Sequential(
nn.Conv2d(channels, channels, 3, padding=1), nn.GELU(),
nn.Conv2d(channels, channels, 3, padding=1), nn.GELU(),
)
def forward(self, x):
return self.net(x)
class GLONET(nn.Module):
"""Two-day to one-day global ocean forecast reference model.
The paper does not publish a complete layer configuration, so all sizing
choices remain explicit constructor parameters rather than hidden claims.
"""
def __init__(self, in_channels, out_channels=None, hidden_channels=32, modes=(6, 8), layers=4):
super().__init__()
out_channels = out_channels or in_channels
self.input_projection = nn.Conv2d(in_channels, hidden_channels, 1)
self.fno = nn.Sequential(*[SpectralBlock(hidden_channels, modes) for _ in range(layers)])
self.cnn = CNNBranch(hidden_channels)
self.output_projection = nn.Sequential(
nn.Conv2d(hidden_channels * 2, hidden_channels, 1), nn.GELU(),
nn.Conv2d(hidden_channels, out_channels, 1),
)
def forward(self, x):
if x.ndim == 5:
x = x.flatten(1, 2)
features = self.input_projection(x)
return self.output_projection(torch.cat((self.fno(features), self.cnn(features)), dim=1))
|