GraphCast完全新手教程:从零跑通AI天气预报,亲手生成全球多日预测
【免费下载链接】weathernext项目地址: https://gitcode.com/GitHub_Trending/gr/weathernext
GraphCast 是 Google DeepMind 开源的 AI 天气预报工具包:它把全球大气观测场建模成图结构,用神经网络直接预测未来几天的气温、风速、降水等要素,单次推理只需几秒到几分钟。本教程带你走完"选环境 → 跑通模型 → 理解原理 → 上云端"的完整动线,最终让你在自己的机器或免费算力上产出一份可评估的全球天气预报。
GraphCast 能干什么:把大气当成一张图来预报
传统的数值天气预报要解流体方程,动辄需要超级计算机跑数小时。GraphCast 换了个思路:把地球表面切成网格节点、节点之间连成边,让图神经网络(GNN)在图上做消息传递,一次前向推理就得到"未来 6 小时"的大气状态,再把这个输出喂回去作为下一次输入,一步步滚出 10 天甚至更远的预报。仓库里同时提供两代模型:
- GraphCast:确定性预报,0.25° 高分辨率,适合对标数值模式的单点预测;
- GenCast:基于扩散模型的集合预报,一次采样出一批样本,能直接给出"不确定性"而不仅仅是一个答案。
适合谁:想入门 AI 气象的算法工程师、需要低成本全球预报的研究者、以及任何想亲手跑一遍 DeepMind 天气模型的初学者。代码遵循 Apache 2.0,模型权重遵循 CC BY-NC-SA 4.0(非商用),使用前建议确认自己的用途在许可范围内。
先跑起来:30 分钟出第一份预报
最短路径是官方提供的 gencast_mini_demo.ipynb——它使用内存占用最小的GenCast 1p0deg Mini模型,在 Colab 的免费 TPU 上就能跑通完整的"加载数据 → 推理 → 画动画 → 算损失 → 算梯度"流程。
路线一:免费 TPU 笔记本(零成本,推荐)
- 用浏览器打开 Colab,导入仓库中的
gencast_mini_demo.ipynb; - 菜单"运行时"→"更改运行时类型",硬件加速器选 TPU,重启运行时;
- 按顺序执行各单元格。首次运行会自动拉取示例数据和预训练权重。
⚠️ 这一步容易踩坑:Colab 自带的 JAX/libtpu 版本偏旧,notebook 里有专门的单元格负责卸载旧版并重装 TPU 版 JAX,请按顺序执行,不要跳过。
执行到 "Autoregressive rollout" 单元格时,核心逻辑就是这段——用 8 个随机数种子并行采样、按 12 小时一步做自回归展开:
chunks = [ c for c in rollout.chunked_prediction_generator_multiple_runs( predictor_fn=run_forward_pmap, rngs=rngs, inputs=eval_inputs, targets_template=eval_targets * np.nan, forcings=eval_forcings, num_steps_per_chunk=1, num_samples=8, pmap_devices=jax.local_devices()) ] predictions = xarray.combine_by_coords(chunks)跑完后你得到的predictions是一个 xarray 数据集:每个样本一路滚 30 步,即约 15 天的全球预报序列,notebook 会把它和真实值并排画出逐日对比动画。
路线二:本地机器跑 GraphCast 小模型
git clone https://gitcode.com/GitHub_Trending/gr/graphcast cd graphcast python -m venv graphcast-env && source graphcast-env/bin/activate pip install -e .然后打开 graphcast_demo.ipynb,把模型参数来源切到 "Random"(随机权重,用于熟悉管线)或选择GraphCast_small(1° 低分辨率预训练版)加载真实权重。JAX 需要按你的 CPU/GPU 型号单独安装,版本兼容问题建议先查 JAX 官方安装指南。本地跑 0.25° 大模型基本不现实,学习阶段用 1° 版本即可。
它是怎么把"现在"变出"十天后"的
白话版只需三句话:先把网格上的气象场打包成"节点特征 + 边特征"的图,网络在图上跑几轮消息传递后一次性输出未来一个步长的全字段预报,然后把这个预报当成新输入继续滚,直到滚满你要的天数。其中 graphcast/model_utils.py 负责网格与图特征之间的转换,graphcast/deep_typed_graph_net.py 是 GraphCast 的消息传递骨干,graphcast/rollout.py 负责推理时的自回归展开,graphcast/gencast.py 则是 GenCast 的扩散式单步预测器。
各模块的分工一表看懂:
| 模块 | 在管线中的角色 |
|---|---|
gencast.py/graphcast.py | 单步预测网络(一个 12h/6h 增量) |
rollout.py | 自回归滚动,拼出多日预报 |
normalization.py | 用历史统计量归一化输入输出 |
sparse_transformer.py | 在三角网格上稀疏注意力的消息传递 |
losses.py | 纬度加权损失,评估用 |
预报有多可靠?仓库自带 Mini 模型与传统集合预报(ENS)的对比评分卡,Mini 版精度"够用但不代表大模型水平",这也是官方反复提示的一点:
按模型版本挑算力路径
四个预训练模型怎么选
权重和示例数据托管在 Google Cloud 存储桶dm_graphcast,官方文档列出了四个 GenCast 版本加上三个 GraphCast 版本,差异集中在分辨率和显存门槛:
| 模型 | 分辨率 | 系统内存 / 加速卡显存(TPU 推理) | 典型场景 |
|---|---|---|---|
| GenCast 1p0deg Mini | 1° | ~21GB / ~8GB HBM | 免费 Colab 演示 |
| GenCast 1p0deg | 1° | ~21GB / ~8GB HBM | 中端研究 |
| GenCast 0p25deg(含运营版) | 0.25° | ~250GB / ~32GB HBM | 高精度业务对标 |
GraphCast 系列另有 0.25°(37 压层)、1° small、0.25° 运营微调三个版本。具体硬件门槛和区域可用性会随时间变化,以官方最新说明为准。
Google Cloud 上开 TPU VM 跑大模型
完整步骤见 docs/cloud_vm_setup.md,主线如下:
- 在 Cloud Console 的 Compute Engine → TPU 里创建实例,选好区域、芯片数(芯片数 = 能并行的集合预报样本数);建议勾选排队(queuing),否则库存不足会直接报
Stockout:
- SSH 进 VM 后安装 TPU 版 JAX 与 Jupyter:
pip install -U "jax[tpu]" -f https://storage.googleapis.com/jax-releases/libtpu_releases.html pip install jupyter- 把示例数据、权重、统计文件拷到 VM:
gcloud storage cp gs://dm_graphcast/gencast/dataset/source-era5_date-2019-03-29_res-1.0_levels-13_steps-30.nc . gcloud storage cp "gs://dm_graphcast/gencast/params/GenCast 1p0deg <2019.npz" . gcloud storage cp --recursive gs://dm_graphcast/gencast/stats/ .- 启动本地运行时,复制终端输出的
http://localhost:8081/...地址,粘贴到 gencast_demo_cloud_vm.ipynb 的 "Connect to local runtime" 对话框即可接管:
⚠️ 用 Spot(抢占式)TPU 可以省 60% 以上费用,但它随时会被回收且无法重启;实验任务用完记得删掉 VM,否则持续计费。各 TPU 型号的单价和可部署区域变动频繁,以官方计费页为准。
没有 TPU 想用 GPU?需要把注意力实现从splash_mha换成triblockdiag_mha,加载检查点后覆盖配置:
cfg = ckpt.denoiser_architecture_config.sparse_transformer_config cfg.attention_type = "triblockdiag_mha" cfg.mask_type = "full"代价是 0.25° 模型需要约 60GB 显存,且精度相比 TPU 路径有约 0.3%~0.4% 的小幅下降,速度也更慢:
报错别慌,先查这几处
Q:Colab 里报libtpu相关错误或 TPU 识别不到?A:几乎总是 JAX 与 libtpu 版本不匹配。按 notebook 里的单元格卸载libtpu/libtpu-nightly后重装jax[tpu],然后重启运行时。
Q:创建 TPU 实例报Stockout或一直排队?A:TPU 库存紧张是常态。勾选"Enable queuing"让请求进入排队队列;若"Queued Resources"里出现过失败任务,先删掉再重建,失败的占位任务会吃掉配额。
Q:推理直接 OOM(内存不足)?A:编译阶段吃的是系统内存而非显存,0.25° 模型需要约 250GB 主机内存,单机放不下时要请求多芯片主机(如 2x2x1 拓扑)。学习阶段建议直接从 1° 或 Mini 模型起步。
Q:GPU 上跑出来的结果和 TPU 上的不一样?A:这是已知现象。两种注意力实现在代数上等价但数值不完全一致,叠加 GPU 与 TPU 的默认 matmul 精度差异,会带来上述约 0.3% 量级的精度偏移,属于正常偏差而非代码 bug。
Q:预测出 NaN 或海温字段异常?A:海表温度在陆地格点是 NaN,模型内部靠 graphcast/nan_cleaning.py 填充后推理、再还原 NaN。若你自行替换了输入数据,需要保证变量清单、压层、时间分辨率与训练数据一致。
跑通之后去哪继续
按"由浅入深"的顺序,延伸路径大致是:
- 换模型对比:同一份初始场分别跑 Mini 和 0.25° 模型,用 notebook 里的 CRPS 计算对比不确定性表现,直观感受分辨率的价值;
- 读单步网络:从 graphcast/gencast.py 的
GenCast类入手,再看denoiser.py的扩散去噪结构,理解集合预报的采样来源; - 接入自己的数据:参考 graphcast/data_utils.py 里
extract_inputs_targets_forcings的切分方式,把自己的 ERA5/HRES 场切进同样的 inputs/targets/forcings 三元组; - 复现训练:notebook 末尾的 loss 与梯度单元格已演示完整反传,扩展到多步训练需要 ERA5 全量数据(约 TB 级,通过 ECMWF/Weatherbench2 渠道获取,注意各自的使用条款)。
入口汇总:README 是模型清单与许可总览,docs/cloud_vm_setup.md 是云端部署唯一权威文档,学术背景读 GraphCast 的 Science 论文与 GenCast 的 arXiv 论文(引用信息见 README 末尾)。官方反馈邮箱在 README 的 Contact 一节,遇到管线级问题可以先翻 GitHub Issues。
【免费下载链接】weathernext项目地址: https://gitcode.com/GitHub_Trending/gr/weathernext
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考