File size: 2,593 Bytes
5c365c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
73
74
75
76
"""ACE's spherical Fourier neural operator implementation.

The paper uses SFNO with Gauss-Legendre transforms. This module builds the
implementation distributed by ``torch_harmonics``. All trainable layers use
random initialization, so a checkpoint can be created from scratch.
"""

from __future__ import annotations

from dataclasses import dataclass

import torch
from torch import nn

try:
    from torch_harmonics.examples.models.sfno import SphericalFourierNeuralOperator
except ImportError as exc:
    raise ImportError(
        "ACE requires torch_harmonics. Activate the configured training "
        "environment or install torch_harmonics before running ACE."
    ) from exc


@dataclass
class SFNOConfig:
    nlat: int = 180
    nlon: int = 360
    in_channels: int = 40
    out_channels: int = 44
    embed_dim: int = 256
    num_layers: int = 8
    filter_type: str = "linear"
    operator_type: str = "dhconv"
    scale_factor: int = 1
    spectral_layers: int = 3
    grid: str = "legendre-gauss"
    grid_internal: str = "legendre-gauss"
    mlp_ratio: float = 2.0
    fallback: bool = False


class SFNOAdapter(nn.Module):
    """Construct the paper-compatible spherical SFNO for one grid size."""

    def __init__(self, config: SFNOConfig) -> None:
        super().__init__()
        if config.fallback:
            raise ValueError("The planar fallback was removed; use the spherical SFNO implementation")
        if config.nlat < 5 or config.nlon < 8:
            raise ValueError("SFNO grid must be at least 5 x 8")
        self.config = config
        self.model = SphericalFourierNeuralOperator(
            img_size=(config.nlat, config.nlon),
            grid=config.grid,
            grid_internal=config.grid_internal,
            scale_factor=config.scale_factor,
            in_chans=config.in_channels,
            out_chans=config.out_channels,
            embed_dim=config.embed_dim,
            num_layers=config.num_layers,
            use_mlp=True,
            mlp_ratio=config.mlp_ratio,
            normalization_layer="none",
            residual_prediction=False,
            pos_embed="none",
            bias=True,
        )

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        if x.ndim != 4:
            raise ValueError(f"SFNO input must be [B,C,H,W], got {tuple(x.shape)}")
        expected = (self.config.in_channels, self.config.nlat, self.config.nlon)
        if tuple(x.shape[1:]) != expected:
            raise ValueError(f"SFNO expects trailing shape {expected}, got {tuple(x.shape[1:])}")
        return self.model(x.float())