#!/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()