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())