OneScience-Group/SatMAE
030
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 