File size: 7,017 Bytes
a13f4b9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
"""Self-contained, paper-aligned Surya forecasting model."""

import torch
from torch import nn
from torch.nn import functional as F


CHANNEL_NAMES = ("AIA_94", "AIA_131", "AIA_171", "AIA_193", "AIA_211", "AIA_304",
                 "AIA_335", "AIA_1600", "HMI_magnetogram", "HMI_continuum",
                 "HMI_doppler", "HMI_vector_x", "HMI_vector_y")


def signum_log(x):
    """Compress signed solar products while retaining their sign."""
    return torch.sign(x) * torch.log1p(torch.abs(x))


class SpectralGating(nn.Module):
    def __init__(self, grid, dim):
        super().__init__()
        self.grid = grid
        self.weight = nn.Parameter(torch.randn(grid, grid // 2 + 1, dim, 2) * 0.02)
        self.norm = nn.LayerNorm(dim)
        self.mlp = nn.Sequential(nn.Linear(dim, dim * 2), nn.GELU(), nn.Linear(dim * 2, dim))

    def forward(self, tokens):
        residual = tokens
        values = self.norm(tokens).reshape(tokens.shape[0], self.grid, self.grid, -1).float()
        spectrum = torch.fft.rfft2(values, dim=(1, 2), norm="ortho")
        spectrum = spectrum * torch.view_as_complex(self.weight.float())
        values = torch.fft.irfft2(spectrum, s=(self.grid, self.grid), dim=(1, 2), norm="ortho")
        return residual + self.mlp(values.reshape_as(tokens).to(tokens.dtype))


class LongShortAttention(nn.Module):
    def __init__(self, grid, dim, heads, window, global_tokens):
        super().__init__()
        self.grid, self.window = grid, window
        self.norm = nn.LayerNorm(dim)
        self.local = nn.MultiheadAttention(dim, heads, batch_first=True)
        self.global_attn = nn.MultiheadAttention(dim, heads, batch_first=True)
        self.projection = nn.Linear(dim * 2, dim)
        self.score = nn.Linear(dim, global_tokens)
        self.global_norm = nn.LayerNorm(dim)
        self.output_norm = nn.LayerNorm(dim)
        self.mlp_norm = nn.LayerNorm(dim)
        self.mlp = nn.Sequential(nn.Linear(dim, dim * 4), nn.GELU(), nn.Linear(dim * 4, dim))

    def forward(self, tokens):
        residual = tokens
        values = self.norm(tokens)
        b, _, d = values.shape
        grid = values.reshape(b, self.grid, self.grid, d)
        w = self.window
        # Unfold with padding creates overlapping neighborhoods at every token.
        padded = F.pad(grid.permute(0, 3, 1, 2), (w // 2, w // 2, w // 2, w // 2), mode="replicate")
        neighborhoods = F.unfold(padded, kernel_size=w, padding=0).transpose(1, 2)
        windows = neighborhoods.reshape(b, self.grid * self.grid, d, w * w)
        windows = windows.permute(0, 1, 3, 2).reshape(-1, w * w, d)
        local = self.local(windows, windows, windows, need_weights=False)[0]
        local = local[:, (w * w) // 2].reshape(b, self.grid * self.grid, d)
        mixing = self.score(values).transpose(1, 2).softmax(-1)
        compressed = self.global_norm(mixing @ values)
        global_context = self.global_attn(values, compressed, compressed, need_weights=False)[0]
        tokens = residual + self.projection(torch.cat((local, global_context), dim=-1))
        tokens = self.output_norm(tokens)
        return tokens + self.mlp(self.mlp_norm(tokens))


class Surya(nn.Module):
    def __init__(self, image_size=32, patch_size=4, channels=13, input_steps=2,
                 embed_dim=64, depth=4, spectral_blocks=1, num_heads=4,
                 window_size=2, global_tokens=4):
        super().__init__()
        if image_size % patch_size:
            raise ValueError("image_size must be divisible by patch_size")
        if spectral_blocks >= depth:
            raise ValueError("spectral_blocks must be smaller than depth")
        if embed_dim % num_heads:
            raise ValueError("embed_dim must be divisible by num_heads")
        self.image_size, self.patch_size = image_size, patch_size
        self.channels, self.input_steps = channels, input_steps
        self.grid = image_size // patch_size
        if channels != 13:
            raise ValueError("Surya requires the 13 SDO AIA/HMI channels")
        patch_dim = channels * patch_size ** 2
        self.patch_embed = nn.Linear(patch_dim * 2, embed_dim)
        self.temporal_embed = nn.Parameter(torch.zeros(1, input_steps, 1, embed_dim))
        self.position = nn.Parameter(torch.zeros(1, self.grid ** 2, embed_dim))
        blocks = [SpectralGating(self.grid, embed_dim) for _ in range(spectral_blocks)]
        blocks += [LongShortAttention(self.grid, embed_dim, num_heads, window_size, global_tokens)
                   for _ in range(depth - spectral_blocks)]
        self.blocks = nn.ModuleList(blocks)
        self.norm = nn.LayerNorm(embed_dim)
        self.decoder = nn.Sequential(
            nn.ConvTranspose2d(embed_dim, embed_dim // 2, patch_size, stride=patch_size),
            nn.GELU(), nn.Conv2d(embed_dim // 2, channels, 3, padding=1))
        nn.init.normal_(self.position, std=0.02)
        nn.init.zeros_(self.decoder[-1].weight)
        nn.init.zeros_(self.decoder[-1].bias)

    def tokenize(self, frames):
        b, t, c, h, w = frames.shape
        if (t, c, h, w) != (self.input_steps, self.channels, self.image_size, self.image_size):
            raise ValueError("Expected BTCHW input matching configured dimensions")
        p = self.patch_size
        current = frames[:, -1]
        delta = current - frames[:, -2]
        values = torch.stack((current, delta), dim=1).reshape(b, 2 * c, h // p, p, w // p, p)
        values = values.permute(0, 2, 4, 1, 3, 5).reshape(b, self.grid ** 2, -1)
        return self.patch_embed(values) + self.position + self.temporal_embed[:, -1]

    def _predict(self, frames):
        tokens = self.tokenize(frames)
        for block in self.blocks:
            tokens = block(tokens)
        tokens = self.norm(tokens)
        features = tokens.mean(1)
        grid = tokens.transpose(1, 2).reshape(tokens.shape[0], -1, self.grid, self.grid)
        prediction = frames[:, -1] + self.decoder(grid)
        if not torch.isfinite(prediction).all():
            raise FloatingPointError("Surya produced a non-finite prediction")
        return {"prediction": prediction, "features": features}

    def forward(self, frames, steps=None):
        if steps is None:
            return self._predict(frames)
        if steps < 1:
            raise ValueError("steps must be positive")
        history, predictions = frames, []
        for _ in range(steps):
            prediction = self._predict(history)["prediction"]
            predictions.append(prediction)
            history = torch.cat((history[:, 1:], prediction[:, None]), dim=1)
        return torch.stack(predictions, dim=1)

    def rollout(self, frames, steps):
        return self(frames, steps=steps)

    def rollout_loss(self, frames, targets):
        predictions = self.rollout(frames, targets.shape[1])
        step_losses = (predictions - targets).square().mean(dim=(0, 2, 3, 4))
        return {"loss": step_losses.mean(), "step_losses": step_losses,
                "predictions": predictions}