CoolFace
Modelpublic

OneScience-Group/GraphDOP

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes16downloads
train.py223 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 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