用随机数据跑通 AuroraSmallPretrained,先 load_checkpoint() 再构造 Batch 并 forward。首次运行会下载约 500 MB 权重。
from datetime import datetime
import torch
from aurora import AuroraSmallPretrained, Batch, Metadata
model = AuroraSmallPretrained()
model.load_checkpoint()
batch = Batch(
surf_vars={k: torch.randn(1, 2, 17, 32) for k in ("2t", "10u", "10v", "msl")},
static_vars={k: torch.randn(17, 32) for k in ("lsm", "z", "slt")},
atmos_vars={k: torch.randn(1, 2, 4, 17, 32) for k in ("z", "u", "v", "t", "q")},
metadata=Metadata(
lat=torch.linspace(90, -90, 17),
lon=torch.linspace(0, 360, 32 + 1)[:-1],
time=(datetime(2020, 6, 1, 12, 0),),
atmos_levels=(100, 250, 500, 850),
),
)
prediction = model.forward(batch)
print(prediction.surf_vars["2t"])