DGMR / scripts /train.py
Zhongning's picture
Upload folder using huggingface_hub
5a5d1a8 verified
Raw
History Blame Contribute Delete
11.3 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.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()
## 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.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 init
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)
# 生成器与判别器使用独立的 Adam 优化器(论文 lr=1e-4)
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')
## 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_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)
## 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, num_context, C, H, W]
outvar = data[1].to(device, dtype=torch.float32) # [B, forecast_steps, C, H, W]
full_real = torch.cat([invar, outvar], dim=1) # 上下文 + 真实未来帧
# --- 判别器(hinge loss,生成图像 detach 不传梯度) ---
gen_images = model.generator(invar) # [B, forecast_steps, C, H, W]
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()
# --- 生成器(hinge + 网格单元正则器,MC 采样估计期望) ---
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) # 取 MC 均值作为生成均值图像
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")
### 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()