MetNet-2 / scripts /train.py
zhangrenchao's picture
Publish MetNet-2 reproduction
efe4fbe verified
Raw
History Blame Contribute Delete
3.03 kB
#!/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()