"""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()