OneScience-Group/FourCastNet
040
1import sys2from pathlib import Path3 4# 获取项目根目录(train.py上级的上级)5root_path = Path(__file__).parent.parent6sys.path.append(str(root_path))7import torch8import os9import numpy as np10import torch.distributed as dist11import logging12import time13 14from torch.nn.parallel import DistributedDataParallel15from model.fourcastnet import FourCastNet16from onescience.datapipes.climate import ERA5Datapipe17from onescience.utils.YParams import YParams18from onescience.utils.fcn.darcy_loss import LpLoss19 20from apex import optimizers21 22 23def main():24 logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")25 logger = logging.getLogger()26 27 ## Model config init28 config_file_path = os.path.join(current_path, "conf/config.yaml")29 cfg = YParams(config_file_path, "model")30 31 ## Distributed config init32 cfg.world_size = 133 if "WORLD_SIZE" in os.environ:34 cfg.world_size = int(os.environ["WORLD_SIZE"])35 world_rank = 036 local_rank = 037 if cfg.world_size > 1:38 dist.init_process_group(backend="nccl", init_method="env://")39 local_rank = int(os.environ["LOCAL_RANK"])40 world_rank = dist.get_rank()41 42 ## DataLoader init43 cfg_data = YParams(config_file_path, "datapipe")44 cfg['N_in_channels'] = len(cfg_data.dataset.channels)45 cfg['N_out_channels'] = len(cfg_data.dataset.channels)46 datapipe = ERA5Datapipe(47 dataset_dir=cfg_data.dataset.data_dir,48 used_variables=cfg_data.dataset.channels,49 used_years=cfg_data.dataset.train_time,50 distributed=dist.is_initialized()51 )52 train_dataloader, train_sampler = datapipe.get_dataloader("train")53 datapipe = ERA5Datapipe(54 dataset_dir=cfg_data.dataset.data_dir,55 used_variables=cfg_data.dataset.channels,56 used_years=cfg_data.dataset.val_time,57 distributed=dist.is_initialized()58 )59 val_dataloader, val_sampler = datapipe.get_dataloader("valid")60 61 # Model init62 model = FourCastNet().to(local_rank)63 optimizer = optimizers.FusedAdam(model.parameters(), lr=cfg.lr)64 scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.2, patience=5, mode='min')65 loss_obj = LpLoss()66 67 ## Train process init68 os.makedirs(cfg.checkpoint_dir, exist_ok=True)69 train_loss_file = f"{cfg.checkpoint_dir}/trloss.npy"70 valid_loss_file = f"{cfg.checkpoint_dir}/valoss.npy"71 best_valid_loss = 1.0e672 best_loss_epoch = 073 train_losses = np.empty((0,), dtype=np.float32)74 valid_losses = np.empty((0,), dtype=np.float32)75 76 ## Get model params count77 if cfg.world_size == 1:78 total_params = sum(p.numel() for p in model.parameters())79 print("\n\n")80 print("-" * 50)81 print(f"📂 now params is {total_params}, {total_params / 1e6:.2f}M, {total_params / 1e9:.2f}B")82 print("-" * 50, "\n")83 84 ## Load model weight if there exist well-trained model 85 if os.path.exists(f"{cfg.checkpoint_dir}/model_bak.pth"):86 if world_rank == 0:87 print("\n\n")88 print("-" * 50)89 print(f"✅ There has a model weight, load and continue training...")90 print(f'If you want to train a new model, ensure there is no *.pth file in {cfg.checkpoint_dir}')91 print("-" * 50, "\n")92 ckpt = torch.load(f"{cfg.checkpoint_dir}/model_bak.pth", map_location=f'cuda:{local_rank}', weights_only=False)93 model.load_state_dict(ckpt["model_state_dict"])94 optimizer.load_state_dict(ckpt["optimizer_state_dict"])95 scheduler.load_state_dict(ckpt["scheduler_state_dict"])96 best_valid_loss = ckpt["best_valid_loss"]97 best_loss_epoch = ckpt["best_loss_epoch"]98 train_losses = np.load(train_loss_file)99 valid_losses = np.load(valid_loss_file)100 101 ## Distributed model102 if cfg.world_size > 1:103 model = DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank)104 world_rank == 0 and logger.info(f"start training ...")105 106 for epoch in range(cfg.max_epoch):107 if dist.is_initialized():108 train_sampler.set_epoch(epoch)109 val_sampler.set_epoch(epoch)110 model.train()111 train_loss = 0112 start_time = time.time()113 for j, data in enumerate(train_dataloader):114 invar = data[0].to(local_rank, dtype=torch.float32)115 outvar = data[1].to(local_rank, dtype=torch.float32)116 invar = invar[:, :, :-1, :]117 outvar = outvar[:, :, :-1, :]118 outvar_pred = model(invar)119 loss = loss_obj(outvar, outvar_pred)120 optimizer.zero_grad()121 loss.backward()122 optimizer.step()123 train_loss += loss.item()124 if world_rank == 0:125 logger.info(f'Train: Epoch {epoch}-{j+1}/{len(train_dataloader)} '126 f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] '127 f'[{(time.time()-start_time)/(j+1): .02f}s/{cfg_data.dataloader.batch_size}batch] '128 f'loss:{train_loss / (j+1): .04f}')129 130 train_loss /= len(train_dataloader)131 132 model.eval()133 valid_loss = 0134 with torch.no_grad():135 start_time = time.time()136 for j, data in enumerate(val_dataloader):137 invar = data[0].to(local_rank, dtype=torch.float32)138 outvar = data[1].to(local_rank, dtype=torch.float32)139 invar = invar[:, :, :-1, :]140 outvar = outvar[:, :, :-1, :]141 outvar_pred = model(invar)142 loss = loss_obj(outvar, outvar_pred)143 144 if cfg.world_size > 1:145 loss_tensor = loss.detach().to(local_rank) # torch.tensor(loss, device=local_rank)146 dist.all_reduce(loss_tensor)147 loss = loss_tensor.item() / cfg.world_size148 valid_loss += loss149 else:150 valid_loss += loss.item() 151 if world_rank == 0:152 logger.info(f'Valid: Epoch {epoch}-{j+1}/{len(val_dataloader)} '153 f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] '154 f'[{(time.time()-start_time)/(j+1): .02f}s/{cfg_data.dataloader.batch_size}batch] '155 f'loss:{valid_loss / (j+1): .04f}')156 157 valid_loss /= len(val_dataloader)158 is_save_ckp = False159 if valid_loss < best_valid_loss:160 best_valid_loss = valid_loss161 best_loss_epoch = epoch162 world_rank == 0 and save_checkpoint(model, optimizer, scheduler, best_valid_loss, best_loss_epoch, cfg.checkpoint_dir)163 is_save_ckp = True164 scheduler.step(valid_loss)165 166 if world_rank == 0:167 logger.info(f"Epoch [{epoch + 1}/{cfg.max_epoch}], "168 f"Train Loss: {train_loss:.4f}, "169 f"Valid Loss: {valid_loss:.4f}, "170 f"Best loss at Epoch: {best_loss_epoch + 1}"171 + (", saving checkpoint" if is_save_ckp else "")172 )173 train_losses = np.append(train_losses, train_loss)174 valid_losses = np.append(valid_losses, valid_loss)175 np.save(train_loss_file, train_losses)176 np.save(valid_loss_file, valid_losses)177 178 if epoch - best_loss_epoch > cfg.patience:179 print(f"Loss has not decrease in {cfg.patience} epochs, stopping training...")180 exit()181 182 183def save_checkpoint(model, optimizer, scheduler, best_valid_loss, best_loss_epoch, model_path):184 model_to_save = model.module if hasattr(model, "module") else model185 state = {"model_state_dict": model_to_save.state_dict(),186 "optimizer_state_dict": optimizer.state_dict(),187 "scheduler_state_dict": scheduler.state_dict(),188 "best_valid_loss": best_valid_loss,189 "best_loss_epoch": best_loss_epoch,190 }191 torch.save(state, f"{model_path}/model.pth")192 ### the weight file saving may interrupted due to DCU queue limit, get a backup to ensure there at least has one model193 os.system(f"mv {model_path}/model.pth {model_path}/model_bak.pth")194 195 196if __name__ == "__main__":197 current_path = os.getcwd()198 sys.path.append(current_path)199 main()200 