Download scripts/train.py from OneScience-Group/Scale-MAE: direct link, hf CLI and curl.
- Browser
- Download file 5.02 kB
-
https://huggingface.co/OneScience-Group/Scale-MAE/resolve/main/scripts/train.py
- Command line
-
hf download hf://OneScience-Group/Scale-MAE/scripts/train.py
-
curl -L -o train.py https://huggingface.co/OneScience-Group/Scale-MAE/resolve/main/scripts/train.py
5.02 kB
| """Train Scale-MAE with DDP, AMP, accumulation, warmup and cosine decay.""" | |
| import argparse, importlib.util, json, os, random | |
| from pathlib import Path | |
| import numpy as np, torch, 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] | |
| class NPZDataset(Dataset): | |
| def __init__(self,path): | |
| a=np.load(path); self.images,self.targets,self.gsd=a["images"],a["targets"],a["gsd"] | |
| def __len__(self): return len(self.images) | |
| def __getitem__(self,i): return torch.from_numpy(self.images[i]), torch.from_numpy(self.targets[i]), torch.tensor(self.gsd[i], dtype=torch.float32) | |
| def model_class(): | |
| spec=importlib.util.spec_from_file_location("scalemae",ROOT/"model/scalemae.py"); mod=importlib.util.module_from_spec(spec); spec.loader.exec_module(mod); return mod.ScaleMAE | |
| def main(): | |
| p=argparse.ArgumentParser(); p.add_argument("--config",type=Path,default=ROOT/"conf/config.yaml"); p.add_argument("--data",type=Path); p.add_argument("--output",type=Path); p.add_argument("--epochs",type=int); p.add_argument("--batch-size",type=int); p.add_argument("--device",choices=("auto","cpu","cuda")); a=p.parse_args(); cfg=yaml.safe_load(a.config.read_text()) | |
| if a.epochs is not None: cfg["training"]["epochs"] = a.epochs | |
| if a.batch_size is not None: cfg["training"]["batch_size"] = a.batch_size | |
| world=int(os.environ.get("WORLD_SIZE","1")); rank=int(os.environ.get("RANK","0")); local=int(os.environ.get("LOCAL_RANK","0")); requested=a.device or cfg["runtime"]["device"] | |
| if requested == "cuda" and not torch.cuda.is_available(): raise RuntimeError("CUDA requested but unavailable") | |
| cuda=torch.cuda.is_available() and requested!="cpu"; device=torch.device(f"cuda:{local}" if cuda else "cpu") | |
| if cuda: torch.cuda.set_device(local) | |
| if world>1: dist.init_process_group("nccl" if cuda else "gloo") | |
| seed=cfg["seed"]+rank; random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) | |
| source=a.data or ROOT/cfg["data"]["root"]/"train.npz" | |
| if not source.exists(): raise FileNotFoundError("Run scripts/fake_data.py first") | |
| ds=NPZDataset(source); sampler=DistributedSampler(ds,num_replicas=world,rank=rank) if world>1 else None; loader=DataLoader(ds,batch_size=cfg["training"]["batch_size"],shuffle=sampler is None,sampler=sampler,num_workers=cfg["training"]["num_workers"]) | |
| model=model_class()(**cfg["model"]).to(device); model=DistributedDataParallel(model,device_ids=[local]) if world>1 else model; base=model.module if hasattr(model,"module") else model | |
| opt=torch.optim.AdamW(model.parameters(),lr=cfg["training"]["learning_rate"],betas=(.9,.95),weight_decay=cfg["training"]["weight_decay"]); steps=max(1,cfg["training"]["epochs"]*((len(ds)+cfg["training"]["batch_size"]-1)//cfg["training"]["batch_size"])); warm=max(1,int(steps*cfg["training"]["warmup_fraction"])); accum=cfg["training"]["gradient_accumulation"] | |
| amp=torch.amp.GradScaler("cuda", enabled=cuda and cfg["training"]["amp"]); history=[]; opt.zero_grad(set_to_none=True); step=0 | |
| for epoch in range(cfg["training"]["epochs"]): | |
| if sampler: sampler.set_epoch(epoch) | |
| model.train(); totals=torch.zeros(3,device=device) | |
| for batch_idx,(images,targets,gsd) in enumerate(loader): | |
| with torch.autocast(device_type="cuda",enabled=cuda and cfg["training"]["amp"]): out=model(images.to(device),gsd.to(device),target=targets.to(device)); loss=out["loss"]/accum | |
| if not torch.isfinite(loss): raise FloatingPointError("non-finite training loss") | |
| amp.scale(loss).backward(); totals += torch.stack([out["loss"].detach(),out["low_loss"].detach(),out["high_loss"].detach()]) | |
| if (batch_idx+1)%accum==0 or batch_idx+1==len(loader): | |
| amp.unscale_(opt); torch.nn.utils.clip_grad_norm_(model.parameters(),1.0); amp.step(opt); amp.update(); opt.zero_grad(set_to_none=True); step+=1; lr=cfg["training"]["learning_rate"]*(step/warm if step<=warm else .5*(1+np.cos(np.pi*(step-warm)/max(1,steps-warm)))); [g.update(lr=lr) for g in opt.param_groups] | |
| if world>1: dist.all_reduce(totals); totals/=world | |
| totals/=max(len(loader),1); history.append({"epoch":epoch+1,"loss":totals[0].item(),"low_frequency_loss":totals[1].item(),"high_frequency_loss":totals[2].item()}) | |
| if rank==0: print(history[-1]) | |
| if rank==0: | |
| ck=a.output or ROOT/cfg["paths"]["checkpoint"]; met=ROOT/cfg["paths"]["training_metrics"]; ck.parent.mkdir(parents=True,exist_ok=True); met.parent.mkdir(parents=True,exist_ok=True); torch.save({"model":base.state_dict(),"config":cfg,"optimizer":opt.state_dict(),"epoch":cfg["training"]["epochs"],"seed":cfg["seed"]},ck); met.write_text(json.dumps({"history":history,"protocol":cfg["data"]["protocol"],"data_source":"synthetic"},indent=2)+"\n"); print("checkpoint=",ck) | |
| if world>1: dist.destroy_process_group() | |
| if __name__=="__main__": main() | |