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
|