| """Train independently initialized DD members with optional DDP.""" |
|
|
| import copy |
| import json |
| import os |
| import sys |
| from pathlib import Path |
|
|
| import numpy as np |
| import torch |
| from torch import nn |
| from torch.nn.parallel import DistributedDataParallel |
| from torch.utils.data import DataLoader, DistributedSampler, TensorDataset |
|
|
|
|
| ROOT = Path(__file__).resolve().parents[1] |
| sys.path.insert(0, str(ROOT)) |
| from model.precipdd import DATA_FORMAT_VERSION, PrecipDD, load_config, seed_all, validate_archive |
|
|
|
|
| def main(): |
| config = load_config(ROOT / "conf/config.yaml") |
| runtime, settings = config["runtime"], config["training"] |
| distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1 |
| local_rank = int(os.environ.get("LOCAL_RANK", "0")) |
| local_world_size = int(os.environ.get("LOCAL_WORLD_SIZE", "1")) |
| use_cuda = (torch.cuda.is_available() and runtime["device"] != "cpu" |
| and torch.cuda.device_count() >= local_world_size) |
| if distributed: |
| torch.distributed.init_process_group(runtime["ddp_backend_gpu"] if use_cuda else runtime["ddp_backend_cpu"]) |
| rank = torch.distributed.get_rank() if distributed else 0 |
| device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu") |
| if use_cuda: |
| torch.cuda.set_device(device) |
| data = np.load(ROOT / config["paths"]["data"]) |
| validate_archive(data) |
| train_mask, val_mask = data["split"] == 0, data["split"] == 1 |
| dataset = TensorDataset(torch.from_numpy(data["precipitation"][train_mask]).float(), torch.from_numpy(data["agmt"][train_mask]).float()) |
| val_x = torch.from_numpy(data["precipitation"][val_mask]).float().to(device) |
| val_y = torch.from_numpy(data["agmt"][val_mask]).float().to(device) |
| ensemble_states, histories = [], [] |
| for member in range(settings["ensemble_members"]): |
| seed_all(config["project"]["seed"] + member) |
| sampler = DistributedSampler(dataset, shuffle=True, seed=member) if distributed else None |
| loader = DataLoader(dataset, batch_size=settings["batch_size"], shuffle=sampler is None, sampler=sampler, |
| num_workers=settings["num_workers"]) |
| model = PrecipDD(config["model"]["filters"], config["model"]["dense_units"]).to(device) |
| wrapped = DistributedDataParallel(model, device_ids=[local_rank] if use_cuda else None) if distributed else model |
| optimizer = torch.optim.Adam(wrapped.parameters(), lr=settings["learning_rate"], weight_decay=settings["l2_weight_decay"]) |
| history, best_loss, best_state = [], float("inf"), None |
| for epoch in range(settings["epochs"]): |
| if sampler is not None: |
| sampler.set_epoch(member * settings["epochs"] + epoch) |
| wrapped.train() |
| total, count = 0.0, 0 |
| for fields, target in loader: |
| fields, target = fields.to(device), target.to(device) |
| optimizer.zero_grad(set_to_none=True) |
| loss = nn.functional.l1_loss(wrapped(fields), target) |
| loss.backward() |
| optimizer.step() |
| total += float(loss.detach()) * len(fields) |
| count += len(fields) |
| summary = torch.tensor([total, count], dtype=torch.float64, device=device) |
| if distributed: |
| torch.distributed.all_reduce(summary) |
| wrapped.eval() |
| with torch.no_grad(): |
| val_loss = float(nn.functional.l1_loss(model(val_x), val_y)) |
| history.append({"epoch": epoch + 1, "train_mae": float(summary[0] / summary[1]), "validation_mae": val_loss}) |
| if val_loss < best_loss: |
| best_loss = val_loss |
| best_state = {key: value.detach().cpu().clone() for key, value in model.state_dict().items()} |
| ensemble_states.append(best_state) |
| histories.append(history) |
| if rank == 0: |
| checkpoint_path = ROOT / config["paths"]["checkpoint"] |
| checkpoint_path.parent.mkdir(parents=True, exist_ok=True) |
| torch.save({"format_version": DATA_FORMAT_VERSION, "ensemble_states": ensemble_states, |
| "model_config": {"filters": config["model"]["filters"], "dense_units": config["model"]["dense_units"], |
| "input_shape": [1, 55, 160], "feature_shape": [16, 14, 40], "flattened_features": 8960}, |
| "training_config": settings, "histories": histories, "world_size": torch.distributed.get_world_size() if distributed else 1}, checkpoint_path) |
| metrics_path = ROOT / config["paths"]["training_metrics"] |
| metrics_path.parent.mkdir(parents=True, exist_ok=True) |
| metrics_path.write_text(json.dumps({"ensemble_members": len(histories), "history": histories}, indent=2) + "\n", encoding="utf-8") |
| print(f"checkpoint={checkpoint_path.relative_to(ROOT)} members={len(ensemble_states)}") |
| if distributed: |
| torch.distributed.barrier() |
| torch.distributed.destroy_process_group() |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|