| # coding=utf-8 | |
| # | |
| # SPDX-FileCopyrightText: Copyright (c) 2022 The torch-harmonics Authors. All rights reserved. | |
| # SPDX-License-Identifier: BSD-3-Clause | |
| # | |
| # This module is adapted from the official Spherical Fourier Neural Operator | |
| # reference implementation of Boris Bonev et al. (ICML 2023), published in the | |
| # NVIDIA/torch-harmonics repository (BSD-3-Clause). Only a thin configurable | |
| # wrapper is added so that a single YAML config can drive the model. | |
| # | |
| # Redistribution and use in source and binary forms, with or without | |
| # modification, are permitted provided that the following conditions are met: | |
| # | |
| # 1. Redistributions of source code must retain the above copyright notice, | |
| # this list of conditions and the following disclaimer. | |
| # | |
| # 2. Redistributions in binary form must reproduce the above copyright notice, | |
| # this list of conditions and the following disclaimer in the documentation | |
| # and/or other materials provided with the distribution. | |
| # | |
| # 3. Neither the name of the copyright holder nor the names of its contributors | |
| # may be used to endorse or promote products derived from this software | |
| # without specific prior written permission. | |
| # | |
| # THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" | |
| # AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE | |
| # IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE | |
| # DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE | |
| # FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL | |
| # DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR | |
| # SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER | |
| # CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, | |
| # OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE | |
| # OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. | |
| import torch | |
| import torch.nn as nn | |
| from torch_harmonics.examples.models.sfno import SphericalFourierNeuralOperator | |
| class SFNO(nn.Module): | |
| """ | |
| Configurable wrapper around the official SFNO (Spherical Fourier Neural | |
| Operator, Bonev et al. 2023, ICML). | |
| The underlying network is provided by ``torch-harmonics`` | |
| (``torch_harmonics.examples.models.sfno.SphericalFourierNeuralOperator``), | |
| which replaces the plain FFT of a vanilla FNO by a Spherical Harmonic | |
| Transform (SHT) so that the learned convolution respects the geometry of | |
| the sphere. | |
| Model inputs / outputs are deterministic global fields of shape | |
| ``(Batch, C, H, W)``: a single 6-hour state ``u_t`` is mapped to the next | |
| state ``u_{t+1}`` (trained with weighted L2 losses and 1-2 step rollout). | |
| """ | |
| def __init__( | |
| self, | |
| img_size=(32, 64), | |
| scale_factor=2, | |
| in_chans=4, | |
| out_chans=4, | |
| embed_dim=16, | |
| num_layers=2, | |
| activation_function="gelu", | |
| use_mlp=True, | |
| mlp_ratio=2.0, | |
| drop_rate=0.0, | |
| drop_path_rate=0.0, | |
| normalization_layer="instance_norm", | |
| hard_thresholding_fraction=1.0, | |
| residual_prediction=False, | |
| pos_embed="none", | |
| bias=False, | |
| ): | |
| super().__init__() | |
| self.img_size = tuple(img_size) | |
| self.in_chans = int(in_chans) | |
| self.out_chans = int(out_chans) | |
| self.model = SphericalFourierNeuralOperator( | |
| img_size=self.img_size, | |
| scale_factor=int(scale_factor), | |
| in_chans=self.in_chans, | |
| out_chans=self.out_chans, | |
| embed_dim=int(embed_dim), | |
| num_layers=int(num_layers), | |
| activation_function=activation_function, | |
| use_mlp=use_mlp, | |
| mlp_ratio=mlp_ratio, | |
| drop_rate=drop_rate, | |
| drop_path_rate=drop_path_rate, | |
| normalization_layer=normalization_layer, | |
| hard_thresholding_fraction=hard_thresholding_fraction, | |
| residual_prediction=residual_prediction, | |
| pos_embed=pos_embed, | |
| bias=bias, | |
| ) | |
| def forward(self, x): | |
| return self.model(x) | |