| from pathlib import Path |
| import sys,copy,os |
| import numpy as np, torch |
| from torch.nn.parallel import DistributedDataParallel |
| ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT)) |
| from model.scale_adaptive_cm import ScaleAdaptiveCM,load_config,write_json |
| c=load_config(ROOT);rank=int(os.environ.get("RANK",0));world=int(os.environ.get("WORLD_SIZE",1));local=int(os.environ.get("LOCAL_RANK",0));distributed=world>1;use_cuda=torch.cuda.is_available() and torch.cuda.device_count()>=world and c["runtime"]["device"]!="cpu" |
| if distributed: torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo") |
| device=torch.device(f"cuda:{local}" if use_cuda else "cpu");torch.manual_seed(c["seed"]);torch.set_num_threads(2) |
| d=np.load(ROOT/c["data"]["path"]); x=torch.from_numpy(np.log(d["high"][d["split"]=="train"]+c["data"]["log_epsilon"])-np.log(c["data"]["log_epsilon"])).float() |
| mean,std=x.mean(),x.std().clamp_min(1e-6);x=(x-mean)/std |
| m=ScaleAdaptiveCM(**c["model"]).to(device);target=copy.deepcopy(m).eval();wrapped=DistributedDataParallel(m,device_ids=[local] if use_cuda else None) if distributed else m;opt=torch.optim.RAdam(wrapped.parameters(),lr=c["train"]["learning_rate"]);hist=[];x=x.to(device) |
| for e in range(c["train"]["epochs"]): |
| noise=torch.randn_like(x);t1=torch.full((len(x),),.35);t2=torch.full((len(x),),.55) |
| with torch.no_grad(): y=target(x+t1[:,None,None,None]*noise,t1) |
| pred=wrapped(x+t2[:,None,None,None]*noise,t2);loss=(pred-y).abs().mean();opt.zero_grad();loss.backward();opt.step() |
| with torch.no_grad(): |
| for p,q in zip(target.parameters(),m.parameters()):p.mul_(c["train"]["ema_decay"]).add_(q,alpha=1-c["train"]["ema_decay"]) |
| hist.append({"epoch":e+1,"loss":float(loss)}) |
| summary=torch.tensor([sum(v["loss"] for v in hist),len(hist)],dtype=torch.float64,device=device) |
| if distributed: torch.distributed.all_reduce(summary) |
| if rank==0: |
| path=ROOT/c["paths"]["checkpoint"];path.parent.mkdir(parents=True,exist_ok=True);torch.save({"model":target.state_dict(),"model_config":c["model"],"format_version":c["data"]["format_version"],"normalization":{"mean":float(mean),"std":float(std)}},path) |
| write_json(ROOT/c["paths"]["training_metrics"],{"history":hist,"global_mean_loss":float(summary[0]/summary[1]),"world_size":world});print(path) |
| if distributed: torch.distributed.destroy_process_group() |
|
|