File size: 4,668 Bytes
2862bae
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
from torch import nn
from torch.nn import functional as F

from onescience.modules.embedding.fuxiembedding import FuxiEmbedding
from onescience.modules.fc.fuxifc import FuxiFC
from onescience.modules.transformer.fuxitransformer import FuxiTransformer


class Fuxi(nn.Module):
    """
    Fuxi 的主模型实现。

    该模型使用以下组件完成输入编码、二维 trunk 特征提取与 patch 级输出恢复:

    - `OneEmbedding(style="FuxiEmbedding")`
      - 将 `(TimeSteps, Height, Width)` 三维时空块映射为 patch 特征
    - `OneTransformer(style="FuxiTransformer")`
      - 在二维特征图上执行下采样、Swin trunk、上采样
    - `OneFC(style="FuxiFC")`
      - 将每个二维网格位置的 embedding 特征映射为 patch 级输出变量

    在当前实现中:

    - 输入包含多个时间步的二维气象场
    - `patch_size[0]` 默认与 `TimeSteps` 相同,使 embedding 后时间维压缩为 1
    - trunk 只处理二维特征图
    - 最终通过 patch 重排与双线性插值恢复到目标空间分辨率

    Args:
        img_size (tuple[int, int, int]):
            输入空间尺寸 `(TimeSteps, Height, Width)`。
        patch_size (tuple[int, int, int]):
            patch 切分尺寸 `(PatchTimeSteps, PatchHeight, PatchWidth)`。
        in_chans (int):
            输入变量通道数。
        out_chans (int):
            输出变量通道数。
        embed_dim (int):
            embedding 特征维度。
        num_groups (int):
            trunk 中采样模块的 `GroupNorm` 分组数。
        num_heads (int):
            `SwinTransformerV2Stage` 的注意力头数。
        window_size (int | tuple[int, int]):
            trunk 局部窗口大小。
    """

    def __init__(
        self,
        img_size=(2, 721, 1440),
        patch_size=(2, 4, 4),
        in_chans=70,
        out_chans=70,
        embed_dim=1536,
        num_groups=32,
        num_heads=8,
        window_size=7,
    ):
        super().__init__()

        TimeSteps, Height, Width = img_size
        PatchTimeSteps, PatchHeight, PatchWidth = patch_size
        if TimeSteps != PatchTimeSteps:
            raise ValueError(
                "Current Fuxi model expects patch_size[0] to equal img_size[0] "
                "so the embedding output time dimension is 1 before squeeze"
            )

        EmbeddedHeight = Height // PatchHeight
        EmbeddedWidth = Width // PatchWidth
        TransformerInputResolution = (
            EmbeddedHeight // 2,
            EmbeddedWidth // 2,
        )

        self.cube_embedding = FuxiEmbedding(
            img_size=img_size,
            patch_size=patch_size,
            in_chans=in_chans,
            embed_dim=embed_dim,
        )
        self.u_transformer = FuxiTransformer(
            embed_dim=embed_dim,
            num_groups=num_groups,
            input_resolution=TransformerInputResolution,
            num_heads=num_heads,
            window_size=window_size,
        )
        self.fc = FuxiFC(
            in_channels=embed_dim,
            out_channels=out_chans * PatchHeight * PatchWidth,
        )

        self.patch_size = patch_size
        self.transformer_input_resolution = TransformerInputResolution
        self.embedded_resolution = (EmbeddedHeight, EmbeddedWidth)
        self.out_chans = out_chans
        self.img_size = img_size

    def forward(self, x):
        """
        Args:
            x (torch.Tensor):
                输入张量,形状为 `(Batch, in_chans, TimeSteps, Height, Width)`。

        Returns:
            torch.Tensor:
                输出张量,形状为 `(Batch, out_chans, Height, Width)`。
        """
        Batch, _, _, _, _ = x.shape
        _, PatchHeight, PatchWidth = self.patch_size
        EmbeddedHeight, EmbeddedWidth = self.embedded_resolution

        x = self.cube_embedding(x)
        if x.shape[2] != 1:
            raise ValueError(
                f"Expected embedding time dimension 1 before squeeze, but received {x.shape[2]}"
            )

        x = x.squeeze(2)
        x = self.u_transformer(x)
        x = self.fc(x.permute(0, 2, 3, 1))
        x = x.reshape(
            Batch,
            EmbeddedHeight,
            EmbeddedWidth,
            PatchHeight,
            PatchWidth,
            self.out_chans,
        ).permute(0, 1, 3, 2, 4, 5)
        x = x.reshape(
            Batch,
            EmbeddedHeight * PatchHeight,
            EmbeddedWidth * PatchWidth,
            self.out_chans,
        )
        x = x.permute(0, 3, 1, 2)

        x = F.interpolate(x, size=self.img_size[1:], mode="bilinear")

        return x