Download scripts/train.py from OneScience-Group/MetNet-2: direct link, hf CLI and curl.
- Browser
- Download file 3.03 kB
-
https://huggingface.co/OneScience-Group/MetNet-2/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/MetNet-2/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/MetNet-2/resolve/main/scripts/train.py
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() | |