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