zhangrenchao commited on
Commit
90cbf31
·
verified ·
1 Parent(s): 36e4eb5

Publish ScaleAdaptiveCM reproduction

Browse files
.gitattributes CHANGED
@@ -1,35 +1,2 @@
1
- *.7z filter=lfs diff=lfs merge=lfs -text
2
- *.arrow filter=lfs diff=lfs merge=lfs -text
3
- *.bin filter=lfs diff=lfs merge=lfs -text
4
- *.bz2 filter=lfs diff=lfs merge=lfs -text
5
- *.ckpt filter=lfs diff=lfs merge=lfs -text
6
- *.ftz filter=lfs diff=lfs merge=lfs -text
7
- *.gz filter=lfs diff=lfs merge=lfs -text
8
- *.h5 filter=lfs diff=lfs merge=lfs -text
9
- *.joblib filter=lfs diff=lfs merge=lfs -text
10
- *.lfs.* filter=lfs diff=lfs merge=lfs -text
11
- *.mlmodel filter=lfs diff=lfs merge=lfs -text
12
- *.model filter=lfs diff=lfs merge=lfs -text
13
- *.msgpack filter=lfs diff=lfs merge=lfs -text
14
- *.npy filter=lfs diff=lfs merge=lfs -text
15
- *.npz filter=lfs diff=lfs merge=lfs -text
16
- *.onnx filter=lfs diff=lfs merge=lfs -text
17
- *.ot filter=lfs diff=lfs merge=lfs -text
18
- *.parquet filter=lfs diff=lfs merge=lfs -text
19
- *.pb filter=lfs diff=lfs merge=lfs -text
20
- *.pickle filter=lfs diff=lfs merge=lfs -text
21
- *.pkl filter=lfs diff=lfs merge=lfs -text
22
  *.pt filter=lfs diff=lfs merge=lfs -text
