ClimEmu-S2L / scripts /train.py
zhangrenchao's picture
Upload folder using huggingface_hub
87f2bd3 verified
Raw
History Blame Contribute Delete
7.22 kB
"""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()