OneScience-Group/MetNet-2
026
1#!/usr/bin/env python32import argparse3import os4from pathlib import Path5 6import torch7import torch.nn.functional as F8from torch.nn.parallel import DistributedDataParallel9from torch.utils.data import DataLoader, DistributedSampler10from model.metnet_2 import (WindowDataset, build_model, categorical_nll_chunked,11 load_config, save_checkpoint, write_json)12 13parser = argparse.ArgumentParser(description="Train MetNet-2 on selected windows")14parser.add_argument("--config", default="conf/config.yaml")15parser.add_argument("--steps", type=int, default=None)16args = parser.parse_args()17config = load_config(args.config)18rank = int(os.environ.get("RANK", "0"))19world_size = int(os.environ.get("WORLD_SIZE", "1"))20local_rank = int(os.environ.get("LOCAL_RANK", "0"))21distributed = world_size > 122requested = config["runtime"]["device"]23use_cuda = (requested != "cpu" and torch.cuda.is_available()24 and (not distributed or torch.cuda.device_count() >= world_size))25if distributed:26 torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo")27torch.manual_seed(config["seed"] + rank)28torch.set_num_threads(config["runtime"]["num_threads"])29device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu" if requested == "auto" else requested)30if use_cuda:31 device = torch.device(f"cuda:{local_rank}")32 torch.cuda.set_device(device)33model = build_model(config).to(device)34if distributed:35 model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None)36optimizer = torch.optim.Adam(model.parameters(), lr=config["training"]["learning_rate"])37dataset = WindowDataset(config["data"]["path"])38sampler = DistributedSampler(dataset, shuffle=True, seed=config["seed"]) if distributed else None39loader = DataLoader(dataset, batch_size=config["training"]["batch_size"], shuffle=sampler is None, sampler=sampler)40steps = args.steps if args.steps is not None else config["training"]["steps"]41losses = []42model.train()43for step, (inputs, target, lead) in enumerate(loader):44 if step >= steps:45 break46 optimizer.zero_grad(set_to_none=True)47 logits = model(inputs.to(device), lead.to(device), config["data"]["window"])48 loss = F.cross_entropy(logits, target.to(device))49 if not torch.isfinite(loss):50 raise FloatingPointError("training loss is not finite")51 loss.backward()52 optimizer.step()53 losses.append(float(loss))54 print(f"step={step} nll={losses[-1]:.6f}")55if not losses:56 raise RuntimeError("training produced no optimization steps")57summary = torch.tensor([sum(losses), len(losses)], dtype=torch.float64, device=device)58if distributed:59 torch.distributed.all_reduce(summary)60if rank == 0:61 save_checkpoint(config["paths"]["checkpoint"], model, config["model"])62 write_json(config["paths"]["training_metrics"],63 {"steps": int(summary[1]), "mean_nll": float(summary[0] / summary[1]), "world_size": world_size})64if distributed:65 torch.distributed.destroy_process_group()66 