SatlasPretrain / scripts /train.py
zhangrenchao's picture
Update SatlasPretrain model package
aed6f6f verified
Raw
History Blame Contribute Delete
5.75 kB
"""Train compact SatlasNet on seven task families; supports AMP and torchrun."""
import argparse
import importlib.util
import json
import os
from contextlib import nullcontext
from functools import partial
from pathlib import Path
import numpy as np
import torch
import yaml
from torch import distributed as dist
from torch.nn.parallel import DistributedDataParallel
from torch.utils.data import DataLoader, Dataset, DistributedSampler
ROOT = Path(__file__).resolve().parents[1]
def load_module():
spec = importlib.util.spec_from_file_location("satlaspretrain", ROOT / "model/satlaspretrain.py")
module = importlib.util.module_from_spec(spec); spec.loader.exec_module(module); return module
class NpzDataset(Dataset):
keys = ("highres_images", "lowres_images", "valid_highres_times", "valid_lowres_times",
"segmentation", "regression", "point", "polygon", "polyline", "property", "classification")
def __init__(self, path, config, module):
archive = np.load(path)
module.validate_npz(archive, config)
self.data = {key: archive[key] for key in self.keys}
self.source = str(archive["source"])
self.protocol = str(archive["protocol"]) if "protocol" in archive else "provided_npz"
def __len__(self): return len(self.data["highres_images"])
def __getitem__(self, index): return {key: torch.as_tensor(value[index]) for key, value in self.data.items()}
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml"); parser.add_argument("--data", type=Path)
parser.add_argument("--checkpoint", type=Path); parser.add_argument("--device", choices=("auto", "cpu", "cuda"))
args = parser.parse_args(); config = yaml.safe_load(args.config.read_text()); train_cfg = config["training"]
world_size, rank, local_rank = int(os.environ.get("WORLD_SIZE", 1)), int(os.environ.get("RANK", 0)), int(os.environ.get("LOCAL_RANK", 0))
requested = args.device or config["runtime"]["device"]; use_cuda = torch.cuda.is_available() and requested != "cpu"
if requested == "cuda" and not use_cuda: raise RuntimeError("CUDA requested but unavailable")
if world_size > 1: dist.init_process_group("nccl" if use_cuda else "gloo")
device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")
if use_cuda: torch.cuda.set_device(local_rank)
torch.manual_seed(config["seed"] + rank)
module = load_module()
data_path = args.data or ROOT / config["data"]["root"] / "train.npz"; dataset = NpzDataset(data_path, config, module)
sampler = DistributedSampler(dataset, shuffle=True) if world_size > 1 else None
loader = DataLoader(dataset, batch_size=train_cfg["batch_size"], sampler=sampler, shuffle=sampler is None,
num_workers=train_cfg["num_workers"], pin_memory=use_cuda)
model = module.SatlasPretrain(**config["model"]).to(device); raw_model = model
if world_size > 1:
model = DistributedDataParallel(model, device_ids=[local_rank] if use_cuda else None); raw_model = model.module
optimizer = torch.optim.AdamW(model.parameters(), lr=train_cfg["learning_rate"], weight_decay=train_cfg["weight_decay"])
amp = bool(config["runtime"]["amp"] and use_cuda); scaler = torch.amp.GradScaler("cuda", enabled=amp); history = []
for epoch in range(train_cfg["epochs"]):
if sampler is not None: sampler.set_epoch(epoch)
model.train(); totals = {"total": 0.0}; steps = 0
for batch in loader:
batch = {key: value.to(device, non_blocking=use_cuda) for key, value in batch.items()}
autocast = partial(torch.amp.autocast, "cuda") if amp else nullcontext
with autocast():
outputs = model(batch["highres_images"], batch["lowres_images"],
batch["valid_highres_times"], batch["valid_lowres_times"])
loss, parts = module.multitask_loss(outputs, batch)
if not torch.isfinite(loss): raise ValueError("non-finite multitask loss")
optimizer.zero_grad(set_to_none=True); scaler.scale(loss).backward(); scaler.step(optimizer); scaler.update()
totals["total"] += loss.detach().item(); steps += 1
for name, value in parts.items(): totals[name] = totals.get(name, 0.0) + value.detach().item()
names = list(totals); statistics = torch.tensor([totals[name] for name in names] + [steps], dtype=torch.float64, device=device)
if world_size > 1: dist.all_reduce(statistics, op=dist.ReduceOp.SUM)
global_steps = max(statistics[-1].item(), 1)
record = {"epoch": epoch + 1, **{f"{name}_loss": statistics[index].item() / global_steps for index, name in enumerate(names)}}; history.append(record)
if rank == 0: print(json.dumps(record))
if rank == 0:
checkpoint_path = args.checkpoint or ROOT / config["paths"]["checkpoint"]; checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
torch.save({"model": raw_model.state_dict(), "optimizer": optimizer.state_dict(), "scaler": scaler.state_dict() if amp else None,
"config": config, "epoch": train_cfg["epochs"], "history": history}, checkpoint_path)
metrics = ROOT / config["paths"]["training_metrics"]; metrics.parent.mkdir(parents=True, exist_ok=True)
metrics.write_text(json.dumps({"history": history, "protocol": dataset.protocol,
"source": dataset.source,
"world_size": world_size, "amp": amp}, indent=2) + "\n")
print(f"checkpoint={checkpoint_path}")
if world_size > 1: dist.destroy_process_group()
if __name__ == "__main__": main()