PrecipDD / scripts /train.py
zhangrenchao's picture
Publish PrecipDD reproduction
950fc23 verified
Raw
History Blame Contribute Delete
5.02 kB
"""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()