File size: 5,800 Bytes
30f7852 | 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 | """Train UnetDif with single-process or torchrun DDP execution."""
import json
import os
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, Dataset, DistributedSampler
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.unetdif import UnetDif, loss_components
class RainfallDataset(Dataset):
def __init__(self, path: Path, config):
archive = np.load(path, allow_pickle=False)
expected = (int(config["time_steps"]), int(config["channels"]), int(config["height"]), int(config["width"]))
if archive["inputs"].shape[1:] != expected or archive["targets"].shape[1:] != (expected[0], expected[2], expected[3]):
raise ValueError(f"expected day shapes {expected} and {(expected[0], expected[2], expected[3])}")
if len(np.unique(archive["group_id"])) != len(archive["group_id"]):
raise ValueError("each rain day must belong to one split group")
if archive["lead_hours"].tolist() != config["lead_hours"]:
raise ValueError("NPZ lead_hours do not match the configured eight 3-hour steps")
self.inputs = archive["inputs"].astype(np.float32, copy=False).reshape(-1, *expected[1:])
self.targets = archive["targets"].astype(np.float32, copy=False).reshape(-1, expected[2], expected[3])
def __len__(self):
return len(self.inputs)
def __getitem__(self, index):
return torch.from_numpy(self.inputs[index]), torch.from_numpy(self.targets[index])
def get_device(local_rank=0, local_world_size=1):
use_cuda = torch.cuda.is_available() and torch.cuda.device_count() >= local_world_size
return torch.device("cuda", local_rank) if use_cuda else torch.device("cpu")
def reduce_values(values, device, distributed):
tensor = torch.tensor(values, device=device)
if distributed:
torch.distributed.all_reduce(tensor)
return tensor.cpu().tolist()
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
torch.manual_seed(int(config["seed"]))
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"))
device = get_device(local_rank, local_world_size)
if distributed:
torch.distributed.init_process_group("nccl" if device.type == "cuda" else "gloo")
rank = torch.distributed.get_rank() if distributed else 0
dataset = RainfallDataset(ROOT / config["data"]["root"] / "train.npz", config["data"])
if distributed and len(dataset) < torch.distributed.get_world_size():
raise ValueError("DDP requires at least one training item per process")
sampler = DistributedSampler(dataset) if distributed else None
loader = DataLoader(dataset, batch_size=int(config["train"]["batch_size"]), shuffle=sampler is None,
sampler=sampler, num_workers=int(config["train"]["num_workers"]))
model = UnetDif(**config["model"]).to(device)
if distributed:
model = DDP(model, device_ids=[local_rank] if device.type == "cuda" else None)
optimizer = torch.optim.Adam(model.parameters(), lr=float(config["train"]["learning_rate"]))
names = ["dry_focal", "false_alarm_focal", "positive_mse", "negative_mse", "dry_mae", "all_mae"]
normalizers = {name: 1.0 for name in names}
history = []
for epoch in range(int(config["train"]["epochs"])):
if sampler is not None:
sampler.set_epoch(epoch)
model.train(); sums = {name: 0.0 for name in names}; batches = 0
for inputs, target in loader:
inputs, target = inputs.to(device), target.to(device)
heads = model(inputs)
raw = loss_components(heads, inputs[:, int(config["data"]["precipitation_channel"])], target,
float(config["data"]["rain_threshold_mm_3h"]),
float(config["loss"]["focal_alpha"]), float(config["loss"]["focal_gamma"]))
loss = sum(raw[name] * normalizers[name] for name in names)
optimizer.zero_grad(set_to_none=True); loss.backward(); optimizer.step()
for name in names:
sums[name] += float(raw[name].detach())
batches += 1
reduced = reduce_values([sums[name] for name in names] + [batches], device, distributed)
total_batches = max(reduced[-1], 1)
averages = {name: reduced[index] / total_batches for index, name in enumerate(names)}
if epoch == 0:
normalizers = {name: 1.0 / max(value, 1e-6) for name, value in averages.items()}
record = {"epoch": epoch + 1, "components": averages, "normalizers": normalizers}
history.append(record)
if rank == 0:
print(json.dumps(record))
checkpoint = ROOT / config["paths"]["checkpoint"]
checkpoint.parent.mkdir(parents=True, exist_ok=True)
state = model.module.state_dict() if distributed else model.state_dict()
torch.save({"epoch": epoch + 1, "model": state, "optimizer": optimizer.state_dict(),
"model_config": config["model"], "loss_normalizers": normalizers,
"loss_components": names,
"format_version": config["data"]["format_version"]}, checkpoint)
if rank == 0:
metrics = ROOT / config["paths"]["training_metrics"]
metrics.parent.mkdir(parents=True, exist_ok=True)
metrics.write_text(json.dumps({"history": history}, indent=2) + "\n")
if distributed:
torch.distributed.destroy_process_group()
if __name__ == "__main__":
main()
|