"""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