File size: 5,018 Bytes
d9ae24d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
"""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()