OneScience-Group/RF-ClimParam
029
1"""Train four scale-specific pairs of joint multi-output random forests."""2 3import json4import os5import random6import sys7from pathlib import Path8 9import numpy as np10import torch11import yaml12 13 14ROOT = Path(__file__).resolve().parents[1]15sys.path.insert(0, str(ROOT))16from model.rf_climparam import FORMAT_VERSION, MODEL_NAME, build_pair17 18 19def load_scale(path, scale):20 data = np.load(path)21 expected = {"tend_inputs": 145, "tend_targets": 144, "diff_inputs": 62, "diff_targets": 17}22 if str(data["format_version"]) != FORMAT_VERSION or str(data["scale"]) != scale:23 raise ValueError(f"invalid metadata in {path}")24 count = len(data["tend_inputs"])25 ny, nx = map(int, data["grid_shape"])26 if count != ny * nx or int(data["snapshot_count"]) != 1:27 raise ValueError(f"{scale} must contain one complete [{ny},{nx}] snapshot")28 linear = data["grid_row"].astype(np.int64) * nx + data["grid_column"].astype(np.int64)29 if not np.array_equal(linear, np.arange(count)):30 raise ValueError(f"{scale} grid cannot be reversibly flattened")31 for name, width in expected.items():32 value = data[name]33 if value.shape != (count, width) or value.dtype != np.float32 or not np.isfinite(value).all():34 raise ValueError(f"{scale}/{name} requires finite float32 [{count},{width}]")35 if np.any(data["diff_targets"][:, :15] < 0):36 raise ValueError("Dbar training targets must be nonnegative")37 return data38 39 40def sample_training_columns(data, columns_per_latitude, seed):41 ny, nx = map(int, data["grid_shape"])42 if not 1 <= columns_per_latitude <= nx:43 raise ValueError("train_columns_per_latitude must be between 1 and grid width")44 rng = np.random.default_rng(seed)45 selected = []46 for row in range(ny):47 selected.extend(row * nx + rng.choice(nx, columns_per_latitude, replace=False))48 return np.asarray(selected, dtype=np.int64)49 50 51def main():52 config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())53 seed = int(config["seed"])54 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)55 distributed = int(os.environ.get("WORLD_SIZE", "1")) > 156 local_rank = int(os.environ.get("LOCAL_RANK", "0"))57 if distributed:58 torch.distributed.init_process_group("gloo")59 rank = torch.distributed.get_rank() if distributed else 060 world = torch.distributed.get_world_size() if distributed else 161 scales = list(config["data"]["scales"])62 assigned = [scale for i, scale in enumerate(scales) if i % world == rank]63 local_states, local_metrics = {}, {}64 for scale_index, scale in enumerate(assigned):65 data = load_scale(ROOT / config["data"]["root"] / f"{scale}.npz", scale)66 selected = sample_training_columns(67 data, int(config["data"]["train_columns_per_latitude"]), seed + scales.index(scale))68 pair = build_pair(config["model"], seed + scales.index(scale) * 100)69 pair["rf_tend"].fit(data["tend_inputs"][selected], data["tend_targets"][selected])70 pair["rf_diff"].fit(data["diff_inputs"][selected], data["diff_targets"][selected])71 local_states[scale] = {name: model.state_dict() for name, model in pair.items()}72 local_metrics[scale] = {"complete_grid_points": len(data["tend_inputs"]),73 "train_columns": len(selected),74 "columns_per_latitude": int(config["data"]["train_columns_per_latitude"])}75 print(f"rank={rank} trained={scale} sampled_columns={len(selected)}")76 if distributed:77 gathered_states, gathered_metrics = [None] * world, [None] * world78 torch.distributed.all_gather_object(gathered_states, local_states)79 torch.distributed.all_gather_object(gathered_metrics, local_metrics)80 states = {key: value for item in gathered_states for key, value in item.items()}81 metrics = {key: value for item in gathered_metrics for key, value in item.items()}82 else:83 states, metrics = local_states, local_metrics84 if rank == 0:85 if set(states) != set(scales):86 raise RuntimeError("not all scales were trained")87 checkpoint = {88 "model": states,89 "model_name": MODEL_NAME,90 "model_config": {"dimensions": {"rf_tend_input": 145, "rf_tend_output": 144,91 "rf_diff_input": 62, "rf_diff_output": 17},92 "engineering": config["model"]["engineering"],93 "paper_model": config["paper_model"], "scales": scales},94 "format_version": FORMAT_VERSION, "seed": seed}95 path = ROOT / config["paths"]["checkpoint"]96 path.parent.mkdir(parents=True, exist_ok=True)97 torch.save(checkpoint, path)98 metrics_path = ROOT / config["paths"]["training_metrics"]99 metrics_path.parent.mkdir(parents=True, exist_ok=True)100 metrics_path.write_text(json.dumps({"format_version": FORMAT_VERSION, "scales": metrics}, indent=2) + "\n")101 print(f"checkpoint={path.relative_to(ROOT)} scales={len(states)}")102 if distributed:103 torch.distributed.barrier(); torch.distributed.destroy_process_group()104 105 106if __name__ == "__main__":107 main()108 