Team Ai
Modelpublic

OneScience-Group/MetNet-2

sourceHugging Faceapache-2.0updated 29d agoView on Hugging Face
0likes26downloads
train.py66 linesDownload Raw Back to scripts
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