FuXi / model /fuxi.py
OneScience's picture
Upload folder using huggingface_hub
2862bae verified
Raw
History Blame Contribute Delete
4.67 kB
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