| from pathlib import Path |
| import sys,os |
| import numpy as np,torch |
| import torch.nn.functional as F |
| from torch.nn.parallel import DistributedDataParallel |
| ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT)) |
| from model.precip_extremes_gan import * |
| c=load_config(ROOT);rank=int(os.environ.get("RANK",0));world=int(os.environ.get("WORLD_SIZE",1));local=int(os.environ.get("LOCAL_RANK",0));distributed=world>1;use_cuda=torch.cuda.is_available() and torch.cuda.device_count()>=world and c["runtime"]["device"]!="cpu" |
| if distributed:torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo") |
| device=torch.device(f"cuda:{local}" if use_cuda else "cpu");torch.manual_seed(c["seed"]+rank);torch.set_num_threads(2);d=np.load(ROOT/c["data"]["path"]);m=d["split"]=="train";x=torch.from_numpy(d["predictors"][m]).to(device);y=torch.log(torch.from_numpy(d["precipitation"][m])+c["data"]["log_epsilon"]).to(device);net=PrecipExtremesGAN(c["data"]["input_channels"],c["model"]["filters"],c["model"]["noise_channels"]).to(device);critic=PatchCritic(c["data"]["input_channels"]).to(device) |
| baseline=DistributedDataParallel(net.baseline,device_ids=[local] if use_cuda else None) if distributed else net.baseline;generator=DistributedDataParallel(net.generator,device_ids=[local] if use_cuda else None) if distributed else net.generator;critic_train=DistributedDataParallel(critic,device_ids=[local] if use_cuda else None) if distributed else critic |
| ob=torch.optim.Adam(baseline.parameters(),lr=c["train"]["baseline_lr"]) |
| for _ in range(c["train"]["baseline_epochs"]): pred=baseline(x,c["data"]["high_grid"]);loss=F.mse_loss(pred,y);ob.zero_grad();loss.backward();ob.step() |
| og=torch.optim.Adam(generator.parameters(),lr=c["train"]["gan_lr"]);oc=torch.optim.Adam(critic_train.parameters(),lr=c["train"]["gan_lr"]) |
| for _ in range(c["train"]["gan_epochs"]): |
| with torch.no_grad():base=baseline(x,c["data"]["high_grid"]) |
| fake=base+generator(torch.cat((x,torch.randn(len(x),1,*x.shape[-2:],device=device)),1),c["data"]["high_grid"]);cl=F.softplus(critic_train(x,fake.detach())).mean()+F.softplus(-critic_train(x,y)).mean();oc.zero_grad();cl.backward();oc.step();adv=F.softplus(-critic_train(x,fake)).mean();intensity=F.mse_loss(fake.mean((2,3)),y.mean((2,3)));gl=F.mse_loss(fake.mean(0),y.mean(0))+c["train"]["adversarial_weight"]*adv+intensity;og.zero_grad();gl.backward();og.step() |
| vals=torch.tensor([float(loss),float(gl),float(cl)],device=device); |
| if distributed:torch.distributed.all_reduce(vals);vals/=world |
| if rank==0: |
| path=ROOT/c["paths"]["checkpoint"];path.parent.mkdir(parents=True,exist_ok=True);torch.save({"model":net.state_dict(),"critic":critic.state_dict(),"model_config":net.model_config,"format_version":c["data"]["format_version"]},path);write_json(ROOT/c["paths"]["training_metrics"],{"baseline_mse":float(vals[0]),"generator_loss":float(vals[1]),"critic_loss":float(vals[2]),"world_size":world});print(path) |
| if distributed:torch.distributed.destroy_process_group() |
|
|