File size: 9,810 Bytes
b300acf | 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 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 | """Fit stable target clusters, NumPy RF, and neural classifiers."""
import argparse
import json
import os
import random
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
from torch import nn
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, TensorDataset
from torch.utils.data.distributed import DistributedSampler
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.cesm_seasonal_ml import (CHECKPOINT_FORMAT_VERSION, DATA_FORMAT_VERSION, FeedForwardNN,
NumpyGradientBoostedTrees, NumpyRandomForest, SeasonalLSTM,
StablePrecipKMeans, classification_metrics)
def setup(config):
distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
local_world = int(os.environ.get("LOCAL_WORLD_SIZE", "1"))
use_cuda = (torch.cuda.is_available() and config["runtime"]["device"] != "cpu"
and torch.cuda.device_count() >= local_world)
if distributed:
torch.distributed.init_process_group(config["runtime"]["ddp_backend_gpu"] if use_cuda else config["runtime"]["ddp_backend_cpu"])
rank = torch.distributed.get_rank() if distributed else 0
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', 0))}" if use_cuda else "cpu")
if use_cuda:
torch.cuda.set_device(device)
return distributed, rank, device
def train_neural(name, model, x_train, y_train, x_val, y_val, settings, distributed, device):
dataset = TensorDataset(torch.from_numpy(x_train).float(), torch.from_numpy(y_train).long())
sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
local_samples = len(sampler) if sampler is not None else len(dataset)
loader = DataLoader(dataset, batch_size=settings["batch_size"], shuffle=sampler is None,
sampler=sampler, drop_last=name == "nn" and local_samples % settings["batch_size"] == 1)
model = model.to(device)
wrapped = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None) if distributed else model
counts = np.bincount(y_train, minlength=4)
weights = len(y_train) / np.maximum(4 * counts, 1)
criterion = nn.CrossEntropyLoss(weight=torch.tensor(weights, dtype=torch.float32, device=device))
optimizer = torch.optim.Adam(wrapped.parameters(), lr=settings["learning_rate"])
history, best_accuracy, best_state = [], -1.0, None
for epoch in range(settings["epochs"]):
if sampler:
sampler.set_epoch(epoch)
wrapped.train()
total_loss, total_samples = 0.0, 0
for features, target in loader:
optimizer.zero_grad(set_to_none=True)
loss = criterion(wrapped(features.to(device)), target.to(device))
loss.backward()
optimizer.step()
total_loss += float(loss.detach()) * len(features)
total_samples += len(features)
loss_summary = torch.tensor([total_loss, total_samples], dtype=torch.float64, device=device)
if distributed:
torch.distributed.all_reduce(loss_summary)
wrapped.eval()
with torch.no_grad():
logits = wrapped(torch.from_numpy(x_val).float().to(device))
correct = (logits.argmax(1).cpu().numpy() == y_val).sum()
count = len(y_val)
if distributed:
summary = torch.tensor([correct, count], dtype=torch.float64, device=device)
torch.distributed.all_reduce(summary)
correct, count = summary.tolist()
accuracy = correct / count
history.append({"epoch": epoch + 1, "loss": float(loss_summary[0] / loss_summary[1]), "validation_accuracy": accuracy})
if accuracy > best_accuracy:
best_accuracy = accuracy
module = wrapped.module if distributed else wrapped
best_state = {key: value.detach().cpu().clone() for key, value in module.state_dict().items()}
return best_state, history
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--season", choices=["NDJ", "JFM"])
parser.add_argument("--models", nargs="+", choices=["rf", "nn", "lstm", "xgboost"])
args = parser.parse_args()
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
season = args.season or config["season"]
enabled = args.models or config["model"]["enabled"]
seed = int(config["seed"])
random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
distributed, rank, device = setup(config)
data = np.load(ROOT / config["data"]["path"])
if str(data["format_version"]) != DATA_FORMAT_VERSION:
raise ValueError("data format version mismatch")
season_index = [str(x) for x in data["seasons"]].index(season)
train, validation, test = data["split"] == 0, data["split"] == 1, data["split"] == 2
cluster = StablePrecipKMeans(seed=seed).fit(data["precipitation"][season_index, train])
labels = cluster.transform(data["precipitation"][season_index])
models, histories, validation_probabilities = {}, {}, {}
if "rf" in enabled:
rf_config = config["model"]["rf"]
forest = NumpyRandomForest(rf_config["trees"], rf_config["mtry"], rf_config["max_depth"], rf_config["min_samples_split"], seed)
if rank == 0:
forest.fit(data["rf_features"][season_index, train], labels[train])
validation_probabilities["rf"] = forest.predict_proba(data["rf_features"][season_index, validation])
importance = forest.permutation_importance(data["rf_features"][season_index, validation], labels[validation], rf_config["permutation_repeats"])
mean_minimum_depth, root_frequency = forest.tree_structure_importance()
models["rf"] = forest
histories["rf"] = {"permutation_importance": importance.tolist(), "mean_minimum_depth": mean_minimum_depth.tolist(),
"root_frequency": root_frequency.tolist(), "feature_names": data["rf_manifest"].tolist()}
if "xgboost" in enabled and rank == 0:
xgb_config = config["model"]["xgboost"]
boosted = NumpyGradientBoostedTrees(xgb_config["n_rounds"], xgb_config["max_depth"],
xgb_config["eta"], xgb_config["gamma"], seed)
boosted.fit(data["rf_features"][season_index, train], labels[train])
models["xgboost"] = boosted
validation_probabilities["xgboost"] = boosted.predict_proba(data["rf_features"][season_index, validation])
histories["xgboost"] = {"rounds": xgb_config["n_rounds"]}
history_length = 4 if season == "NDJ" else 12
specifications = {
"nn": (FeedForwardNN(hidden=tuple(config["model"]["nn"]["hidden"]), dropout=config["model"]["nn"]["dropout"]), data["nn_features"][season_index], config["model"]["nn"]),
"lstm": (SeasonalLSTM(hidden_size=config["model"]["lstm"]["hidden_size"], dense_size=config["model"]["lstm"]["dense_size"], dropout=config["model"]["lstm"]["dropout"]), data["eof_sequence"][season_index, :, -history_length:], config["model"]["lstm"])
}
for name in ("nn", "lstm"):
if name in enabled:
model, features, settings = specifications[name]
state, history = train_neural(name, model, features[train], labels[train], features[validation], labels[validation], settings, distributed, device)
model.load_state_dict(state); model.to(device).eval()
with torch.no_grad():
probability = torch.softmax(model(torch.from_numpy(features[validation]).float().to(device)), dim=1).cpu().numpy()
models[name] = state
histories[name] = history
validation_probabilities[name] = probability
if rank == 0:
metrics = {name: classification_metrics(labels[validation], probability, config["evaluation"]["random_trials"], seed) for name, probability in validation_probabilities.items()}
model_config = {"enabled": list(enabled), "primary": config["model"]["primary"], "settings": config["model"],
"history_length": history_length, "classes": 4, "rf_features": 103,
"nn_features": 416, "eof_channels": 28}
data_spec = {"format_version": str(data["format_version"]), "season": season,
"seasons": data["seasons"].tolist(), "rf_manifest": data["rf_manifest"].tolist(),
"nn_manifest": data["nn_manifest"].tolist(), "rf_shape": list(data["rf_features"].shape),
"nn_shape": list(data["nn_features"].shape), "eof_shape": list(data["eof_sequence"].shape),
"precipitation_shape": list(data["precipitation"].shape), "class_names": data["class_names"].tolist()}
checkpoint = {"format_version": CHECKPOINT_FORMAT_VERSION, "model": models, "model_config": model_config,
"data_spec": data_spec, "season": season, "cluster": cluster.state_dict(),
"training_history": histories, "validation_metrics": metrics, "paper_model": config["paper_model"]}
path = ROOT / config["paths"]["checkpoint"]
path.parent.mkdir(parents=True, exist_ok=True)
torch.save(checkpoint, path)
metrics_path = ROOT / config["paths"]["training_metrics"]
metrics_path.parent.mkdir(parents=True, exist_ok=True)
metrics_path.write_text(json.dumps({"season": season, "models": enabled, "validation": metrics}, indent=2) + "\n")
print(f"checkpoint={path.relative_to(ROOT)} season={season} models={','.join(enabled)}")
if distributed:
torch.distributed.barrier(); torch.distributed.destroy_process_group()
if __name__ == "__main__":
main()
|