File size: 5,017 Bytes
a3441f1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 | from __future__ import annotations
import argparse
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP
from common import (
DEFAULT_CONFIG,
build_datapipe,
build_model,
cleanup_distributed,
get_attr,
initialize_distributed,
load_config,
load_model_state,
predict_batch,
prepare_config,
)
def train_one_epoch(model, loader, optimizer, loss_fn, device, cfg) -> float:
model.train()
losses = []
for x, y, grid in loader:
x = x.to(device)
y = y.to(device)
grid = grid.to(device)
pred, target = predict_batch(model, x, y, grid, cfg)
if pred.numel() == 0:
continue
loss = loss_fn(pred.reshape(pred.size(0), -1), target.reshape(target.size(0), -1))
optimizer.zero_grad()
loss.backward()
optimizer.step()
losses.append(float(loss.item()))
return float(np.mean(losses)) if losses else 0.0
@torch.no_grad()
def validate(model, loader, loss_fn, device, cfg) -> float:
model.eval()
losses = []
for x, y, grid in loader:
x = x.to(device)
y = y.to(device)
grid = grid.to(device)
pred, target = predict_batch(model, x, y, grid, cfg)
if pred.numel() == 0:
continue
loss = loss_fn(pred.reshape(pred.size(0), -1), target.reshape(target.size(0), -1))
losses.append(float(loss.item()))
return float(np.mean(losses)) if losses else 0.0
def main() -> int:
parser = argparse.ArgumentParser(description="Train the standard PDENNEval FNO package.")
parser.add_argument("--config", default=str(DEFAULT_CONFIG), help="Path to conf/config.yaml")
parser.add_argument("--data-dir", default=None, help="Override datapipe.source.data_dir")
parser.add_argument("--output-dir", default=None, help="Override training.output_dir")
parser.add_argument("--force-local-datapipe", action="store_true")
args = parser.parse_args()
dist = initialize_distributed()
device = dist.device
cfg = prepare_config(load_config(args.config), data_dir=args.data_dir, output_dir=args.output_dir)
seed = int(get_attr(cfg.training, "seed", 0))
torch.manual_seed(seed + dist.rank)
np.random.seed(seed + dist.rank)
output_dir = Path(cfg.training.output_dir)
if dist.rank == 0:
output_dir.mkdir(parents=True, exist_ok=True)
print(f"Output directory: {output_dir}")
print(f"Device: {device}, world_size={dist.world_size}")
datapipe = build_datapipe(
cfg,
distributed=(dist.world_size > 1),
force_local=args.force_local_datapipe,
)
train_loader, train_sampler = datapipe.train_dataloader()
val_loader, _ = datapipe.val_dataloader()
model = build_model(datapipe.spatial_dim, cfg).to(device)
if dist.world_size > 1:
model = DDP(model, device_ids=[dist.local_rank], output_device=dist.local_rank)
optimizer_cfg = cfg.training.optimizer
scheduler_cfg = cfg.training.scheduler
optimizer = getattr(torch.optim, optimizer_cfg.name)(
model.parameters(),
lr=float(optimizer_cfg.lr),
weight_decay=float(optimizer_cfg.weight_decay),
)
scheduler = getattr(torch.optim.lr_scheduler, scheduler_cfg.name)(
optimizer,
step_size=int(scheduler_cfg.step_size),
gamma=float(scheduler_cfg.gamma),
)
loss_fn = nn.MSELoss()
if bool(get_attr(cfg.training, "continue_training", False)) and cfg.training.model_path:
model_to_load = model.module if hasattr(model, "module") else model
model_to_load.load_state_dict(load_model_state(Path(cfg.training.model_path), device))
best_val = float("inf")
epochs = int(cfg.training.epochs)
save_period = max(1, int(cfg.training.save_period))
for epoch in range(epochs):
if train_sampler is not None:
train_sampler.set_epoch(epoch)
train_loss = train_one_epoch(model, train_loader, optimizer, loss_fn, device, cfg)
scheduler.step()
val_loss = validate(model, val_loader, loss_fn, device, cfg)
if dist.rank == 0:
print(f"epoch={epoch} train_loss={train_loss:.6e} val_loss={val_loss:.6e}")
model_to_save = model.module if hasattr(model, "module") else model
state = {
"epoch": epoch,
"model_state_dict": model_to_save.state_dict(),
"optimizer_state_dict": optimizer.state_dict(),
"loss": val_loss,
}
torch.save(state, output_dir / "latest_model.pt")
if (epoch + 1) % save_period == 0:
torch.save(state, output_dir / f"model_epoch_{epoch}.pt")
if val_loss <= best_val:
best_val = val_loss
torch.save(state, output_dir / "best_model.pt")
cleanup_distributed()
return 0
if __name__ == "__main__":
raise SystemExit(main())
|