23
- *.pth filter=lfs diff=lfs merge=lfs -text
24
- *.rar filter=lfs diff=lfs merge=lfs -text
25
- *.safetensors filter=lfs diff=lfs merge=lfs -text
26
- saved_model/**/* filter=lfs diff=lfs merge=lfs -text
27
- *.tar.* filter=lfs diff=lfs merge=lfs -text
28
- *.tar filter=lfs diff=lfs merge=lfs -text
29
- *.tflite filter=lfs diff=lfs merge=lfs -text
30
- *.tgz filter=lfs diff=lfs merge=lfs -text
31
- *.wasm filter=lfs diff=lfs merge=lfs -text
32
- *.xz filter=lfs diff=lfs merge=lfs -text
33
- *.zip filter=lfs diff=lfs merge=lfs -text
34
- *.zst filter=lfs diff=lfs merge=lfs -text
35
- *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  *.pt filter=lfs diff=lfs merge=lfs -text
2
+ *.npz filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
README.md ADDED
@@ -0,0 +1,130 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ language:
4
+ - en
5
+ tags:
6
+ - OneScience
7
+ - Earth Science
8
+ - Probabilistic Downscaling
9
+ - Consistency Model
10
+ frameworks: PyTorch
11
+ ---
12
+
13
+ <p align="center"><strong><span style="font-size: 30px;">ScaleAdaptiveCM</span></strong></p>
14
+
15
+ # Model Introduction
16
+
17
+ ScaleAdaptiveCM uses a single-step consistency model to transform coarse Earth-system-model precipitation into high-resolution probabilistic fields for fast, scale-adaptive climate downscaling and uncertainty analysis.
18
+
19
+ Paper: Fast, scale-adaptive and uncertainty-aware downscaling of Earth system model fields with generative machine learning
20
+ https://doi.org/10.1038/s42256-025-00980-5
21
+
22
+ # Model Description
23
+
24
+ The method was proposed by teams from the Potsdam Institute for Climate Impact Research, the Technical University of Munich, Nanjing University of Information Science and Technology, and collaborating institutions. The paper trains on daily ERA5 precipitation and evaluates POEM, GFDL-ESM4, and SpeedyWeather.jl simulations. The model generates `240×384` precipitation ensembles from `60×96` coarse fields.
25
+
26
+ # Use Cases
27
+
28
+ | Use Case | Description |
29
+ | :---: | :--- |
30
+ | Probabilistic precipitation downscaling | Convert coarse ESM precipitation fields into high-resolution ensembles. |
31
+ | Scale-adaptive generation | Control retained large scales and generated small-scale structure through noise scale. |
32
+ | Local engineering validation | Validate consistency training, single-step sampling, ensemble mean, and spread. |
33
+ | ModelScope/OneCode execution | Validate structured data, training, inference, downscaling metrics, and visualization in ModelScope or OneCode environments. |
34
+ | Multi-GPU training | Validate distributed training and the checkpoint workflow through `torchrun`. |
35
+
36
+ # Usage Instructions
37
+
38
+ ## 1.OneCode
39
+
40
+ [Try intelligent, one-click AI4S programming](https://web-2069360198568017922-iaaj.ksai.scnet.cn:58043/home)
41
+
42
+ ## 2. Download and Installation
43
+
44
+ ```bash
45
+ hf download OneScience-Group/ScaleAdaptiveCM --local-dir ./ScaleAdaptiveCM
46
+ cd ScaleAdaptiveCM
47
+ ```
48
+
49
+ ### Environment Dependencies
50
+
51
+ **Hardware Requirements**
52
+
53
+ - A GPU or DCU is recommended.
54
+ - A CPU can be used for connectivity validation with the default small-sample configuration.
55
+ - DCU users must install DTK in advance. DTK 25.04.2 or later, or the OneScience-recommended version matching the current cluster, is recommended.
56
+
57
+ **DCU Environment**
58
+
59
+ ```bash
60
+ # Activate DTK and Conda first
61
+ conda create -n onescience311 python=3.11 -y
62
+ conda activate onescience311
63
+ pip install onescience[earth-dcu] -i http://mirrors.onescience.ai:3141/pypi/simple/ --trusted-host mirrors.onescience.ai
64
+ ```
65
+
66
+ **GPU Environment**
67
+
68
+ ```bash
69
+ # Activate Conda first
70
+ conda create -n onescience311 python=3.11 -y libstdcxx-ng=12 libgcc-ng=12 gcc_linux-64=12 gxx_linux-64=12
71
+ conda activate onescience311
72
+ pip install onescience[earth-gpu] -i http://mirrors.onescience.ai:3141/pypi/simple/ --trusted-host mirrors.onescience.ai
73
+ ```
74
+
75
+ ### Training Data
76
+
77
+ The paper uses daily ERA5 precipitation from 1940-2018 on a `240×384` target grid and ESM fields on a `60×96` native grid. This repository generates structured precipitation with ITCZ, weather-system, and intermittent small-scale patterns while preserving the fourfold downscaling relation. Synthetic data validate engineering only and do not represent ERA5 or ESM distributions, training scale, or paper performance.
78
+
79
+ ```bash
80
+ python scripts/fake_data.py
81
+ ```
82
+
83
+ ### Training
84
+
85
+ For single-GPU training, use:
86
+
87
+ ```bash
88
+ python scripts/train.py
89
+ ```
90
+
91
+ For multi-GPU training, use:
92
+
93
+ ```bash
94
+ torchrun --nproc_per_node=8 --nnodes=1 --rdzv_id=1000 --rdzv_backend=c10d --max_restarts=0 --master_addr="localhost" --master_port=29500 scripts/train.py
95
+ ```
96
+
97
+ The default reduces samples, network width, and epochs without reducing the `60×96 → 240×384` protocol. Training artifacts are saved to `result/checkpoints/scale_adaptive_cm.pt` and `result/training/metrics.json`.
98
+
99
+ ### Trained Weights
100
+
101
+ The paper does not provide directly loadable official model weights, and this repository bundles no weights under `weight/`.
102
+
103
+ ### Inference
104
+
105
+ ```bash
106
+ python scripts/inference.py
107
+ ```
108
+
109
+ Inference adds scale-guided noise to the coarse field and generates a high-resolution ensemble in one network evaluation per member. Results are saved to `result/output/predictions.npz`.
110
+
111
+ ### Evaluation and Visualization
112
+
113
+ ```bash
114
+ python scripts/result.py
115
+ ```
116
+
117
+ Evaluation computes MAE, RMSE, large-scale correlation, power-spectrum error, and ensemble CRPS and writes `result/evaluation/metrics.json` and `result/evaluation/comparison.png`.
118
+
119
+ # Official OneScience Information
120
+
121
+ | Platform | OneScience Main Repository | Skills Repository |
122
+ | --- | --- | --- |
123
+ | Gitee | https://gitee.com/onescience-ai/onescience | https://gitee.com/onescience-ai/oneskills |
124
+ | GitHub | https://github.com/onescience-ai/OneScience | https://github.com/onescience-ai/oneskills |
125
+
126
+ # Citation and License
127
+
128
+ This repository is an independent engineering reproduction of the public ScaleAdaptiveCM specifications.
129
+
130
+ Use of this repository's code, official model weights, and data remains subject to the licenses and terms of their respective projects.
README_zh.md ADDED
@@ -0,0 +1,136 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: apache-2.0
3
+ language:
4
+ - zh
5
+ - en
6
+ tags:
7
+ - OneScience
8
+ - 地球科学
9
+ - 概率降尺度
10
+ - 一致性模型
11
+ frameworks: PyTorch
12
+ ---
13
+
14
+ <p align="center"><strong><span style="font-size: 30px;">ScaleAdaptiveCM</span></strong></p>
15
+
16
+ # 模型介绍
17
+
18
+ ScaleAdaptiveCM 使用单步一致性生成模型将粗分辨率地球系统模式降水场转换为高分辨率概率降水场,用于快速、尺度自适应的气候降尺度和不确定性分析。
19
+
20
+ 论文:Fast, scale-adaptive and uncertainty-aware downscaling of Earth system model fields with generative machine learning
21
+ https://doi.org/10.1038/s42256-025-00980-5
22
+
23
+ # 模型描述
24
+
25
+ 该方法由波茨坦气候影响研究所、慕尼黑工业大学和南京信息工程大学等机构的研究团队提出。论文使用 ERA5 逐日降水训练一致性模型,并使用 POEM、GFDL-ESM4 和 SpeedyWeather.jl 等模式数据评估。模型适用于从 `60×96` 粗网格生成 `240×384` 高分辨率降水集合。
26
+
27
+ # 适用场景
28
+
29
+ | 场景 | 说明 |
30
+ | :---: | :--- |
31
+ | 概率降水降尺度 | 将粗分辨率 ESM 降水场转换为高分辨率集合。 |
32
+ | 尺度自适应生成 | 通过噪声尺度控制保留的大尺度模式和生成的小尺度结构。 |
33
+ | 本地工程验证 | 验证一致性训练、单步采样、集合均值和离散度。 |
34
+ | ModelScope/OneCode 运行 | 在 ModelScope 或 OneCode 环境中验证结构化数据、训练、推理、降尺度指标和可视化流程。 |
35
+ | 多卡训练 | 通过 `torchrun` 验证分布式训练和 checkpoint 流程。 |
36
+
37
+ # 使用说明
38
+
39
+ ## 1.OneCode
40
+
41
+ [点击体验智能化一键式 AI4S 编程](https://web-2069360198568017922-iaaj.ksai.scnet.cn:58043/home)
42
+
43
+ ## 2.下载安装
44
+
45
+ ```bash
46
+ modelscope download --model OneScience/ScaleAdaptiveCM --local_dir ./ScaleAdaptiveCM
47
+ cd ScaleAdaptiveCM
48
+ ```
49
+
50
+ ### 环境依赖
51
+
52
+ **硬件要求**
53
+
54
+ - 推荐使用 GPU 或 DCU 运行。
55
+ - CPU 可用于默认小样本配置的连通性验证。
56
+ - DCU 用户需预先安装 DTK,建议使用 DTK 25.04.2 以上版本或与当前集群匹配的 OneScience 推荐版本。
57
+
58
+ **DCU环境**
59
+
60
+ ```bash
61
+ # 请首先激活 DTK 及 Conda
62
+ conda create -n onescience311 python=3.11 -y
63
+ conda activate onescience311
64
+ pip install onescience[earth-dcu] -i http://mirrors.onescience.ai:3141/pypi/simple/ --trusted-host mirrors.onescience.ai
65
+ ```
66
+
67
+ **GPU环境**
68
+
69
+ ```bash
70
+ # 请首先激活 Conda
71
+ conda create -n onescience311 python=3.11 -y libstdcxx-ng=12 libgcc-ng=12 gcc_linux-64=12 gxx_linux-64=12
72
+ conda activate onescience311
73
+ pip install onescience[earth-gpu] -i http://mirrors.onescience.ai:3141/pypi/simple/ --trusted-host mirrors.onescience.ai
74
+ ```
75
+
76
+ ### 训练数据介绍
77
+
78
+ 论文使用 1940-2018 年 ERA5 逐日降水,目标网格为 `240×384`,ESM 原生网格为 `60×96`。本仓库生成具有 ITCZ、天气系统和小尺度间歇性的结构化虚拟降水场,并保持真实四倍降尺度关系。虚拟数据仅用于验证工程流程,不代表 ERA5 或 ESM 的真实分布、训练规模或论文性能。
79
+
80
+ ```bash
81
+ python scripts/fake_data.py
82
+ ```
83
+
84
+ ### 训练
85
+
86
+ 单卡训练可使用:
87
+
88
+ ```bash
89
+ python scripts/train.py
90
+ ```
91
+
92
+ 多卡训练可使用:
93
+
94
+ ```bash
95
+ torchrun --nproc_per_node=8 --nnodes=1 --rdzv_id=1000 --rdzv_backend=c10d --max_restarts=0 --master_addr="localhost" --master_port=29500 scripts/train.py
96
+ ```
97
+
98
+ 默认配置缩小样本数、网络宽度和训练轮数,但不缩小 `60×96 → 240×384` 空间协议。正式实验需要真实 ERA5 与 ESM 数据和论文规模模型,训练产物保存到:
99
+
100
+ ```text
101
+ result/checkpoints/scale_adaptive_cm.pt
102
+ result/training/metrics.json
103
+ ```
104
+
105
+ ### 训练权重
106
+
107
+ 论文未提供可直接加载的官方模型权重,本仓库不在 `weight/` 中内置权重。本地 checkpoint 不得描述为官方预训练权重。
108
+
109
+ ### 推理
110
+
111
+ ```bash
112
+ python scripts/inference.py
113
+ ```
114
+
115
+ 推理对粗分辨率降水场施加尺度引导噪声,并在单次网络调用中生成高分辨率集合。结果保存到 `result/output/predictions.npz`。
116
+
117
+ ### 评估和可视化
118
+
119
+ ```bash
120
+ python scripts/result.py
121
+ ```
122
+
123
+ 评估计算 MAE、RMSE、大尺度相关性、功率谱误差和 ensemble CRPS,并生成输入、目标、集合均值与误差图。结果保存到 `result/evaluation/metrics.json` 和 `result/evaluation/comparison.png`。
124
+
125
+ # OneScience官方信息
126
+
127
+ | 平台 | OneScience 主仓库 | Skills 仓库 |
128
+ | --- | --- | --- |
129
+ | Gitee | https://gitee.com/onescience-ai/onescience | https://gitee.com/onescience-ai/oneskills |
130
+ | GitHub | https://github.com/onescience-ai/OneScience | https://github.com/onescience-ai/oneskills |
131
+
132
+ # 引用与许可证
133
+
134
+ 本仓库为 ScaleAdaptiveCM 公开规格的独立工程复现版本。
135
+
136
+ 本仓库代码、官方模型权重和数据的使用仍应以各自项目中的许可证及使用条款为准。
conf/config.yaml ADDED
@@ -0,0 +1,40 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ seed: 42
2
+ data:
3
+ path: data/scale_adaptive_cm_fake.npz
4
+ format_version: scale_adaptive_cm_v1
5
+ samples: 4
6
+ low_grid: [60, 96]
7
+ high_grid: [240, 384]
8
+ scale_factor: 4
9
+ log_epsilon: 0.0001
10
+ model:
11
+ channels: [4, 8, 16]
12
+ time_dim: 16
13
+ sigma_data: 0.5
14
+ train:
15
+ epochs: 1
16
+ batch_size: 1
17
+ learning_rate: 0.0002
18
+ ema_decay: 0.9
19
+ runtime:
20
+ device: cpu
21
+ paths:
22
+ checkpoint: result/checkpoints/scale_adaptive_cm.pt
23
+ training_metrics: result/training/metrics.json
24
+ predictions: result/output/predictions.npz
25
+ evaluation: result/evaluation/metrics.json
26
+ figure: result/evaluation/comparison.png
27
+ evaluation:
28
+ ensemble_members: 3
29
+ guidance_sigma: 0.468
30
+ paper_model:
31
+ high_grid: [240, 384]
32
+ low_grid: [60, 96]
33
+ channels: [128, 128, 256, 256]
34
+ epochs: 150
35
+ batch_size: 1
36
+ learning_rate: 0.0002
37
+ t_min: 0.002
38
+ t_max: 80.0
39
+ rho: 7
40
+ ema_initial_decay: 0.9
config.json ADDED
@@ -0,0 +1,19 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_name": "ScaleAdaptiveCM",
3
+ "model_type": "scale_adaptive_cm",
4
+ "architectures": ["ScaleAdaptiveCM"],
5
+ "framework": "PyTorch",
6
+ "domain": "climate",
7
+ "task": "probabilistic-precipitation-downscaling",
8
+ "implementation": {
9
+ "entry_point": "model/scale_adaptive_cm.py",
10
+ "scope": "core-method, full-spatial-dimension reduced-model engineering reproduction",
11
+ "train_script": "scripts/train.py",
12
+ "inference_script": "scripts/inference.py",
13
+ "evaluation_script": "scripts/result.py",
14
+ "synthetic_data_script": "scripts/fake_data.py"
15
+ },
16
+ "architecture": {"input_channels": 1, "output_channels": 1, "engineering_channels": [4, 8, 16], "paper_channels": [128, 128, 256, 256], "scale_factor": 4},
17
+ "data": {"dataset": "ERA5 and ESM precipitation", "format_version": "scale_adaptive_cm_v1", "low_shape": ["B", 1, 60, 96], "high_shape": ["B", 1, 240, 384], "unit": "mm day-1", "synthetic": true},
18
+ "configuration_sources": ["conf/config.yaml", "model/scale_adaptive_cm.py", "scripts/fake_data.py", "scripts/train.py", "scripts/inference.py", "scripts/result.py"]
19
+ }
configuration.json ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "framework": "PyTorch",
3
+ "task": "probabilistic precipitation downscaling",
4
+ "model": "ScaleAdaptiveCM",
5
+ "input_format": "BCHW: low-resolution [B,1,60,96] and target [B,1,240,384]",
6
+ "protocol": "unconditional consistency training and scale-guided single-step downscaling",
7
+ "default_config": "conf/config.yaml",
8
+ "training": "scripts/train.py",
9
+ "inference": "scripts/inference.py",
10
+ "evaluation": "scripts/result.py",
11
+ "visualization": "scripts/result.py"
12
+ }
model/scale_adaptive_cm.py ADDED
@@ -0,0 +1,85 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Compact consistency model for full-grid precipitation downscaling."""
2
+ import json
3
+ from pathlib import Path
4
+
5
+ import numpy as np
6
+ import torch
7
+ from torch import nn
8
+ import torch.nn.functional as F
9
+ import yaml
10
+
11
+
12
+ def load_config(root):
13
+ return yaml.safe_load((Path(root) / "conf/config.yaml").read_text())
14
+
15
+
16
+ class TimeBlock(nn.Module):
17
+ def __init__(self, cin, cout, time_dim):
18
+ super().__init__()
19
+ groups = min(4, cout)
20
+ self.conv = nn.Conv2d(cin, cout, 3, padding=1)
21
+ self.norm = nn.GroupNorm(groups, cout)
22
+ self.time = nn.Linear(time_dim, cout)
23
+
24
+ def forward(self, x, embedding):
25
+ return F.silu(self.norm(self.conv(x)) + self.time(embedding)[:, :, None, None])
26
+
27
+
28
+ class ScaleAdaptiveCM(nn.Module):
29
+ def __init__(self, channels=(4, 8, 16), time_dim=16, sigma_data=0.5):
30
+ super().__init__()
31
+ self.sigma_data = float(sigma_data)
32
+ self.time_dim = int(time_dim)
33
+ self.time_mlp = nn.Sequential(nn.Linear(time_dim, time_dim), nn.SiLU(), nn.Linear(time_dim, time_dim))
34
+ self.enc1 = TimeBlock(1, channels[0], time_dim)
35
+ self.enc2 = TimeBlock(channels[0], channels[1], time_dim)
36
+ self.mid = TimeBlock(channels[1], channels[2], time_dim)
37
+ self.dec2 = TimeBlock(channels[2] + channels[1], channels[1], time_dim)
38
+ self.dec1 = TimeBlock(channels[1] + channels[0], channels[0], time_dim)
39
+ self.out = nn.Conv2d(channels[0], 1, 1)
40
+ self.model_config = {"channels": list(channels), "time_dim": time_dim, "sigma_data": sigma_data}
41
+
42
+ def embed_time(self, t):
43
+ half = self.time_dim // 2
44
+ freq = torch.exp(torch.linspace(0, -7, half, device=t.device))
45
+ emb = torch.cat((torch.sin(t[:, None] * freq), torch.cos(t[:, None] * freq)), 1)
46
+ return self.time_mlp(emb)
47
+
48
+ def forward(self, noisy, t):
49
+ if noisy.ndim != 4 or noisy.shape[1] != 1:
50
+ raise ValueError("expected [B,1,H,W]")
51
+ emb = self.embed_time(t.float())
52
+ e1 = self.enc1(noisy, emb)
53
+ e2 = self.enc2(F.avg_pool2d(e1, 2), emb)
54
+ mid = self.mid(F.avg_pool2d(e2, 2), emb)
55
+ d2 = self.dec2(torch.cat((F.interpolate(mid, e2.shape[-2:], mode="bilinear", align_corners=False), e2), 1), emb)
56
+ d1 = self.dec1(torch.cat((F.interpolate(d2, e1.shape[-2:], mode="bilinear", align_corners=False), e1), 1), emb)
57
+ raw = self.out(d1)
58
+ sigma2 = self.sigma_data ** 2
59
+ cskip = sigma2 / ((t[:, None, None, None] - 0.002).square() + sigma2)
60
+ cout = self.sigma_data * t[:, None, None, None] / torch.sqrt(t[:, None, None, None].square() + sigma2)
61
+ return cskip * noisy + cout * raw
62
+
63
+
64
+ def structured_fields(samples, high_h, high_w, seed):
65
+ rng = np.random.default_rng(seed)
66
+ yy, xx = np.mgrid[-1:1:complex(high_h), -1:1:complex(high_w)]
67
+ fields = []
68
+ for i in range(samples):
69
+ phase = 2 * np.pi * i / max(samples, 4)
70
+ itcz = 9 * np.exp(-((yy - .12 * np.sin(phase)) / .16) ** 2)
71
+ storms = 18 * np.exp(-((xx - .45 * np.cos(phase)) ** 2 + (yy - .3 * np.sin(phase)) ** 2) / .025)
72
+ texture = 2 * np.maximum(0, np.sin(18 * xx + phase) * np.cos(13 * yy - phase))
73
+ fields.append(np.maximum(0, itcz + storms + texture + rng.normal(0, .15, yy.shape)))
74
+ return np.asarray(fields, np.float32)[:, None]
75
+
76
+
77
+ def radial_spectrum(field):
78
+ power = np.abs(np.fft.fftshift(np.fft.fft2(field))) ** 2
79
+ y, x = np.indices(field.shape); r = np.sqrt((y-field.shape[0]/2)**2 + (x-field.shape[1]/2)**2).astype(int)
80
+ return np.bincount(r.ravel(), power.ravel()) / np.maximum(np.bincount(r.ravel()), 1)
81
+
82
+
83
+ def write_json(path, value):
84
+ path = Path(path); path.parent.mkdir(parents=True, exist_ok=True)
85
+ path.write_text(json.dumps(value, indent=2) + "\n")
scripts/fake_data.py ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys
3
+ import numpy as np
4
+ import torch.nn.functional as F
5
+ import torch
6
+ ROOT = Path(__file__).resolve().parents[1]; sys.path.insert(0, str(ROOT))
7
+ from model.scale_adaptive_cm import load_config, structured_fields
8
+
9
+ c = load_config(ROOT); h,w=c["data"]["high_grid"]; lh,lw=c["data"]["low_grid"]
10
+ high=structured_fields(c["data"]["samples"],h,w,c["seed"])
11
+ low=F.avg_pool2d(torch.from_numpy(high),c["data"]["scale_factor"]).numpy()
12
+ assert low.shape[-2:]==(lh,lw)
13
+ path=ROOT/c["data"]["path"]; path.parent.mkdir(parents=True,exist_ok=True)
14
+ np.savez_compressed(path,format_version=np.array(c["data"]["format_version"]),low=low,high=high,split=np.array(["train"]*(len(high)-1)+["test"]),unit=np.array("mm day-1"))
15
+ print(path,low.shape,high.shape)
scripts/inference.py ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys
3
+ import numpy as np,torch
4
+ import torch.nn.functional as F
5
+ ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
6
+ from model.scale_adaptive_cm import ScaleAdaptiveCM,load_config
7
+ c=load_config(ROOT);d=np.load(ROOT/c["data"]["path"]);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);assert ck["format_version"]==str(d["format_version"])
8
+ m=ScaleAdaptiveCM(**ck["model_config"]);m.load_state_dict(ck["model"]);m.eval(); low=torch.from_numpy(d["low"][d["split"]=="test"]).float(); up=F.interpolate(low,size=c["data"]["high_grid"],mode="bilinear",align_corners=False)
9
+ base=(torch.log(up+c["data"]["log_epsilon"])-np.log(c["data"]["log_epsilon"])-ck["normalization"]["mean"])/ck["normalization"]["std"];members=[];t=torch.full((len(base),),c["evaluation"]["guidance_sigma"])
10
+ with torch.no_grad():
11
+ for i in range(c["evaluation"]["ensemble_members"]): members.append(m(base+t[:,None,None,None]*torch.randn_like(base),t))
12
+ z=torch.stack(members)*ck["normalization"]["std"]+ck["normalization"]["mean"];pred=torch.exp(z+np.log(c["data"]["log_epsilon"]))-c["data"]["log_epsilon"];pred=pred.clamp_min(0).numpy()
13
+ path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,low=low.numpy(),target=d["high"][d["split"]=="test"],members=pred,mean=pred.mean(0),std=pred.std(0),unit=d["unit"]);print(path)
scripts/result.py ADDED
@@ -0,0 +1,13 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys
3
+ import numpy as np
4
+ import matplotlib;matplotlib.use("Agg");import matplotlib.pyplot as plt
5
+ ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
6
+ from model.scale_adaptive_cm import load_config,radial_spectrum,write_json
7
+ c=load_config(ROOT);d=np.load(ROOT/c["paths"]["predictions"]);p,t=d["mean"],d["target"];err=p-t
8
+ mae=float(np.mean(abs(err)));rmse=float(np.sqrt(np.mean(err**2)));corr=float(np.corrcoef(p.ravel(),t.ravel())[0,1]);psp,psT=radial_spectrum(p[0,0]),radial_spectrum(t[0,0]);n=min(len(psp),len(psT));psd=float(np.mean(abs(np.log1p(psp[:n])-np.log1p(psT[:n]))))
9
+ members=d["members"][:,0,0];obs=t[0,0];crps=float(np.mean(abs(members-obs))-0.5*np.mean(abs(members[:,None]-members[None,:])))
10
+ write_json(ROOT/c["paths"]["evaluation"],{"mae_mm_day":mae,"rmse_mm_day":rmse,"large_scale_correlation":corr,"log_psd_mae":psd,"ensemble_crps":crps,"synthetic":True})
11
+ fig,ax=plt.subplots(1,4,figsize=(14,3.4));fields=(d["low"][0,0],t[0,0],p[0,0],abs(err[0,0]));titles=("Low resolution","Target","CM mean","Absolute error")
12
+ for a,f,title in zip(ax,fields,titles):im=a.imshow(f,cmap="viridis");a.set_title(title);a.axis("off");fig.colorbar(im,ax=a,shrink=.7)
13
+ fig.tight_layout();path=ROOT/c["paths"]["figure"];path.parent.mkdir(parents=True,exist_ok=True);fig.savefig(path,dpi=140);plt.close(fig);print(path)
scripts/train.py ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys,copy,os
3
+ import numpy as np, torch
4
+ from torch.nn.parallel import DistributedDataParallel
5
+ ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
6
+ from model.scale_adaptive_cm import ScaleAdaptiveCM,load_config,write_json
7
+ c=load_config(ROOT);rank=int(os.environ.get("RANK",0));world=int(os.environ.get("WORLD_SIZE",1));local=int(os.environ.get("LOCAL_RANK",0));distributed=world>1;use_cuda=torch.cuda.is_available() and torch.cuda.device_count()>=world and c["runtime"]["device"]!="cpu"
8
+ if distributed: torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo")
9
+ device=torch.device(f"cuda:{local}" if use_cuda else "cpu");torch.manual_seed(c["seed"]);torch.set_num_threads(2)
10
+ d=np.load(ROOT/c["data"]["path"]); x=torch.from_numpy(np.log(d["high"][d["split"]=="train"]+c["data"]["log_epsilon"])-np.log(c["data"]["log_epsilon"])).float()
11
+ mean,std=x.mean(),x.std().clamp_min(1e-6);x=(x-mean)/std
12
+ m=ScaleAdaptiveCM(**c["model"]).to(device);target=copy.deepcopy(m).eval();wrapped=DistributedDataParallel(m,device_ids=[local] if use_cuda else None) if distributed else m;opt=torch.optim.RAdam(wrapped.parameters(),lr=c["train"]["learning_rate"]);hist=[];x=x.to(device)
13
+ for e in range(c["train"]["epochs"]):
14
+ noise=torch.randn_like(x);t1=torch.full((len(x),),.35);t2=torch.full((len(x),),.55)
15
+ with torch.no_grad(): y=target(x+t1[:,None,None,None]*noise,t1)
16
+ pred=wrapped(x+t2[:,None,None,None]*noise,t2);loss=(pred-y).abs().mean();opt.zero_grad();loss.backward();opt.step()
17
+ with torch.no_grad():
18
+ for p,q in zip(target.parameters(),m.parameters()):p.mul_(c["train"]["ema_decay"]).add_(q,alpha=1-c["train"]["ema_decay"])
19
+ hist.append({"epoch":e+1,"loss":float(loss)})
20
+ summary=torch.tensor([sum(v["loss"] for v in hist),len(hist)],dtype=torch.float64,device=device)
21
+ if distributed: torch.distributed.all_reduce(summary)
22
+ if rank==0:
23
+ path=ROOT/c["paths"]["checkpoint"];path.parent.mkdir(parents=True,exist_ok=True);torch.save({"model":target.state_dict(),"model_config":c["model"],"format_version":c["data"]["format_version"],"normalization":{"mean":float(mean),"std":float(std)}},path)
24
+ write_json(ROOT/c["paths"]["training_metrics"],{"history":hist,"global_mean_loss":float(summary[0]/summary[1]),"world_size":world});print(path)
25
+ if distributed: torch.distributed.destroy_process_group()
weight/.gitkeep ADDED
File without changes