WeatherNext2 / scripts /train.py
Zhongning's picture
Upload folder using huggingface_hub
9c16f7b verified
Raw
History Blame Contribute Delete
9.36 kB
import sys
from pathlib import Path
# 获取项目根目录(train.py上级的上级)
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.fgn import FGN
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()
## Model config init
config_file_path = os.path.join(current_path, "conf/config.yaml")
cfg = YParams(config_file_path, "model")
## Distributed config init
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"
## DataLoader init
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.input_steps,
output_steps=cfg.output_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.input_steps,
output_steps=cfg.output_steps,
batch_size=cfg_data.dataloader.batch_size,
num_workers=cfg_data.dataloader.num_workers,
)
val_dataloader, val_sampler = datapipe.get_dataloader("valid")
# Model init
model = FGN(
in_channels=cfg['N_in_channels'],
out_channels=cfg['N_out_channels'],
input_steps=cfg.input_steps,
output_steps=cfg.output_steps,
grid_shape=cfg.grid_shape,
mesh_shape=cfg.mesh_shape,
latent_dim=cfg.latent_dim,
num_encoder_layers=cfg.num_encoder_layers,
num_decoder_layers=cfg.num_decoder_layers,
num_processor_blocks=cfg.num_processor_blocks,
n_heads=cfg.n_heads,
hidden_dim=cfg.hidden_dim,
noise_dim=cfg.noise_dim,
channel_weights=cfg.channel_weights,
).to(device)
if _FUSED_ADAM:
optimizer = optimizers.FusedAdam(model.parameters(), lr=cfg.lr)
else:
optimizer = torch.optim.Adam(model.parameters(), lr=cfg.lr)
scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.2, patience=5, mode='min')
## Train process init
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)
## Get model params count
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")
## Load model weight if there exist well-trained model
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.load_state_dict(ckpt["optimizer_state_dict"])
scheduler.load_state_dict(ckpt["scheduler_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)
## Distributed model
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) # [B, input_steps, C, H, W]
outvar = data[1].to(device, dtype=torch.float32) # [B, output_steps, C, H, W]
outvar_pred = model(invar, num_members=cfg.num_members) # [B, M, output_steps, C, H, W]
loss = model.crps_loss(outvar_pred, outvar) # 论文式(4) 公平 CRPS
optimizer.zero_grad()
loss.backward()
optimizer.step()
train_loss += loss.item()
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'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)
outvar_pred = model(invar, num_members=cfg.num_members)
loss = model.crps_loss(outvar_pred, outvar)
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, scheduler, best_valid_loss, best_loss_epoch, cfg.checkpoint_dir)
is_save_ckp = True
scheduler.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, scheduler, 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_state_dict": optimizer.state_dict(),
"scheduler_state_dict": scheduler.state_dict(),
"best_valid_loss": best_valid_loss,
"best_loss_epoch": best_loss_epoch,
}
torch.save(state, f"{model_path}/model.pth")
### the weight file saving may interrupted due to DCU queue limit, get a backup to ensure there at least has one model
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()