planning-baselines / code /train_models.py
greycat111's picture
Add checkpoint provenance, configurations, licenses and five-seed results
9375a59 verified
Raw
History Blame Contribute Delete
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
@torch.inference_mode()
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]