CoolFace
Modelpublic

OneScience-Group/SatMAE

sourceHugging Facecc-by-nc-4.0updated 21d agoView on Hugging Face
0likes30downloads
train.py237 linesDownload Raw Back to scripts
1"""Pre-train SatMAE with masked reconstruction; supports torchrun."""2 3import argparse4import importlib.util5import json6import math7import os8import random9from contextlib import nullcontext10from functools import partial11from pathlib import Path12 13import numpy as np14import torch15import yaml16from torch import distributed as dist17from torch.nn.parallel import DistributedDataParallel18from torch.utils.data import DataLoader, Dataset, DistributedSampler19 20 21ROOT = Path(__file__).resolve().parents[1]22 23 24class NPZDataset(Dataset):25    def __init__(self, path, mode):26        archive = np.load(path)27        self.images = archive["images"]28        self.timestamps = archive["timestamps"] if "timestamps" in archive else None29        if mode == "temporal" and self.timestamps is None:30            raise ValueError("temporal datasets must contain timestamps")31 32    def __len__(self):33        return len(self.images)34 35    def __getitem__(self, index):36        images = torch.from_numpy(self.images[index])37        if self.timestamps is None:38            return images, torch.empty(0)39        return images, torch.from_numpy(self.timestamps[index])40 41 42def load_model_class():43    spec = importlib.util.spec_from_file_location("satmae", ROOT / "model/satmae.py")44    module = importlib.util.module_from_spec(spec)45    spec.loader.exec_module(module)46    return module.SatMAE47 48 49def model_config(config):50    return {51        key: value for key, value in config["model"].items()52        if key not in {"architecture", "runtime_profile"}53    }54 55 56def parse_args():57    parser = argparse.ArgumentParser(description=__doc__)58    parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml")59    parser.add_argument("--data", type=Path, default=None)60    parser.add_argument("--output", type=Path, default=None)61    parser.add_argument("--resume", type=Path, default=None)62    parser.add_argument("--epochs", type=int, default=None)63    parser.add_argument("--batch-size", type=int, default=None)64    parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default=None)65    return parser.parse_args()66 67 68def cosine_learning_rate(progress, config, peak_lr):69    warmup = config["warmup_epochs"]70    if warmup > 0 and progress < warmup:71        return peak_lr * progress / warmup72    span = max(config["epochs"] - warmup, 1)73    phase = min(max((progress - warmup) / span, 0.0), 1.0)74    return config["min_learning_rate"] + 0.5 * (75        peak_lr - config["min_learning_rate"]76    ) * (1.0 + math.cos(math.pi * phase))77 78 79def main():80    args = parse_args()81    config = yaml.safe_load(args.config.read_text())82    train_config = config["training"]83    if args.epochs is not None:84        train_config["epochs"] = args.epochs85    if args.batch_size is not None:86        train_config["batch_size"] = args.batch_size87 88    world_size = int(os.environ.get("WORLD_SIZE", "1"))89    local_rank = int(os.environ.get("LOCAL_RANK", "0"))90    rank = int(os.environ.get("RANK", "0"))91    distributed = world_size > 192    requested_device = args.device or config["runtime"]["device"]93    use_cuda = torch.cuda.is_available() and requested_device != "cpu"94    if requested_device == "cuda" and not torch.cuda.is_available():95        raise RuntimeError("CUDA was requested but is unavailable")96    if distributed:97        dist.init_process_group("nccl" if use_cuda else "gloo")98    device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")99    if use_cuda:100        torch.cuda.set_device(local_rank)101 102    seed = config["seed"] + rank103    random.seed(seed)104    np.random.seed(seed)105    torch.manual_seed(seed)106    data_path = args.data or ROOT / config["data"]["root"] / "train.npz"107    if not data_path.exists():108        raise FileNotFoundError(f"training data not found: {data_path}")109    dataset = NPZDataset(data_path, config["model"]["mode"])110    sampler = DistributedSampler(dataset, shuffle=True) if distributed else None111    loader = DataLoader(112        dataset,113        batch_size=train_config["batch_size"],114        shuffle=sampler is None,115        sampler=sampler,116        num_workers=train_config["num_workers"],117        pin_memory=use_cuda,118        drop_last=False,119    )120 121    model = load_model_class()(**model_config(config)).to(device)122    model_without_ddp = model123    if distributed:124        model = DistributedDataParallel(125            model, device_ids=[local_rank] if use_cuda else None126        )127        model_without_ddp = model.module128 129    effective_batch = (130        train_config["batch_size"] * train_config["accum_iter"] * world_size131    )132    peak_lr = train_config["learning_rate"]133    if peak_lr is None:134        peak_lr = train_config["base_learning_rate"] * effective_batch / 256135    decay, no_decay = [], []136    for name, parameter in model_without_ddp.named_parameters():137        if not parameter.requires_grad:138            continue139        (no_decay if parameter.ndim == 1 or name.endswith("bias") else decay).append(parameter)140    optimizer = torch.optim.AdamW(141        [142            {"params": decay, "weight_decay": train_config["weight_decay"]},143            {"params": no_decay, "weight_decay": 0.0},144        ],145        lr=peak_lr,146        betas=(0.9, 0.95),147    )148    amp_enabled = bool(config["runtime"].get("amp", True) and use_cuda)149    scaler = torch.amp.GradScaler("cuda", enabled=amp_enabled)150    start_epoch = 0151    history = []152    resume_path = args.resume153    if resume_path is None and train_config.get("resume"):154        resume_path = ROOT / train_config["resume"]155    if resume_path is not None:156        checkpoint = torch.load(resume_path, map_location="cpu", weights_only=False)157        model_without_ddp.load_state_dict(checkpoint["model"])158        optimizer.load_state_dict(checkpoint["optimizer"])159        if checkpoint.get("scaler") is not None:160            scaler.load_state_dict(checkpoint["scaler"])161        start_epoch = checkpoint["epoch"] + 1162        history = checkpoint.get("history", [])163 164    checkpoint_path = args.output or ROOT / config["paths"]["checkpoint"]165    metrics_path = ROOT / config["paths"]["training_metrics"]166    optimizer.zero_grad(set_to_none=True)167    for epoch in range(start_epoch, train_config["epochs"]):168        if sampler is not None:169            sampler.set_epoch(epoch)170        model.train()171        total_loss = 0.0172        steps = len(loader)173        for step, (images, timestamps) in enumerate(loader):174            progress = epoch + step / max(steps, 1)175            learning_rate = cosine_learning_rate(progress, train_config, peak_lr)176            for group in optimizer.param_groups:177                group["lr"] = learning_rate178            images = images.to(device, non_blocking=use_cuda)179            timestamps = timestamps.to(device, non_blocking=use_cuda)180            timestamps = timestamps if timestamps.numel() else None181            autocast = partial(torch.amp.autocast, "cuda") if amp_enabled else nullcontext182            with autocast():183                output = model(images, timestamps=timestamps)184                loss = output["loss"] / train_config["accum_iter"]185            if not torch.isfinite(loss):186                raise ValueError(f"non-finite loss at epoch {epoch}, step {step}")187            scaler.scale(loss).backward()188            update = (step + 1) % train_config["accum_iter"] == 0 or step + 1 == steps189            if update:190                scaler.step(optimizer)191                scaler.update()192                optimizer.zero_grad(set_to_none=True)193            total_loss += output["loss"].detach().item()194 195        epoch_loss = total_loss / max(steps, 1)196        record = {197            "epoch": epoch + 1,198            "reconstruction_loss": epoch_loss,199            "learning_rate": optimizer.param_groups[0]["lr"],200        }201        history.append(record)202        if rank == 0:203            print(204                f"epoch={epoch + 1} reconstruction_loss={epoch_loss:.6f} "205                f"lr={record['learning_rate']:.3e}"206            )207            if (epoch + 1) % train_config["save_every"] == 0 or epoch + 1 == train_config["epochs"]:208                checkpoint_path.parent.mkdir(parents=True, exist_ok=True)209                torch.save(210                    {211                        "model": model_without_ddp.state_dict(),212                        "optimizer": optimizer.state_dict(),213                        "scaler": scaler.state_dict() if amp_enabled else None,214                        "epoch": epoch,215                        "history": history,216                        "config": config,217                    },218                    checkpoint_path,219                )220 221    if rank == 0:222        metrics_path.parent.mkdir(parents=True, exist_ok=True)223        metrics_path.write_text(json.dumps({224            "history": history,225            "protocol": config["data"]["protocol"],226            "data_source": "synthetic" if "synthetic" in data_path.name or (data_path.parent / "format.json").exists() else "provided",227            "effective_batch_size": effective_batch,228            "peak_learning_rate": peak_lr,229        }, indent=2) + "\n")230        print("checkpoint=", checkpoint_path)231    if distributed:232        dist.destroy_process_group()233 234 235if __name__ == "__main__":236    main()237