Earthformer / script /train.py
yzt15806542928's picture
Upload folder using huggingface_hub
8bfc737 verified
Raw
History Blame Contribute Delete
6.82 kB
from __future__ import annotations
import argparse
import json
import os
import random
import sys
from pathlib import Path
import numpy as np
import torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel
if __package__ in (None, ""):
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from model.earthformer import Earthformer
from script.data_loader import make_loader
from script.metrics import metric_sums, metric_sums_light, metrics_from_light_sums, metrics_from_sums, mse
from script.utils import DEFAULT_CONFIG, atomic_torch_save, clean_state_dict, load_checkpoint_payload, load_config, resolve_device
def seed_everything(seed: int) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)
def validate(
model: torch.nn.Module,
loader,
device: torch.device,
max_batches: int | None = None,
full_metrics: bool = False,
) -> dict[str, float]:
"""Run validation with only cheap MSE/MAE by default.
The training loss is already the per-sample MSE, so recomputing MSE during
training is redundant monitoring. SSIM and per-threshold CSI are expensive
on full-resolution 384x384 data, so they are computed only when
`train.compute_full_metrics` is enabled; the authoritative full evaluation
lives in script/result.py.
"""
model.eval()
sums = torch.zeros(22 if full_metrics else 3, dtype=torch.float64, device=device)
with torch.no_grad():
for batch_index, (inputs, targets) in enumerate(loader):
predictions = model(inputs.to(device, non_blocking=True)).clamp(0.0, 1.0)
targets_device = targets.to(device, non_blocking=True)
if full_metrics:
sums += metric_sums(predictions, targets_device)
else:
sums += metric_sums_light(predictions, targets_device)
if max_batches is not None and batch_index + 1 >= max_batches:
break
if dist.is_initialized():
dist.all_reduce(sums, op=dist.ReduceOp.SUM)
count_index = 3 if full_metrics else 2
if sums[count_index].item() == 0:
raise ValueError("validation loader is empty")
return (metrics_from_sums if full_metrics else metrics_from_light_sums)(sums.cpu())
def run_training(config: dict, requested_device: str = "auto", resume: str | None = None) -> tuple[Path, dict[str, float]]:
world_size = int(os.environ.get("WORLD_SIZE", "1"))
rank = int(os.environ.get("RANK", "0"))
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
distributed = world_size > 1
device = resolve_device(requested_device, local_rank)
backend = config["distributed"].get("backend", "auto")
if backend == "auto":
backend = "nccl" if device.type == "cuda" else "gloo"
if distributed:
dist.init_process_group(backend=backend, rank=rank, world_size=world_size)
try:
seed_everything(int(config["train"]["seed"]))
model = Earthformer(config).to(device)
optimizer = torch.optim.AdamW(
model.parameters(),
lr=float(config["train"]["learning_rate"]),
weight_decay=float(config["train"]["weight_decay"]),
)
start_epoch, step = 0, 0
if resume:
payload = load_checkpoint_payload(resume, device)
model.load_state_dict(clean_state_dict(payload["model"]))
if "optimizer" in payload:
optimizer.load_state_dict(payload["optimizer"])
start_epoch = int(payload.get("epoch", -1)) + 1
step = int(payload.get("step", 0))
if distributed:
model = DistributedDataParallel(
model,
device_ids=[local_rank] if device.type == "cuda" else None,
find_unused_parameters=True,
)
train_loader, train_sampler = make_loader(config, "train", distributed, rank, world_size)
val_loader, _ = make_loader(config, "val", distributed, rank, world_size, shuffle=False)
full_metrics = bool(config["train"].get("compute_full_metrics", False))
last_epoch = max(start_epoch - 1, 0)
metrics: dict[str, float] = {}
for epoch in range(start_epoch, int(config["train"]["epochs"])):
last_epoch = epoch
if train_sampler is not None:
train_sampler.set_epoch(epoch)
model.train()
epoch_loss_sum, epoch_steps = 0.0, 0
for inputs, targets in train_loader:
optimizer.zero_grad(set_to_none=True)
loss = mse(model(inputs.to(device, non_blocking=True)), targets.to(device, non_blocking=True))
if not torch.isfinite(loss):
raise FloatingPointError("training loss is not finite")
loss.backward()
optimizer.step()
epoch_loss_sum += float(loss.detach().item())
epoch_steps += 1
step += 1
if epoch_steps == 0:
break
metrics = validate(model, val_loader, device, int(config["train"].get("validation_steps", 1)), full_metrics)
if rank == 0:
epoch_line = {"epoch": epoch, "train_loss": epoch_loss_sum / epoch_steps, "validation": metrics}
print(json.dumps(epoch_line))
checkpoint = Path(config["train"]["output_dir"]) / "earthformer.pt"
if rank == 0:
raw_model = model.module if isinstance(model, DistributedDataParallel) else model
atomic_torch_save(
{
"model": raw_model.state_dict(),
"optimizer": optimizer.state_dict(),
"config": config,
"metrics": metrics,
"epoch": last_epoch,
"step": step,
"world_size": world_size,
},
checkpoint,
)
print(json.dumps({"checkpoint": str(checkpoint), "step": step, "world_size": world_size, "metrics": metrics}, indent=2))
if distributed:
dist.barrier()
return checkpoint, metrics
finally:
if dist.is_initialized():
dist.destroy_process_group()
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Train Earthformer with one process or DDP")
parser.add_argument("--config", default=str(DEFAULT_CONFIG))
parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
parser.add_argument("--resume")
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
run_training(load_config(args.config), args.device, args.resume)