File size: 7,221 Bytes
87f2bd3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 | """Train all 21 LOSO folds; DDP ranks independently shard outer folds."""
import json
import argparse
import os
import random
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.climemu_s2l import (FEATURE_COUNT, FORMAT_VERSION, GRID_SHAPE,
MODEL_NAME, DualRidge, SharedKernelGPR,
area_weights, select_ridge_alpha, weighted_rmse)
def load_data(path):
data = np.load(path)
short, long = data["short_response"], data["long_response"]
if str(data["format_version"]) != FORMAT_VERSION:
raise ValueError("data format version mismatch")
if short.shape != (21, *GRID_SHAPE) or long.shape != short.shape:
raise ValueError("responses must both have shape [21,145,192]")
ids = [str(value) for value in data["scenario_ids"]]
if len(set(ids)) != 21 or not np.isfinite(short).all() or not np.isfinite(long).all():
raise ValueError("scenario IDs must be unique and fields finite")
return data, short.reshape(21, FEATURE_COUNT), long.reshape(21, FEATURE_COUNT), ids
def train_fold(fold, x, y, scenario_ids, config, weights):
training = np.asarray([index for index in range(21) if index != fold])
alpha, cv_scores = select_ridge_alpha(
x[training], y[training], config["model"]["ridge"]["alphas"],
int(config["model"]["ridge"]["inner_folds"]), weights)
ridge = DualRidge(alpha).fit(x[training], y[training])
ridge_prediction = ridge.predict(x[fold:fold + 1]).detach().cpu().numpy()[0]
gpr_config = config["model"]["gpr"]
gpr = SharedKernelGPR(gpr_config["kernel_mode"], gpr_config["jitter"])
optimizer_trace = gpr.fit(x[training], y[training], gpr_config["optimizer_steps"], gpr_config["learning_rate"])
gpr_prediction = gpr.predict(x[fold:fold + 1]).detach().cpu().numpy()[0]
state = {"fold": fold, "held_out_scenario_id": scenario_ids[fold], "train_indices": training.tolist(),
"ridge_alpha": alpha, "ridge_cv_scores": cv_scores,
"gpr_kernel_mode": gpr.kernel_mode, "gpr_hyperparameters": gpr.hyperparameters(),
"gpr_optimizer": optimizer_trace}
metrics = {"scenario_id": scenario_ids[fold],
"ridge_weighted_rmse": weighted_rmse(y[fold], ridge_prediction, weights),
"gpr_weighted_rmse": weighted_rmse(y[fold], gpr_prediction, weights),
"selected_alpha": alpha, "gpr_final_loss": optimizer_trace["loss"][-1],
"gpr_backward": optimizer_trace["gradient_seen"],
"gpr_parameter_update": optimizer_trace["parameter_updated"]}
return state, metrics
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--resume", action="store_true", help="resume completed LOSO folds from the checkpoint")
args = parser.parse_args()
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
seed = int(config["seed"])
random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
torch.set_default_dtype(torch.float64)
distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
if distributed:
torch.distributed.init_process_group(config["runtime"]["ddp_backend"])
rank = torch.distributed.get_rank() if distributed else 0
world = torch.distributed.get_world_size() if distributed else 1
data, x_numpy, y_numpy, scenario_ids = load_data(ROOT / config["data"]["path"])
x, y = torch.from_numpy(x_numpy).double(), torch.from_numpy(y_numpy).double()
weights = area_weights(data["latitude_deg"], GRID_SHAPE[1])
checkpoint_path = ROOT / config["paths"]["checkpoint"]
resumed_states, resumed_metrics = {}, {}
if args.resume and checkpoint_path.exists():
try:
previous = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
except TypeError:
previous = torch.load(checkpoint_path, map_location="cpu")
if previous.get("format_version") != FORMAT_VERSION or tuple(previous.get("grid_shape", ())) != GRID_SHAPE:
raise ValueError("checkpoint protocol mismatch")
restored_folds = previous.get("model", {}).get("folds", previous.get("folds", []))
resumed_states = {state["fold"]: state for state in restored_folds}
resumed_metrics = {index: value for index, value in enumerate(previous.get("training_metrics", []))}
local_states, local_metrics = {}, {}
for fold in range(rank, 21, world):
if fold in resumed_states and fold in resumed_metrics:
continue
state, metrics = train_fold(fold, x, y, scenario_ids, config, weights)
local_states[fold], local_metrics[fold] = state, metrics
print(f"rank={rank} fold={fold:02d} scenario={scenario_ids[fold]} alpha={state['ridge_alpha']}")
if distributed:
gathered_states, gathered_metrics = [None] * world, [None] * world
torch.distributed.all_gather_object(gathered_states, local_states)
torch.distributed.all_gather_object(gathered_metrics, local_metrics)
states = {**resumed_states, **{key: value for item in gathered_states for key, value in item.items()}}
metrics = {**resumed_metrics, **{key: value for item in gathered_metrics for key, value in item.items()}}
else:
states, metrics = {**resumed_states, **local_states}, {**resumed_metrics, **local_metrics}
if rank == 0:
if sorted(states) != list(range(21)):
raise RuntimeError("all 21 LOSO folds must be trained")
fold_states = [states[index] for index in range(21)]
fold_metrics = [metrics[index] for index in range(21)]
checkpoint = {"epoch": 1,
"model": {"folds": fold_states},
"model_config": config["model"],
"optimizer_state_dict": {"folds": [state["gpr_optimizer"] for state in fold_states]},
"loss": float(np.mean([item["gpr_final_loss"] for item in fold_metrics])),
"config": config,
"model_name": MODEL_NAME, "format_version": FORMAT_VERSION,
"grid_shape": GRID_SHAPE, "scenario_ids": scenario_ids,
"short_response": torch.from_numpy(x_numpy), "long_response": torch.from_numpy(y_numpy),
"training_metrics": fold_metrics,
"gpr_ard": False, "gpr_ard_gap": config["model"]["gpr"]["ard_gap"]}
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
torch.save(checkpoint, checkpoint_path)
metrics_path = ROOT / config["paths"]["training_metrics"]
metrics_path.parent.mkdir(parents=True, exist_ok=True)
metrics_path.write_text(json.dumps({"fold_count": 21, "world_size": world,
"folds": fold_metrics}, indent=2) + "\n")
print(f"checkpoint={checkpoint_path.relative_to(ROOT)} folds=21 world_size={world} resumed={len(resumed_states)}")
if distributed:
torch.distributed.barrier(); torch.distributed.destroy_process_group()
if __name__ == "__main__":
main()
|