IMPROVER-AIFS / scripts /train.py
zhangrenchao's picture
Upload folder using huggingface_hub
7a7efc9 verified
Raw
History Blame Contribute Delete
5.6 kB
"""Fit two valid-time folds without caching full-station threshold probabilities."""
import json
import os
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.improver_aifs import ImproverAIFS
from fake_data import generate_chunk
def fit_fold(config, meta, fold, rank):
model = ImproverAIFS(config["model"])
stations = int(config["data"]["station_count"])
chunk_size = int(config["train"]["station_chunk_size"])
# Bias uses every station but only station-centre values, so no full-grid cache exists.
for start in range(0, stations, chunk_size):
section = slice(start, min(start + chunk_size, stations))
centres, analyses, elevation = generate_chunk(
meta, fold, section, True, int(config["seed"]), include_patch=False
)
model.fit_bias_chunk(
torch.from_numpy(centres), torch.from_numpy(analyses),
torch.from_numpy(elevation), section,
)
sample_count = max(int(config["train"]["calibration_stations"]), int(config["train"]["blend_stations"]))
sample_stations = np.linspace(0, stations - 1, sample_count, dtype=np.int64)
patches, analyses, elevation = generate_chunk(
meta, fold, sample_stations, True, int(config["seed"]), include_patch=True
)
truth = torch.from_numpy(analyses)
thresholds = [torch.tensor(values, dtype=torch.float32) for values in config["model"]["thresholds"]]
source_expected = []
for source in range(3):
patch = torch.from_numpy(patches[:, :, :, source])
patch[:, :, 0] += (-0.0098 * torch.from_numpy(elevation)).view(1, 1, -1, 1, 1)
patch -= model.bias[source, :, :, sample_stations].unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
source_expected.append(patch[..., 1, 1])
probabilities = []
for variable, values in enumerate(thresholds):
width = model.fuzzy_widths[variable]
probability_patch = ((patch[:, :, variable].unsqueeze(2) - values.view(1, 1, -1, 1, 1, 1) + width) / (2 * width)).clamp(0, 1)
probabilities.append(model.neighborhood(model.recursive_filter(probability_patch)))
model.fit_reliability(probabilities, truth, thresholds, source)
del probabilities
blend_count = int(config["train"]["blend_stations"])
blend_input = torch.stack(source_expected, dim=3)[..., :blend_count]
blend_truth = truth[..., :blend_count]
optimizer = torch.optim.Adam([model.blend_logits], lr=float(config["train"]["learning_rate"]))
history = []
for epoch in range(int(config["train"]["epochs"])):
loss = torch.mean((model.blend_expected(blend_input) - blend_truth) ** 2)
if not torch.isfinite(loss):
raise ValueError("non-finite blend loss")
optimizer.zero_grad(set_to_none=True)
loss.backward()
optimizer.step()
history.append({"epoch": epoch + 1, "blend_mse_loss": float(loss.detach())})
record = {
"fold": fold, "held_out_valid_date": str(meta["valid_dates"][fold]),
"history_days": 30, "bias_stations": stations,
"calibration_stations": int(config["train"]["calibration_stations"]),
"blend_stations": blend_count, "rank": rank, "history": history,
}
print(f"rank={rank} fold={fold} history_days=30 bias_stations=569 calibration_stations={sample_count} loss={history[-1]['blend_mse_loss']:.6f}")
return model.state_dict(), record
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
meta = np.load(ROOT / config["data"]["root"] / "protocol.npz", allow_pickle=True)
if str(meta["format_version"]) != config["data"]["format_version"] or meta["history_dates"].shape != (2, 30):
raise ValueError("protocol requires two valid dates and complete 30-day histories")
distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
if distributed:
torch.distributed.init_process_group("gloo")
rank = torch.distributed.get_rank() if distributed else 0
world_size = torch.distributed.get_world_size() if distributed else 1
local = []
for fold in range(2):
if fold % world_size == rank:
state, record = fit_fold(config, meta, fold, rank)
local.append((fold, state, record))
if distributed:
gathered = [None] * world_size
torch.distributed.all_gather_object(gathered, local)
combined = [item for rank_items in gathered for item in rank_items]
else:
combined = local
if rank == 0:
combined.sort(key=lambda item: item[0])
if [item[0] for item in combined] != [0, 1]:
raise RuntimeError("DDP ranks did not produce both valid-time folds")
fold_states = [item[1] for item in combined]
records = [item[2] for item in combined]
checkpoint = ROOT / config["paths"]["checkpoint"]
metrics = ROOT / config["paths"]["training_metrics"]
checkpoint.parent.mkdir(parents=True, exist_ok=True)
metrics.parent.mkdir(parents=True, exist_ok=True)
torch.save({
"model": {"fold_states": fold_states, "fold_semantics": "valid-time 2-fold"},
"model_config": config["model"],
"format_version": config["data"]["format_version"],
}, checkpoint)
metrics.write_text(json.dumps({"folds": records, "distributed_world_size": world_size}, indent=2) + "\n")
if distributed:
torch.distributed.destroy_process_group()
if __name__ == "__main__":
main()