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