OneScience-Group/GraphDOP
016
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 model.graphdop import GraphDOP15from onescience.datapipes.climate import ERA5Datapipe16from onescience.utils.YParams import YParams17 18try:19 from apex import optimizers20 _FUSED_ADAM = True21except Exception:22 _FUSED_ADAM = False23 24 25def main():26 logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")27 logger = logging.getLogger()28 29 ## Model config init30 config_file_path = os.path.join(current_path, "conf/config.yaml")31 cfg = YParams(config_file_path, "model")32 33 ## Distributed config init34 cfg.world_size = 135 if "WORLD_SIZE" in os.environ:36 cfg.world_size = int(os.environ["WORLD_SIZE"])37 world_rank = 038 local_rank = 039 if cfg.world_size > 1 and torch.cuda.is_available():40 dist.init_process_group(backend="nccl", init_method="env://")41 local_rank = int(os.environ["LOCAL_RANK"])42 world_rank = dist.get_rank()43 device = f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu"44 45 ## DataLoader init46 cfg_data = YParams(config_file_path, "datapipe")47 cfg['N_in_channels'] = len(cfg_data.dataset.channels)48 cfg['N_out_channels'] = len(cfg_data.dataset.channels)49 datapipe = ERA5Datapipe(50 dataset_dir=cfg_data.dataset.data_dir,51 used_variables=cfg_data.dataset.channels,52 used_years=cfg_data.dataset.train_time,53 distributed=dist.is_initialized(),54 input_steps=cfg.input_steps,55 output_steps=cfg.output_steps,56 batch_size=cfg_data.dataloader.batch_size,57 num_workers=cfg_data.dataloader.num_workers,58 )59 train_dataloader, train_sampler = datapipe.get_dataloader("train")60 datapipe = ERA5Datapipe(61 dataset_dir=cfg_data.dataset.data_dir,62 used_variables=cfg_data.dataset.channels,63 used_years=cfg_data.dataset.val_time,64 distributed=dist.is_initialized(),65 input_steps=cfg.input_steps,66 output_steps=cfg.output_steps,67 batch_size=cfg_data.dataloader.batch_size,68 num_workers=cfg_data.dataloader.num_workers,69 )70 val_dataloader, val_sampler = datapipe.get_dataloader("valid")71 72 # Model init73 model = GraphDOP(74 in_channels=cfg['N_in_channels'],75 out_channels=cfg['N_out_channels'],76 input_steps=cfg.input_steps,77 output_steps=cfg.output_steps,78 grid_shape=cfg.grid_shape,79 mesh_shape=cfg.mesh_shape,80 latent_dim=cfg.latent_dim,81 num_encoder_layers=cfg.num_encoder_layers,82 num_decoder_layers=cfg.num_decoder_layers,83 num_processor_blocks=cfg.num_processor_blocks,84 n_heads=cfg.n_heads,85 hidden_dim=cfg.hidden_dim,86 channel_weights=cfg.channel_weights,87 ).to(device)88 89 if _FUSED_ADAM:90 optimizer = optimizers.FusedAdam(model.parameters(), lr=cfg.lr)91 else:92 optimizer = torch.optim.Adam(model.parameters(), lr=cfg.lr)93 scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.2, patience=5, mode='min')94 95 ## Train process init96 os.makedirs(cfg.checkpoint_dir, exist_ok=True)97 train_loss_file = f"{cfg.checkpoint_dir}/trloss.npy"98 valid_loss_file = f"{cfg.checkpoint_dir}/valoss.npy"99 best_valid_loss = float("inf")100 best_loss_epoch = 0101 train_losses = np.empty((0,), dtype=np.float32)102 valid_losses = np.empty((0,), dtype=np.float32)103 104 ## Get model params count105 if cfg.world_size == 1:106 total_params = sum(p.numel() for p in model.parameters())107 print("\n\n")108 print("-" * 50)109 print(f"📂 now params is {total_params}, {total_params / 1e6:.2f}M, {total_params / 1e9:.2f}B")110 print("-" * 50, "\n")111 112 ## Load model weight if there exist well-trained model113 if os.path.exists(f"{cfg.checkpoint_dir}/model_bak.pth"):114 if world_rank == 0:115 print("\n\n")116 print("-" * 50)117 print(f"✅ There has a model weight, load and continue training...")118 print(f'If you want to train a new model, ensure there is no *.pth file in {cfg.checkpoint_dir}')119 print("-" * 50, "\n")120 ckpt = torch.load(f"{cfg.checkpoint_dir}/model_bak.pth", map_location=device, weights_only=False)121 model.load_state_dict(ckpt["model_state_dict"])122 optimizer.load_state_dict(ckpt["optimizer_state_dict"])123 scheduler.load_state_dict(ckpt["scheduler_state_dict"])124 best_valid_loss = ckpt["best_valid_loss"]125 best_loss_epoch = ckpt["best_loss_epoch"]126 train_losses = np.load(train_loss_file)127 valid_losses = np.load(valid_loss_file)128 129 ## Distributed model130 if dist.is_initialized():131 model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank)132 world_rank == 0 and logger.info(f"start training ...")133 134 for epoch in range(cfg.max_epoch):135 if dist.is_initialized():136 train_sampler.set_epoch(epoch)137 val_sampler.set_epoch(epoch)138 model.train()139 train_loss = 0140 start_time = time.time()141 for j, data in enumerate(train_dataloader):142 invar = data[0].to(device, dtype=torch.float32) # [B, input_steps, C, H, W]143 outvar = data[1].to(device, dtype=torch.float32) # [B, output_steps, C, H, W]144 outvar_pred = model(invar) # [B, output_steps, C, H, W]145 loss = model.wmse_loss(outvar_pred, outvar) # 论文式(1) WMSE146 optimizer.zero_grad()147 loss.backward()148 optimizer.step()149 train_loss += loss.item()150 if world_rank == 0:151 logger.info(f'Train: Epoch {epoch}-{j+1}/{len(train_dataloader)} '152 f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] '153 f'[{(time.time()-start_time)/(j+1): .02f}s/{cfg_data.dataloader.batch_size}batch] '154 f'loss:{train_loss / (j+1): .04f}')155 156 train_loss /= len(train_dataloader)157 158 model.eval()159 valid_loss = 0160 with torch.no_grad():161 start_time = time.time()162 for j, data in enumerate(val_dataloader):163 invar = data[0].to(device, dtype=torch.float32)164 outvar = data[1].to(device, dtype=torch.float32)165 outvar_pred = model(invar)166 loss = model.wmse_loss(outvar_pred, outvar)167 168 if dist.is_initialized():169 loss_tensor = loss.detach().to(device)170 dist.all_reduce(loss_tensor)171 loss = loss_tensor.item() / cfg.world_size172 valid_loss += loss173 else:174 valid_loss += loss.item()175 if world_rank == 0:176 logger.info(f'Valid: Epoch {epoch}-{j+1}/{len(val_dataloader)} '177 f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] '178 f'loss:{valid_loss / (j+1): .04f}')179 180 valid_loss /= len(val_dataloader)181 is_save_ckp = False182 if valid_loss < best_valid_loss:183 best_valid_loss = valid_loss184 best_loss_epoch = epoch185 world_rank == 0 and save_checkpoint(model, optimizer, scheduler, best_valid_loss, best_loss_epoch, cfg.checkpoint_dir)186 is_save_ckp = True187 scheduler.step(valid_loss)188 189 if world_rank == 0:190 logger.info(f"Epoch [{epoch + 1}/{cfg.max_epoch}], "191 f"Train Loss: {train_loss:.4f}, "192 f"Valid Loss: {valid_loss:.4f}, "193 f"Best loss at Epoch: {best_loss_epoch + 1}"194 + (", saving checkpoint" if is_save_ckp else "")195 )196 train_losses = np.append(train_losses, train_loss)197 valid_losses = np.append(valid_losses, valid_loss)198 np.save(train_loss_file, train_losses)199 np.save(valid_loss_file, valid_losses)200 201 if epoch - best_loss_epoch > cfg.patience:202 print(f"Loss has not decrease in {cfg.patience} epochs, stopping training...")203 exit()204 205 206def save_checkpoint(model, optimizer, scheduler, best_valid_loss, best_loss_epoch, model_path):207 model_to_save = model.module if hasattr(model, "module") else model208 state = {"model_state_dict": model_to_save.state_dict(),209 "optimizer_state_dict": optimizer.state_dict(),210 "scheduler_state_dict": scheduler.state_dict(),211 "best_valid_loss": best_valid_loss,212 "best_loss_epoch": best_loss_epoch,213 }214 torch.save(state, f"{model_path}/model.pth")215 ### the weight file saving may interrupted due to DCU queue limit, get a backup to ensure there at least has one model216 os.system(f"mv {model_path}/model.pth {model_path}/model_bak.pth")217 218 219if __name__ == "__main__":220 current_path = os.getcwd()221 sys.path.append(current_path)222 main()223 