| import sys |
| from pathlib import Path |
|
|
| |
| root_path = Path(__file__).parent.parent |
| sys.path.append(str(root_path)) |
| import torch |
| import os |
| import numpy as np |
| import torch.distributed as dist |
| import logging |
| import time |
|
|
| from model.dgmr import DGMR |
| from model.dgmr_official.losses import loss_hinge_disc, loss_hinge_gen |
| from onescience.datapipes.climate import ERA5Datapipe |
| from onescience.utils.YParams import YParams |
|
|
| try: |
| from apex import optimizers |
| _FUSED_ADAM = True |
| except Exception: |
| _FUSED_ADAM = False |
|
|
|
|
| def main(): |
| logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s") |
| logger = logging.getLogger() |
|
|
| |
| config_file_path = os.path.join(current_path, "conf/config.yaml") |
| cfg = YParams(config_file_path, "model") |
|
|
| |
| cfg.world_size = 1 |
| if "WORLD_SIZE" in os.environ: |
| cfg.world_size = int(os.environ["WORLD_SIZE"]) |
| world_rank = 0 |
| local_rank = 0 |
| if cfg.world_size > 1 and torch.cuda.is_available(): |
| dist.init_process_group(backend="nccl", init_method="env://") |
| local_rank = int(os.environ["LOCAL_RANK"]) |
| world_rank = dist.get_rank() |
| device = f"cuda:{local_rank}" if torch.cuda.is_available() else "cpu" |
|
|
| |
| cfg_data = YParams(config_file_path, "datapipe") |
| cfg['N_in_channels'] = len(cfg_data.dataset.channels) |
| cfg['N_out_channels'] = len(cfg_data.dataset.channels) |
| datapipe = ERA5Datapipe( |
| dataset_dir=cfg_data.dataset.data_dir, |
| used_variables=cfg_data.dataset.channels, |
| used_years=cfg_data.dataset.train_time, |
| distributed=dist.is_initialized(), |
| input_steps=cfg.num_context, |
| output_steps=cfg.forecast_steps, |
| batch_size=cfg_data.dataloader.batch_size, |
| num_workers=cfg_data.dataloader.num_workers, |
| ) |
| train_dataloader, train_sampler = datapipe.get_dataloader("train") |
| datapipe = ERA5Datapipe( |
| dataset_dir=cfg_data.dataset.data_dir, |
| used_variables=cfg_data.dataset.channels, |
| used_years=cfg_data.dataset.val_time, |
| distributed=dist.is_initialized(), |
| input_steps=cfg.num_context, |
| output_steps=cfg.forecast_steps, |
| batch_size=cfg_data.dataloader.batch_size, |
| num_workers=cfg_data.dataloader.num_workers, |
| ) |
| val_dataloader, val_sampler = datapipe.get_dataloader("valid") |
|
|
| |
| model = DGMR( |
| forecast_steps=cfg.forecast_steps, |
| num_context=cfg.num_context, |
| input_channels=cfg.input_channels, |
| output_shape=cfg.output_shape, |
| conv_type=cfg.conv_type, |
| latent_channels=cfg.latent_channels, |
| context_channels=cfg.context_channels, |
| generation_steps=cfg.generation_steps, |
| grid_lambda=cfg.grid_lambda, |
| precip_weight_cap=cfg.precip_weight_cap, |
| ).to(device) |
|
|
| |
| if _FUSED_ADAM: |
| optimizer_g = optimizers.FusedAdam(model.generator.parameters(), lr=cfg.lr) |
| optimizer_d = optimizers.FusedAdam(model.discriminator.parameters(), lr=cfg.lr_disc) |
| else: |
| optimizer_g = torch.optim.Adam(model.generator.parameters(), lr=cfg.lr) |
| optimizer_d = torch.optim.Adam(model.discriminator.parameters(), lr=cfg.lr_disc) |
| scheduler_g = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer_g, factor=0.2, patience=5, mode='min') |
| scheduler_d = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer_d, factor=0.2, patience=5, mode='min') |
|
|
| |
| os.makedirs(cfg.checkpoint_dir, exist_ok=True) |
| train_loss_file = f"{cfg.checkpoint_dir}/trloss.npy" |
| valid_loss_file = f"{cfg.checkpoint_dir}/valoss.npy" |
| best_valid_loss = float("inf") |
| best_loss_epoch = 0 |
| train_losses = np.empty((0,), dtype=np.float32) |
| valid_losses = np.empty((0,), dtype=np.float32) |
|
|
| |
| if cfg.world_size == 1: |
| total_params = sum(p.numel() for p in model.parameters()) |
| print("\n\n") |
| print("-" * 50) |
| print(f"📂 now params is {total_params}, {total_params / 1e6:.2f}M, {total_params / 1e9:.2f}B") |
| print("-" * 50, "\n") |
|
|
| |
| if os.path.exists(f"{cfg.checkpoint_dir}/model_bak.pth"): |
| if world_rank == 0: |
| print("\n\n") |
| print("-" * 50) |
| print(f"✅ There has a model weight, load and continue training...") |
| print(f'If you want to train a new model, ensure there is no *.pth file in {cfg.checkpoint_dir}') |
| print("-" * 50, "\n") |
| ckpt = torch.load(f"{cfg.checkpoint_dir}/model_bak.pth", map_location=device, weights_only=False) |
| model.load_state_dict(ckpt["model_state_dict"]) |
| optimizer_g.load_state_dict(ckpt["optimizer_g_state_dict"]) |
| optimizer_d.load_state_dict(ckpt["optimizer_d_state_dict"]) |
| scheduler_g.load_state_dict(ckpt["scheduler_g_state_dict"]) |
| scheduler_d.load_state_dict(ckpt["scheduler_d_state_dict"]) |
| best_valid_loss = ckpt["best_valid_loss"] |
| best_loss_epoch = ckpt["best_loss_epoch"] |
| train_losses = np.load(train_loss_file) |
| valid_losses = np.load(valid_loss_file) |
|
|
| |
| if dist.is_initialized(): |
| model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[local_rank], output_device=local_rank) |
| world_rank == 0 and logger.info(f"start training ...") |
|
|
| for epoch in range(cfg.max_epoch): |
| if dist.is_initialized(): |
| train_sampler.set_epoch(epoch) |
| val_sampler.set_epoch(epoch) |
| model.train() |
| train_loss = 0 |
| start_time = time.time() |
| for j, data in enumerate(train_dataloader): |
| invar = data[0].to(device, dtype=torch.float32) |
| outvar = data[1].to(device, dtype=torch.float32) |
| full_real = torch.cat([invar, outvar], dim=1) |
|
|
| |
| gen_images = model.generator(invar) |
| full_fake = torch.cat([invar, gen_images], dim=1) |
| score_real = model.discriminator(full_real) |
| score_generated = model.discriminator(full_fake.detach()) |
| disc_loss = loss_hinge_disc(score_generated, score_real) |
| optimizer_d.zero_grad() |
| disc_loss.backward() |
| optimizer_d.step() |
|
|
| |
| score_generated = model.discriminator(full_fake) |
| gen_loss = loss_hinge_gen(score_generated) |
| gen_samples = torch.stack( |
| [model.generator(invar) for _ in range(cfg.generation_steps)], dim=0 |
| ).mean(dim=0) |
| grid_loss = model.grid_regularizer(gen_samples, outvar) |
| gen_loss = gen_loss + cfg.grid_lambda * grid_loss |
| optimizer_g.zero_grad() |
| gen_loss.backward() |
| optimizer_g.step() |
|
|
| loss = gen_loss.item() + disc_loss.item() |
| train_loss += loss |
| if world_rank == 0: |
| logger.info(f'Train: Epoch {epoch}-{j+1}/{len(train_dataloader)} ' |
| f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] ' |
| f'[{(time.time()-start_time)/(j+1): .02f}s/{cfg_data.dataloader.batch_size}batch] ' |
| f'gen_loss:{gen_loss.item(): .04f} disc_loss:{disc_loss.item(): .04f} ' |
| f'loss:{train_loss / (j+1): .04f}') |
|
|
| train_loss /= len(train_dataloader) |
|
|
| |
| model.eval() |
| valid_loss = 0 |
| with torch.no_grad(): |
| start_time = time.time() |
| for j, data in enumerate(val_dataloader): |
| invar = data[0].to(device, dtype=torch.float32) |
| outvar = data[1].to(device, dtype=torch.float32) |
| gen_images = model.generator(invar) |
| grid_loss = model.grid_regularizer(gen_images, outvar) |
| loss = grid_loss |
|
|
| if dist.is_initialized(): |
| loss_tensor = loss.detach().to(device) |
| dist.all_reduce(loss_tensor) |
| loss = loss_tensor.item() / cfg.world_size |
| valid_loss += loss |
| else: |
| valid_loss += loss.item() |
| if world_rank == 0: |
| logger.info(f'Valid: Epoch {epoch}-{j+1}/{len(val_dataloader)} ' |
| f'[cost {int((time.time()-start_time) // 60):02}:{int((time.time()-start_time) % 60):02}] ' |
| f'loss:{valid_loss / (j+1): .04f}') |
|
|
| valid_loss /= len(val_dataloader) |
| is_save_ckp = False |
| if valid_loss < best_valid_loss: |
| best_valid_loss = valid_loss |
| best_loss_epoch = epoch |
| world_rank == 0 and save_checkpoint(model, optimizer_g, optimizer_d, scheduler_g, scheduler_d, best_valid_loss, best_loss_epoch, cfg.checkpoint_dir) |
| is_save_ckp = True |
| scheduler_g.step(valid_loss) |
| scheduler_d.step(valid_loss) |
|
|
| if world_rank == 0: |
| logger.info(f"Epoch [{epoch + 1}/{cfg.max_epoch}], " |
| f"Train Loss: {train_loss:.4f}, " |
| f"Valid Loss: {valid_loss:.4f}, " |
| f"Best loss at Epoch: {best_loss_epoch + 1}" |
| + (", saving checkpoint" if is_save_ckp else "") |
| ) |
| train_losses = np.append(train_losses, train_loss) |
| valid_losses = np.append(valid_losses, valid_loss) |
| np.save(train_loss_file, train_losses) |
| np.save(valid_loss_file, valid_losses) |
|
|
| if epoch - best_loss_epoch > cfg.patience: |
| print(f"Loss has not decrease in {cfg.patience} epochs, stopping training...") |
| exit() |
|
|
|
|
| def save_checkpoint(model, optimizer_g, optimizer_d, scheduler_g, scheduler_d, best_valid_loss, best_loss_epoch, model_path): |
| model_to_save = model.module if hasattr(model, "module") else model |
| state = {"model_state_dict": model_to_save.state_dict(), |
| "optimizer_g_state_dict": optimizer_g.state_dict(), |
| "optimizer_d_state_dict": optimizer_d.state_dict(), |
| "scheduler_g_state_dict": scheduler_g.state_dict(), |
| "scheduler_d_state_dict": scheduler_d.state_dict(), |
| "best_valid_loss": best_valid_loss, |
| "best_loss_epoch": best_loss_epoch, |
| } |
| torch.save(state, f"{model_path}/model.pth") |
| |
| os.system(f"mv {model_path}/model.pth {model_path}/model_bak.pth") |
|
|
|
|
| if __name__ == "__main__": |
| current_path = os.getcwd() |
| sys.path.append(current_path) |
| main() |
|
|