File size: 8,172 Bytes
9375a59
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
"""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()