Download scripts/train.py from OneScience-Group/IMPROVER-AIFS: direct link, hf CLI and curl.
- Browser
- Download file 5.6 kB
-
https://huggingface.co/OneScience-Group/IMPROVER-AIFS/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/IMPROVER-AIFS/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/IMPROVER-AIFS/resolve/main/scripts/train.py
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() | |