CoolFace
Modelpublic

OneScience-Group/FourCastNet

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes40downloads
train.py200 linesDownload Raw Back to scripts
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