Download code/train_baseline.py from FidelityWM/planning-baselines: direct link, hf CLI and curl.
- Browser
- Download file 8.17 kB
-
https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/train_baseline.py
- Command line
-
hf download hf://FidelityWM/planning-baselines/code/train_baseline.py
-
curl -L -o train_baseline.py https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/train_baseline.py
8.17 kB
| """Ten full epochs on the supplied dataset's seeded 90% training split.""" | |
| import argparse,json,math,os,random,sys,time | |
| from pathlib import Path | |
| import numpy as np | |
| ROOT=Path(__file__).resolve().parents[1] | |
| sys.path.insert(0,str(ROOT/'runtime')) | |
| import torch | |
| from torch.utils.data import DataLoader,Subset | |
| from sklearn.preprocessing import StandardScaler | |
| import stable_worldmodel as swm | |
| from stable_worldmodel.wm.loss import PLDMLoss,SIGReg | |
| from torchvision.transforms import v2 as T | |
| from train_models import build_model | |
| from reporting import atomic_json,save_sheet | |
| FILES={'tworoom':'tworoom.h5','cube':'cube_single_expert.h5','pusht':'pusht_expert_train.h5'} | |
| def main(): | |
| ap=argparse.ArgumentParser();ap.add_argument('--method',choices=['pldm','dinowm'],required=True);ap.add_argument('--task',choices=FILES,required=True);ap.add_argument('--epochs',type=int,default=10);ap.add_argument('--batch',type=int);ap.add_argument('--smoke',action='store_true');ap.add_argument('--smoke-dynamics',action='store_true');args=ap.parse_args() | |
| torch.set_num_threads(4);torch.manual_seed(42);np.random.seed(42);random.seed(42) | |
| directory=ROOT/'checkpoints/trained'/f'{args.method}_{args.task}';directory.mkdir(parents=True,exist_ok=True) | |
| state=ROOT/'results'/f'training_{args.method}_{args.task}.json' | |
| ds=swm.data.HDF5Dataset(path='/workspace/datasets/'+FILES[args.task],frameskip=5,num_steps=4,keys_to_load=['pixels','action'],keys_to_cache=['action']) | |
| action=ds.get_col_data('action');normalizer=StandardScaler().fit(action[np.isfinite(action).all(1)]);action_dim=action.shape[-1] | |
| generator=torch.Generator().manual_seed(42);indices=torch.randperm(len(ds),generator=generator);ntrain=int(.9*len(ds)) | |
| train=Subset(ds,indices[:ntrain].tolist()) | |
| batch=args.batch or (128 if args.method in ['pldm','lejepa'] else 32) | |
| loader=DataLoader(train,batch_size=batch,shuffle=True,num_workers=2,multiprocessing_context="spawn",persistent_workers=True,prefetch_factor=2,pin_memory=True,drop_last=True,generator=generator) | |
| model=build_model(args.method,action_dim).cuda() | |
| opt=torch.optim.AdamW([p for p in model.parameters() if p.requires_grad],lr=5e-5 if args.method=='pldm' else 5e-4,weight_decay=1e-3 if args.method=='pldm' else 0.) | |
| sigreg=SIGReg().cuda();augment=T.Compose([T.RandomResizedCrop((224,224),scale=(.3,1.)),T.RandomHorizontalFlip(),T.ColorJitter(.4,.4,.2,.1),T.RandomGrayscale(p=.2)]) | |
| reg=PLDMLoss().cuda();loss_weights={'std_loss':18.,'std_t_loss':.7,'cov_loss':12.,'cov_t_loss':0.,'temp_align_loss':.2} | |
| means=torch.tensor(normalizer.mean_,device='cuda',dtype=torch.float32);stds=torch.tensor(normalizer.scale_,device='cuda',dtype=torch.float32) | |
| image_mean=torch.tensor([.485,.456,.406],device='cuda').view(1,1,3,1,1);image_std=torch.tensor([.229,.224,.225],device='cuda').view(1,1,3,1,1) | |
| metadata=dict(method=args.method,task=args.task,variant='PLDM stable-worldmodel recipe' if args.method=='pldm' else ('LeJEPA planning adaptation: 5 epochs two-view image SSL + 5 epochs frozen-encoder dynamics' if args.method=='lejepa' else 'DINO-WM without proprioception; stable-worldmodel PreJEPA'),epochs_required=args.epochs,training_seed=42,train_clips=ntrain,total_clips=len(ds),batch_size=batch,batches_per_epoch=len(loader),frameskip=5,history_size=3,image_size=224,precision='bf16-mixed',train_split=.9,action_mean=normalizer.mean_.tolist(),action_std=normalizer.scale_.tolist(),optimizer='AdamW',lr=opt.param_groups[0]['lr'],weight_decay=opt.param_groups[0]['weight_decay'],loss_weights=loss_weights if args.method=='pldm' else {},note='Full epochs of the 90% clip split, no capped training subset. Evaluation pairs remain the shared existing dataset pairs, not a held-out test set.') | |
| start_epoch=5 if args.smoke and args.smoke_dynamics and args.method=='lejepa' else 0;total_steps=0;history=[] | |
| latest=directory/'latest.pt' | |
| if latest.exists() and not args.smoke: | |
| ck=torch.load(latest,map_location='cpu',weights_only=False);model.load_state_dict(ck['model']);opt.load_state_dict(ck['optimizer']);start_epoch=ck['epochs_completed'];total_steps=ck['total_steps'];history=ck['history'];generator.set_state(ck['loader_rng']);torch.set_rng_state(ck['torch_rng']);torch.cuda.set_rng_state_all(ck['cuda_rng']) | |
| if not args.smoke:atomic_json(directory/'config.json',metadata) | |
| base_lr=metadata['lr'];all_steps=args.epochs*len(loader);warmup=max(1,int(.01*all_steps)) | |
| for epoch in range(start_epoch,args.epochs): | |
| model.train() | |
| if args.method=='dinowm':model.backbone.eval() | |
| if args.method=='lejepa': | |
| ssl=epoch<5 | |
| for module in [model.encoder,model.projector]:module.requires_grad_(ssl);module.train(ssl) | |
| running=0.;t0=time.time() | |
| for i,data in enumerate(loader): | |
| pixels=data['pixels'].cuda(non_blocking=True).float() | |
| if pixels.shape[-1]==3:pixels=pixels.permute(0,1,4,2,3) | |
| pixels=(pixels/255.-image_mean)/image_std | |
| a=data['action'].cuda(non_blocking=True).float().reshape(-1,4,5,action_dim) | |
| a=torch.nan_to_num((a-means)/stds).flatten(-2) | |
| factor=min(1.,(total_steps+1)/warmup)*(.5*(1+math.cos(math.pi*max(0,total_steps-warmup)/max(1,all_steps-warmup)))) | |
| for group in opt.param_groups:group['lr']=base_lr*factor | |
| opt.zero_grad(set_to_none=True) | |
| with torch.autocast('cuda',dtype=torch.bfloat16): | |
| if args.method=='lejepa' and epoch<5: | |
| raw=pixels[:,0]*image_std[:,0]+image_mean[:,0] | |
| views=(torch.stack([augment(raw),augment(raw)],1)-image_mean)/image_std | |
| z=model.encode({'pixels':views})['emb'].float() | |
| prediction=(z-z.mean(1,keepdim=True)).square().mean() | |
| loss=.95*prediction+.05*sigreg(z.transpose(0,1)) | |
| else: | |
| encoded=model.encode({'pixels':pixels,'action':a});z=encoded['emb'] | |
| if args.method in ['pldm','lejepa']: | |
| pred=model.predict(z[:,:3],encoded['act_emb'][:,:3]);prediction=(pred.float()-z[:,1:].float()).square().mean() | |
| terms=reg(z.float()) if args.method=='pldm' else {} | |
| loss=prediction+sum(terms[k]*w for k,w in loss_weights.items()) if args.method=='pldm' else prediction | |
| else: | |
| pred=model.predict(z[:,:3]);prediction=(pred[...,:384].float()-z[:,1:,...,:384].detach().float()).square().mean();loss=prediction | |
| if not torch.isfinite(loss):raise RuntimeError('Nonfinite training loss') | |
| loss.backward();norm=torch.nn.utils.clip_grad_norm_(model.parameters(),1.) | |
| if not torch.isfinite(norm):raise RuntimeError('Nonfinite gradient') | |
| opt.step();total_steps+=1;running+=float(loss.detach()) | |
| if args.smoke: | |
| print(json.dumps(dict(smoke='passed',method=args.method,task=args.task,loss=float(loss.detach()),gradient_norm=float(norm),batch_size=batch,train_clips=ntrain,batches_per_epoch=len(loader),peak_gpu_gb=torch.cuda.max_memory_allocated()/1e9)),flush=True) | |
| # Save temporary model for an end-to-end planning smoke check, never as a measured final checkpoint. | |
| model.eval();torch.save({'model':model.cpu(),'metadata':dict(metadata,epochs_completed=0,smoke=True)},directory/'smoke.pt');return | |
| if i%100==0: | |
| atomic_json(state,dict(metadata,status='training',epoch=epoch+1,epochs_completed=epoch,batch=i+1,total_steps=total_steps,mean_loss=running/(i+1),epoch_elapsed_seconds=time.time()-t0)) | |
| save_sheet() | |
| print(f'TRAIN {args.method} {args.task} epoch={epoch+1}/{args.epochs} batch={i+1}/{len(loader)} loss={running/(i+1):.6f}',flush=True) | |
| history.append(dict(epoch=epoch+1,loss=running/len(loader),seconds=time.time()-t0)) | |
| temp=latest.with_suffix('.tmp');torch.save(dict(model=model.state_dict(),optimizer=opt.state_dict(),epochs_completed=epoch+1,total_steps=total_steps,history=history,loader_rng=generator.get_state(),torch_rng=torch.get_rng_state(),cuda_rng=torch.cuda.get_rng_state_all()),temp);temp.replace(latest) | |
| atomic_json(state,dict(metadata,status='training' if epoch+1<args.epochs else 'trained',epochs_completed=epoch+1,total_steps=total_steps,history=history)) | |
| model.eval();model.cpu();out=directory/'model.pt';temp=out.with_suffix('.tmp');torch.save({'model':model,'metadata':dict(metadata,epochs_completed=args.epochs,history=history)},temp);temp.replace(out) | |
| latest.unlink(missing_ok=True) | |
| print(f'TRAINED {args.method} {args.task}: {args.epochs} epochs',flush=True) | |
| if __name__=='__main__':main() | |