Download code/train_models.py from FidelityWM/planning-baselines: direct link, hf CLI and curl.
- Browser
- Download file 4.12 kB
-
https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/train_models.py
- Command line
-
hf download hf://FidelityWM/planning-baselines/code/train_models.py
-
curl -L -o train_models.py https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/train_models.py
4.12 kB
| """Model definitions for locally trained, explicitly labeled baseline checkpoints.""" | |
| import sys | |
| from pathlib import Path | |
| from types import SimpleNamespace | |
| import torch | |
| from torch import nn | |
| ROOT=Path(__file__).resolve().parents[1] | |
| sys.path.insert(0,str(ROOT/'runtime')) | |
| for source in (ROOT/'checkpoints/dino-hub').glob('facebookresearch_dinov2*'):sys.path.insert(0,str(source)) | |
| class DinoBackbone(nn.Module): | |
| def __init__(self): | |
| super().__init__() | |
| hub=ROOT/'checkpoints/dino-hub' | |
| torch.hub.set_dir(str(hub)) | |
| source=next(hub.glob('facebookresearch_dinov2*')) | |
| self.model=torch.hub.load(str(source),'dinov2_vits14',source='local',pretrained=True) | |
| self.requires_grad_(False) | |
| def forward(self,x,**kwargs): | |
| with torch.no_grad():z=self.model.forward_features(x)['x_norm_patchtokens'] | |
| return SimpleNamespace(last_hidden_state=torch.cat([z.new_zeros(z.shape[0],1,z.shape[-1]),z],1)) | |
| def build_model(method,action_dim): | |
| from stable_worldmodel.wm.pldm.pldm import PLDM | |
| from stable_worldmodel.wm.pldm.module import Predictor,Embedder,MLP | |
| if method in ['pldm','lejepa']: | |
| from stable_pretraining.backbone.utils import vit_hf | |
| return PLDM(encoder=vit_hf(size='tiny',patch_size=14,image_size=224,pretrained=False,use_mask_token=False),predictor=Predictor(num_frames=3,input_dim=192,hidden_dim=192,output_dim=192,depth=6,heads=16,mlp_dim=2048,dim_head=64,dropout=.1,emb_dropout=0),action_encoder=Embedder(input_dim=5*action_dim,emb_dim=192),projector=MLP(input_dim=192,output_dim=192,hidden_dim=2048,norm_fn=nn.BatchNorm1d),pred_proj=MLP(input_dim=192,output_dim=192,hidden_dim=2048,norm_fn=nn.BatchNorm1d)) | |
| from stable_worldmodel.wm.prejepa.prejepa import PreJEPA | |
| from stable_worldmodel.wm.prejepa.module import CausalPredictor,Embedder as PatchEmbedder | |
| return PreJEPA(encoder=DinoBackbone(),predictor=CausalPredictor(num_patches=256,num_frames=3,dim=394,depth=6,heads=16,mlp_dim=2048,dim_head=64,dropout=.1,emb_dropout=0),extra_encoders=nn.ModuleDict({'action':PatchEmbedder(in_chans=action_dim*5,emb_dim=10)}),history_size=3,num_pred=1) | |
| class TrainedCost(nn.Module): | |
| """Cache visual encoding, then use the trained predictor at every plan step.""" | |
| def __init__(self,model,method): | |
| super().__init__();self.model=model;self.method=method;self._info=None | |
| def get_cost(self,info,candidates): | |
| assert candidates.shape[0]==1 | |
| if self._info is not info: | |
| self._info=info | |
| if self.method=='dinowm': | |
| self.initial=self.model._encode_image(info['pixels'][:,0]) | |
| self.goal=self.model._encode_image(info['goal'][:,0])[:,-1:] | |
| else: | |
| self.initial=self.model.encode({'pixels':info['pixels'][:,0]})['emb'] | |
| self.goal=self.model.encode({'pixels':info['goal'][:,0]})['emb'][:,-1:] | |
| costs=[] | |
| for actions in candidates[0].split(100 if self.method=='dinowm' else 300): | |
| z=self.initial.expand(actions.shape[0],*self.initial.shape[1:]) | |
| if self.method=='dinowm': | |
| # Single initial frame in the common evaluator; each action is a five-step block. | |
| ae=self.model.extra_encoders['action'](actions) | |
| z=torch.cat([z,ae[:,:1,None,:].expand(-1,-1,z.shape[2],-1)],-1) | |
| for t in range(actions.shape[1]): | |
| pred=self.model.predict(z[:,-3:])[:,-1:] | |
| if t+1<actions.shape[1]: | |
| pred=torch.cat([pred[...,:384],ae[:,t+1:t+2,None,:].expand(-1,-1,z.shape[2],-1)],-1) | |
| z=torch.cat([z,pred],1) | |
| error=pred[...,:384]-self.goal | |
| else: | |
| ae=self.model.action_encoder(actions) | |
| for t in range(actions.shape[1]): | |
| pred=self.model.predict(z[:,-3:],ae[:,max(0,t-2):t+1])[:,-1:] | |
| z=torch.cat([z,pred],1) | |
| error=pred-self.goal | |
| costs.append(error.square().flatten(1).mean(1)) | |
| return torch.cat(costs)[None] | |