CoolFace
Modelpublic

OneScience-Group/Pangu_Weather

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes18downloads
train.py249 linesDownload Raw Back to scripts
1import sys2from pathlib import Path3 4# 获取项目根目录(train.py上级的上级)5root_path = Path(__file__).parent.parent6sys.path.append(str(root_path))7 8import torch9import os10import numpy as np11import torch.distributed as dist12import logging13import time14import torch.nn.functional as F15from torch.nn.parallel import DistributedDataParallel16from model.pangu import Pangu17from onescience.datapipes.climate import ERA5Datapipe18from onescience.utils.YParams import YParams19from onescience.memory.checkpoint import replace_function20from apex import optimizers21 22 23 24 25def loss_func(x, y, weights, level_weight=1.0):26    return level_weight * (F.l1_loss(x, y, reduction='none') * weights).mean()27 28def main():29 30    logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")31    logger = logging.getLogger()32 33    ## Model config init34    config_file_path = os.path.join(current_path, "conf/config.yaml")35    cfg = YParams(config_file_path, "model")36 37    ## Distributed config init38    cfg.world_size = 139    if "WORLD_SIZE" in os.environ:40        cfg.world_size = int(os.environ["WORLD_SIZE"])41    world_rank = 042    local_rank = 043    if cfg.world_size > 1:44        dist.init_process_group(backend="nccl", init_method="env://")45        local_rank = int(os.environ["LOCAL_RANK"])46        world_rank = dist.get_rank()47 48    ## DataLoader init49    cfg_data = YParams(config_file_path, "datapipe")50    datapipe = ERA5Datapipe(51        dataset_dir=cfg_data.dataset.data_dir,52        used_variables=cfg_data.dataset.channels,53        used_years=cfg_data.dataset.train_time,54        distributed=dist.is_initialized(),55        batch_size=cfg_data.dataloader.batch_size,56        num_workers=cfg_data.dataloader.num_workers57    )58    train_dataloader, train_sampler = datapipe.get_dataloader("train")59    datapipe = ERA5Datapipe(60        dataset_dir=cfg_data.dataset.data_dir,61        used_variables=cfg_data.dataset.channels,62        used_years=cfg_data.dataset.val_time,63        distributed=dist.is_initialized(),64        batch_size=cfg_data.dataloader.batch_size,65        num_workers=cfg_data.dataloader.num_workers66    )67    val_dataloader, val_sampler = datapipe.get_dataloader("valid")68 69    surface_weights = torch.as_tensor(cfg_data.dataset.weights[:4], device=local_rank, dtype=torch.float32).view(1, -1, 1, 1)70    pressure_weights = torch.as_tensor(cfg_data.dataset.weights[4:], device=local_rank, dtype=torch.float32).view(1, -1, 1, 1)71 72    static_dir = os.path.join(cfg_data.dataset.data_dir, "static")73    74    land_mask = torch.from_numpy(np.load(os.path.join(static_dir, "land_mask.npy")).astype(np.float32))75    soil_type = torch.from_numpy(np.load(os.path.join(static_dir, "soil_type.npy")).astype(np.float32))76    topography = torch.from_numpy(np.load(os.path.join(static_dir, "topography.npy")).astype(np.float32))77    topography = (topography - topography.mean()) / (topography.std(unbiased=False) + 1e-6)78    surface_mask = torch.stack([land_mask, soil_type, topography], dim=0).to(local_rank)79    surface_mask = surface_mask.unsqueeze(0).repeat(cfg_data.dataloader.batch_size, 1, 1, 1)80 81    ## Model init82    model = Pangu(img_size=cfg_data.dataset.img_size,83                  patch_size=cfg.patch_size,84                  embed_dim=cfg.embed_dim,85                  num_heads=cfg.num_heads,86                  window_size=cfg.window_size,87                  ).to(local_rank)88    optimizer = optimizers.FusedAdam(model.parameters(), betas=(0.9, 0.999), lr=5e-4, weight_decay=3e-6)89    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=100)90    91    ## Train process init92    os.makedirs(cfg.checkpoint_dir, exist_ok=True)93    train_loss_file = f"{cfg.checkpoint_dir}/trloss.npy"94    valid_loss_file = f"{cfg.checkpoint_dir}/valoss.npy"95    best_valid_loss = 1.0e696    best_loss_epoch = 097    train_losses = np.empty((0,), dtype=np.float32)98    valid_losses = np.empty((0,), dtype=np.float32)99    current_epoch = 0100 101    ## Get model params count102    if cfg.world_size == 1:103        total_params = sum(p.numel() for p in model.parameters())104        print("\n\n")105        print("-" * 50)106        print(f"📂 now params is {total_params}, {total_params / 1e6:.2f}M, {total_params / 1e9:.2f}B")107        print("-" * 50, "\n")108 109    ## Load model weight if there exist well-trained model 110    if os.path.exists(f"{cfg.checkpoint_dir}/model_bak.pth"):111        if world_rank == 0:112            print("\n\n")113            print("-" * 50)114            print(f"✅ There has a model weight, load and continue training...")115            print(f'If you want to train a new model, ensure there is no *.pth file in {cfg.checkpoint_dir}')116            print("-" * 50, "\n")117        ckpt = torch.load(f"{cfg.checkpoint_dir}/model_bak.pth", map_location=f'cuda:{local_rank}', weights_only=False)118        model.load_state_dict(ckpt["model_state_dict"])119        optimizer.load_state_dict(ckpt["optimizer_state_dict"])120        scheduler.load_state_dict(ckpt["scheduler_state_dict"])121        best_valid_loss = ckpt["best_valid_loss"]122        best_loss_epoch = ckpt["best_loss_epoch"]123        current_epoch = ckpt["current_epoch"]124        train_losses = np.load(train_loss_file)125        valid_losses = np.load(valid_loss_file)126 127    ## Distributed model128    if cfg.world_size > 1:129        model = DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank)130 131    world_rank == 0 and logger.info(f"start training ...")132 133    for epoch in range(current_epoch, cfg.max_epoch):134        if dist.is_initialized():135            train_sampler.set_epoch(epoch)136            val_sampler.set_epoch(epoch)137 138        model.train()139        train_loss = 0140        start_time = time.time()141        for j, data in enumerate(train_dataloader):142            invar = data[0]143            outvar = data[1]144            invar_surface = invar[:, :4, :, :].to(local_rank, dtype=torch.float32)145            invar_upper_air = invar[:, 4:, :, :].to(local_rank, dtype=torch.float32)146            invar = torch.concat([invar_surface, surface_mask, invar_upper_air], dim=1)147            tar_surface = outvar[:, :4, :, :].to(local_rank, dtype=torch.float32)148            tar_upper_air = outvar[:, 4:, :, :].to(local_rank, dtype=torch.float32)149 150            with replace_function(model,["layer2", "layer3"],cfg.world_size > 1):151                out_surface, out_upper_air = model(invar)152 153            out_upper_air = out_upper_air.reshape(tar_upper_air.shape)154            loss1 = loss_func(out_surface, tar_surface, surface_weights,  level_weight=0.25)155            loss2 = loss_func(out_upper_air, tar_upper_air, pressure_weights, level_weight=1.0)156            # 总 loss157            loss = loss1 + loss2158            optimizer.zero_grad()159            loss.backward()160            optimizer.step()161            train_loss += loss.item()162            if world_rank == 0:163                logger.info(f'Train: Epoch {epoch}-{j+1}/{len(train_dataloader)} '164                            f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] '165                            f'[{(time.time()-start_time)/(j+1): .02f}s/{cfg_data.dataloader.batch_size}batch] '166                            f'loss:{train_loss / (j+1): .04f}')167            168        train_loss /= len(train_dataloader)169 170        model.eval()171        valid_loss = 0172        with torch.no_grad():173            start_time = time.time()174            for j, data in enumerate(val_dataloader):175                invar = data[0]176                outvar = data[1]177                invar_surface = invar[:, :4, :, :].to(local_rank, dtype=torch.float32)178                invar_upper_air = invar[:, 4:, :, :].to(local_rank, dtype=torch.float32)179                invar = torch.concat([invar_surface, surface_mask, invar_upper_air], dim=1)180                tar_surface = outvar[:, :4, :, :].to(local_rank, dtype=torch.float32)181                tar_upper_air = outvar[:, 4:, :, :].to(local_rank, dtype=torch.float32)182 183                with replace_function(model,["layer2", "layer3"],cfg.world_size > 1):184                    out_surface, out_upper_air = model(invar)185 186                out_upper_air = out_upper_air.reshape(tar_upper_air.shape)187                loss1 = loss_func(out_surface, tar_surface, surface_weights,  level_weight=0.25).item()188                loss2 = loss_func(out_upper_air, tar_upper_air, pressure_weights, level_weight=1.0).item()189                # 总 loss190                loss = loss1 + loss2191 192                if cfg.world_size > 1:193                    loss_tensor = torch.tensor(loss, device=local_rank)194                    dist.all_reduce(loss_tensor)195                    loss = loss_tensor.item() / cfg.world_size196                valid_loss += loss197                if world_rank == 0:198                    logger.info(f'Valid: Epoch {epoch}-{j+1}/{len(val_dataloader)} '199                            f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] '200                            f'[{(time.time()-start_time)/(j+1): .02f}s/{cfg_data.dataloader.batch_size}batch] '201                            f'loss:{valid_loss / (j+1): .04f}')202                203        valid_loss /= len(val_dataloader)204        is_save_ckp = False205        if valid_loss < best_valid_loss:206            best_valid_loss = valid_loss207            best_loss_epoch = epoch208            world_rank == 0 and save_checkpoint(model, optimizer, scheduler, best_valid_loss, best_loss_epoch, cfg.checkpoint_dir, epoch)209            is_save_ckp = True210 211        scheduler.step()212 213        if world_rank == 0:214            logger.info(f"Epoch [{epoch}/{cfg.max_epoch}], "215                        f"Train Loss: {train_loss:.4f}, "216                        f"Valid Loss: {valid_loss:.4f}, "217                        f"Best loss at Epoch: {best_loss_epoch}"218                        + (", saving checkpoint" if is_save_ckp else "")219                        )220            train_losses = np.append(train_losses, train_loss)221            valid_losses = np.append(valid_losses, valid_loss)222 223            np.save(train_loss_file, train_losses)224            np.save(valid_loss_file, valid_losses)225 226        if epoch - best_loss_epoch > cfg.patience:227            print(f"Loss has not decrease in {cfg.patience} epochs, stopping training...")228            exit()229 230 231def save_checkpoint(model, optimizer, scheduler, best_valid_loss, best_loss_epoch, model_path, epoch):232    model_to_save = model.module if hasattr(model, "module") else model233    state = {"model_state_dict": model_to_save.state_dict(),234             "optimizer_state_dict": optimizer.state_dict(),235             "scheduler_state_dict": scheduler.state_dict(),236             "best_valid_loss": best_valid_loss,237             "best_loss_epoch": best_loss_epoch,238             "current_epoch": epoch239            }240    torch.save(state, f"{model_path}/model.pth")241    ### the weight file saving may interrupted due to DCU queue limit, get a backup to ensure there at least has one model 242    os.system(f"mv {model_path}/model.pth {model_path}/model_bak.pth")243 244 245if __name__ == "__main__":246    current_path = os.getcwd()247    sys.path.append(current_path)248    main()249