File size: 5,159 Bytes
6ff9439 | 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 | """Train independent PPNN replicas with Gaussian CRPS and optional rank allocation."""
import json
import os
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
from torch.utils.data import DataLoader, TensorDataset
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.ppnn import FORMAT_VERSION, PPNN, ensemble_features, gaussian_crps
def load_data(config):
data = np.load(ROOT / config["data"]["file"])
expected = (len(data["target"]), 50, 18)
if str(data["format_version"]) != FORMAT_VERSION or data["ensemble"].shape != expected:
raise ValueError(f"dataset version/shape mismatch: expected {FORMAT_VERSION} and {expected}")
if len(data["station_id"]) != 537 or int(data["lead_hours"]) != 48:
raise ValueError("dataset must retain 537 stations and 48h lead")
arrays = [data["ensemble"], data["auxiliary"], data["target"]]
if not all(np.isfinite(x).all() for x in arrays):
raise ValueError("dataset contains non-finite values")
return data
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
data = load_data(config)
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
device_name = config["runtime"]["device"]
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
enough_accelerators = torch.cuda.is_available() and torch.cuda.device_count() >= world_size
device = torch.device("cuda", local_rank) if device_name == "auto" and enough_accelerators else torch.device("cpu")
ensemble = torch.from_numpy(data["ensemble"])
auxiliary = torch.from_numpy(data["auxiliary"])
station = torch.from_numpy(data["station_index"])
target = torch.from_numpy(data["target"])
continuous_raw = ensemble_features(ensemble, auxiliary)
feature_mean = continuous_raw.mean(0)
feature_std = continuous_raw.std(0, correction=1).clamp_min(1e-6)
target_mean, target_std = target.mean(), target.std(correction=1).clamp_min(1e-6)
continuous = (continuous_raw - feature_mean) / feature_std
target_scaled = (target - target_mean) / target_std
dataset = TensorDataset(continuous, station, target_scaled)
local = []
replicas = int(config["train"]["replicas"])
for replica in range(replicas):
if replica % world_size != rank:
continue
seed = int(config["seed"]) + replica
torch.manual_seed(seed)
generator = torch.Generator().manual_seed(seed)
loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=True, generator=generator)
model = PPNN(int(config["model"]["hidden_size"]), eps=float(config["model"]["sigma_epsilon"])).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=float(config["train"]["learning_rate"]))
history = []
for epoch in range(int(config["train"]["epochs"])):
total = 0.0
for x, s, y in loader:
mu, sigma = model(x.to(device), s.to(device))
loss = gaussian_crps(mu, sigma, y.to(device)).mean()
if not torch.isfinite(loss):
raise RuntimeError("non-finite Gaussian CRPS loss")
optimizer.zero_grad(set_to_none=True)
loss.backward()
optimizer.step()
total += float(loss.detach()) * len(y)
history.append(total / len(dataset))
local.append({"replica": replica, "state_dict": {k: v.cpu() for k, v in model.state_dict().items()}, "loss": history})
print(f"rank={rank} replica={replica} final_scaled_crps={history[-1]:.6f}")
if distributed:
gathered = [None] * world_size
torch.distributed.all_gather_object(gathered, local)
trained = [item for rank_items in gathered for item in rank_items]
else:
trained = local
if rank == 0:
trained.sort(key=lambda x: x["replica"])
if [x["replica"] for x in trained] != list(range(replicas)):
raise RuntimeError("distributed ranks did not return every replica")
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)
scaling = {"feature_mean": feature_mean, "feature_std": feature_std, "target_mean": target_mean, "target_std": target_std}
states = [x["state_dict"] for x in trained]
torch.save({"format_version": FORMAT_VERSION, "model": states,
"model_config": config["model"], "config": config, "scaling": scaling}, checkpoint)
metrics.write_text(json.dumps({"replicas": [{"id": x["replica"], "loss": x["loss"]} for x in trained], "world_size": world_size}, indent=2) + "\n")
if distributed:
torch.distributed.destroy_process_group()
if __name__ == "__main__":
main()
|