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