File size: 4,415 Bytes
1aeffbb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import argparse
import os
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

import torch
import torch.distributed as dist
import yaml
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler

from data_loader import SyntheticOceanDataset
from model.glonet import GLONET


def setup_distributed():
    world_size = int(os.environ.get("WORLD_SIZE", "1"))
    if world_size == 1:
        return 0, 0, torch.device("cuda" if torch.cuda.is_available() else "cpu")
    local_rank = int(os.environ["LOCAL_RANK"])
    force_cpu = os.environ.get("GLONET_FORCE_CPU", "0") == "1"
    if torch.cuda.is_available() and not force_cpu:
        device_count = torch.cuda.device_count()
        if local_rank >= device_count:
            raise RuntimeError(
                f"LOCAL_RANK={local_rank} but only {device_count} accelerator(s) are visible; "
                "reduce --nproc_per_node or fix CUDA_VISIBLE_DEVICES."
            )
        torch.cuda.set_device(local_rank)
        device = torch.device("cuda", local_rank)
        backend = "nccl"
    else:
        device = torch.device("cpu")
        backend = "gloo"
    dist.init_process_group(backend=backend, init_method="env://")
    return dist.get_rank(), local_rank, device


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default=str(ROOT / "conf/config.yaml"))
    parser.add_argument("--phase", choices=["pretrain", "finetune"], default="pretrain")
    args = parser.parse_args()
    with open(args.config, encoding="utf-8") as handle:
        config = yaml.safe_load(handle)
    rank, local_rank, device = setup_distributed()
    torch.manual_seed(config["project"]["seed"] + rank)
    channels = len(config["data"]["channels"])
    rollout_steps = config["training"][f"{args.phase}_rollout_steps"]
    dataset = SyntheticOceanDataset(config["data"]["synthetic_samples"], channels, config["data"]["grid"],
                                    input_steps=config["data"]["input_steps"],
                                    output_steps=config["data"]["output_steps"],
                                    data_dir=str(ROOT / config["data"]["data_dir"]))
    sampler = DistributedSampler(dataset, shuffle=True) if dist.is_initialized() else None
    loader = DataLoader(dataset, batch_size=config["data"]["batch_size"], shuffle=sampler is None, sampler=sampler)
    model = GLONET(channels * config["data"]["input_steps"], out_channels=channels,
                   hidden_channels=config["model"]["hidden_channels"], modes=config["model"]["modes"],
                   layers=config["model"]["layers"]).to(device)
    checkpoint = ROOT / config["training"]["checkpoint"]
    if args.phase == "finetune" and checkpoint.exists():
        state = torch.load(checkpoint, map_location=device, weights_only=False)
        model.load_state_dict(state["model"])
    if dist.is_initialized():
        model = DDP(model, device_ids=[local_rank] if device.type == "cuda" else None)
    optimizer = torch.optim.Adam(model.parameters(), lr=config["training"]["learning_rate"])
    for epoch in range(config["training"]["epochs"]):
        if sampler is not None:
            sampler.set_epoch(epoch)
        model.train()
        total = 0.0
        for inputs, targets in loader:
            inputs, targets = inputs.to(device), targets.to(device)
            optimizer.zero_grad(set_to_none=True)
            loss = 0.0
            state = inputs
            for step in range(rollout_steps):
                prediction = model(state)
                loss = loss + torch.nn.functional.mse_loss(prediction, targets[:, step])
                state = torch.cat((state[:, 1:], prediction.unsqueeze(1)), dim=1)
            loss = loss / rollout_steps
            loss.backward()
            optimizer.step()
            total += loss.item()
        if rank == 0:
            print(f"epoch={epoch + 1} loss={total / len(loader):.6f}")
    if rank == 0:
        checkpoint.parent.mkdir(parents=True, exist_ok=True)
        torch.save({"model": model.module.state_dict() if hasattr(model, "module") else model.state_dict(),
                    "config": config}, checkpoint)
        print(f"saved={checkpoint}")
    if dist.is_initialized():
        dist.destroy_process_group()


if __name__ == "__main__":
    main()