File size: 3,030 Bytes
efe4fbe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
import argparse
import os
from pathlib import Path

import torch
import torch.nn.functional as F
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, DistributedSampler
from model.metnet_2 import (WindowDataset, build_model, categorical_nll_chunked,
                            load_config, save_checkpoint, write_json)

parser = argparse.ArgumentParser(description="Train MetNet-2 on selected windows")
parser.add_argument("--config", default="conf/config.yaml")
parser.add_argument("--steps", type=int, default=None)
args = parser.parse_args()
config = load_config(args.config)
rank = int(os.environ.get("RANK", "0"))
world_size = int(os.environ.get("WORLD_SIZE", "1"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
distributed = world_size > 1
requested = config["runtime"]["device"]
use_cuda = (requested != "cpu" and torch.cuda.is_available()
            and (not distributed or torch.cuda.device_count() >= world_size))
if distributed:
    torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo")
torch.manual_seed(config["seed"] + rank)
torch.set_num_threads(config["runtime"]["num_threads"])
device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu" if requested == "auto" else requested)
if use_cuda:
    device = torch.device(f"cuda:{local_rank}")
    torch.cuda.set_device(device)
model = build_model(config).to(device)
if distributed:
    model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None)
optimizer = torch.optim.Adam(model.parameters(), lr=config["training"]["learning_rate"])
dataset = WindowDataset(config["data"]["path"])
sampler = DistributedSampler(dataset, shuffle=True, seed=config["seed"]) if distributed else None
loader = DataLoader(dataset, batch_size=config["training"]["batch_size"], shuffle=sampler is None, sampler=sampler)
steps = args.steps if args.steps is not None else config["training"]["steps"]
losses = []
model.train()
for step, (inputs, target, lead) in enumerate(loader):
    if step >= steps:
        break
    optimizer.zero_grad(set_to_none=True)
    logits = model(inputs.to(device), lead.to(device), config["data"]["window"])
    loss = F.cross_entropy(logits, target.to(device))
    if not torch.isfinite(loss):
        raise FloatingPointError("training loss is not finite")
    loss.backward()
    optimizer.step()
    losses.append(float(loss))
    print(f"step={step} nll={losses[-1]:.6f}")
if not losses:
    raise RuntimeError("training produced no optimization steps")
summary = torch.tensor([sum(losses), len(losses)], dtype=torch.float64, device=device)
if distributed:
    torch.distributed.all_reduce(summary)
if rank == 0:
    save_checkpoint(config["paths"]["checkpoint"], model, config["model"])
    write_json(config["paths"]["training_metrics"],
               {"steps": int(summary[1]), "mean_nll": float(summary[0] / summary[1]), "world_size": world_size})
if distributed:
    torch.distributed.destroy_process_group()