zhangrenchao's picture
Publish CRAI-ClimateExtremes reproduction
b20ca9c verified
Raw
History Blame Contribute Delete
4.88 kB
"""Train independent CRAI ensemble members, optionally under torchrun DDP."""
from pathlib import Path
import argparse
import json
import os
import random
import sys
import numpy as np
import torch
from torch import distributed as dist
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler
import yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.crai_climateextremes import CRAIClimateExtremes
class ClimateDataset(Dataset):
def __init__(self, archive, limit):
self.observed = torch.from_numpy(archive["observed"][:limit])
self.valid = torch.from_numpy(archive["valid_mask"][:limit])
self.target = torch.from_numpy(archive["target"][:limit])
land = torch.from_numpy(archive["europe_mask"])[None, None]
self.missing = land * (1.0 - self.valid)
def __len__(self):
return len(self.target)
def __getitem__(self, index):
return torch.cat((self.observed[index], self.valid[index]), 0), self.target[index], self.missing[index]
def load_config(path):
with open(path, encoding="utf-8") as handle:
return yaml.safe_load(handle)
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml")
parser.add_argument("--paper-model", action="store_true")
args = parser.parse_args()
config_path = args.config if args.config.is_absolute() else ROOT / args.config
cfg = load_config(config_path)
use_paper = args.paper_model or cfg["paper_model"]
batch_size = cfg["paper_batch_size"] if use_paper else cfg["batch_size"]
iterations = cfg["paper_iterations"] if use_paper else cfg["max_iterations"]
members = cfg["paper_ensemble_members"] if use_paper else cfg["ensemble_members"]
rank, world = int(os.getenv("RANK", 0)), int(os.getenv("WORLD_SIZE", 1))
if world > 1:
dist.init_process_group("nccl" if torch.cuda.is_available() else "gloo")
device = torch.device(f"cuda:{int(os.getenv('LOCAL_RANK', 0))}" if torch.cuda.is_available() else "cpu")
if device.type == "cuda":
torch.cuda.set_device(device)
archive = np.load(ROOT / cfg["data_path"])
dataset = ClimateDataset(archive, cfg["num_samples"])
checkpoint_path = ROOT / cfg["checkpoint_path"]
if rank == 0:
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
(ROOT / "result/training").mkdir(parents=True, exist_ok=True)
if world > 1:
dist.barrier()
records, member_states = [], []
for member in range(members):
seed = int(cfg["seed"]) + member
random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
sampler = DistributedSampler(dataset, shuffle=True, seed=seed) if world > 1 else None
loader = DataLoader(dataset, batch_size=batch_size, shuffle=sampler is None, sampler=sampler)
model = CRAIClimateExtremes(cfg["base_channels"]).to(device)
if world > 1:
model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None)
optimizer = torch.optim.Adam(model.parameters(), lr=cfg["learning_rate"])
step, losses = 0, []
for epoch in range(cfg["epochs"] if not use_paper else 10**9):
if sampler is not None:
sampler.set_epoch(epoch)
for inputs, target, missing in loader:
inputs, target, missing = inputs.to(device), target.to(device), missing.to(device)
prediction = model(inputs)
loss = (torch.abs(prediction - target) * missing).sum() / missing.sum().clamp_min(1)
optimizer.zero_grad(); loss.backward(); optimizer.step()
losses.append(float(loss.detach()))
step += 1
if step >= iterations:
break
if step >= iterations:
break
raw_model = model.module if isinstance(model, DistributedDataParallel) else model
if rank == 0:
member_states.append({key: value.detach().cpu() for key, value in raw_model.state_dict().items()})
records.append({"member": member, "iterations": step, "final_missing_mae": losses[-1]})
if rank == 0:
torch.save(
{
"format_version": "1.0",
"model_config": {"base_channels": cfg["base_channels"]},
"model": member_states,
},
checkpoint_path,
)
payload = {"paper_model": bool(use_paper), "world_size": world, "members": records}
(ROOT / "result/training/metrics.json").write_text(json.dumps(payload, indent=2) + "\n")
print(json.dumps(payload))
if world > 1:
dist.destroy_process_group()
if __name__ == "__main__":
main()