"""Original released DINO-WM PushT model with the shared LeWM evaluator. Keep common CEM coordinates; translate inputs to the checkpoint's native normalization. Chunk candidates only for memory, without changing their count. """ import sys,types from pathlib import Path import torch from torch import nn from torch.nn import functional as F ROOT=Path(__file__).resolve().parents[1] def sdpa_forward(self,x): b,t,_=x.shape q,k,v=[a.reshape(b,t,self.heads,-1).transpose(1,2) for a in self.to_qkv(self.norm(x)).chunk(3,dim=-1)] mask=self.bias[:,:,:t,:t].to(device=x.device,dtype=torch.bool) y=F.scaled_dot_product_attention(q,k,v,attn_mask=mask,dropout_p=0.0) return self.to_out(y.transpose(1,2).reshape(b,t,-1)) class DinoAdapter(nn.Module): def __init__(self,path,process,chunk_size=300): super().__init__() sys.path.insert(0,str(ROOT/'repos/dino_wm')) from models.dino import DinoV2Encoder from models.visual_world_model import VWorldModel from models.vit import Attention torch.hub.set_dir(str(ROOT/'checkpoints/dino-hub')) payload=torch.load(path,map_location='cpu',weights_only=False) print('DINO checkpoint keys:',list(payload),flush=True) encoder=DinoV2Encoder('dinov2_vits14','x_norm_patchtokens') self.wm=VWorldModel(image_size=224,num_hist=3,num_pred=1,encoder=encoder, proprio_encoder=payload['proprio_encoder'],action_encoder=payload['action_encoder'], decoder=None,predictor=payload['predictor'],proprio_dim=10,action_dim=10, concat_dim=1,num_action_repeat=1,num_proprio_repeat=1, train_encoder=False,train_predictor=False,train_decoder=False) self.wm.eval();self.wm.cuda();self.wm.requires_grad_(False) # Compare the optimized attention against the released implementation. for layer in self.wm.predictor.modules(): if isinstance(layer,Attention): layer.bias=layer.bias.cuda() if not hasattr(self,'attention_max_error'): with torch.no_grad(): x=torch.randn(1,32,layer.norm.normalized_shape[0],device='cuda') reference=layer(x);replacement=sdpa_forward(layer,x) torch.testing.assert_close(reference,replacement,rtol=2e-4,atol=2e-5) self.attention_max_error=float((reference-replacement).abs().max()) layer.forward=types.MethodType(sdpa_forward,layer) for k in ['action','proprio']: self.register_buffer(k+'_mean',torch.tensor(process[k].mean_,dtype=torch.float32,device='cuda')) self.register_buffer(k+'_std',torch.tensor(process[k].scale_,dtype=torch.float32,device='cuda')) self.register_buffer('native_action_mean',torch.tensor([-.0087,.0068],device='cuda')) self.register_buffer('native_action_std',torch.tensor([.2019,.2002],device='cuda')) self.register_buffer('native_proprio_mean',torch.tensor([236.6155,264.5674,-2.93032027,2.54307914],device='cuda')) self.register_buffer('native_proprio_std',torch.tensor([101.1202,87.0112,74.84556075,74.14009094],device='cuda')) self.chunk_size=chunk_size;self._info=None def _obs(self,info,goal=False): prop=info['goal_proprio' if goal else 'proprio'][:,0].cuda() prop=(prop*self.proprio_std+self.proprio_mean-self.native_proprio_mean)/self.native_proprio_std return {'visual':info['goal' if goal else 'pixels'][:,0].cuda(),'proprio':prop} def latent_rollout(self,encoded,actions): b,t,_=actions.shape visual=encoded['visual'].expand(b,-1,-1,-1) proprio=encoded['proprio'].expand(b,-1,-1) ae=self.wm.encode_act(actions[:,:1]) patches=visual.shape[2] z=torch.cat([visual,proprio[:,:,None].expand(-1,-1,patches,-1),ae[:,:,None].expand(-1,-1,patches,-1)],dim=-1) for step in range(1,t): new=self.wm.predict(z[:,-3:])[:,-1:] new=self.wm.replace_actions_from_z(new,actions[:,step:step+1]) z=torch.cat([z,new],dim=1) pred=self.wm.predict(z[:,-3:])[:,-1:] return self.wm.separate_emb(pred)[0] @torch.inference_mode() def get_cost(self,info,candidates): assert candidates.shape[0]==1,'The shared solver must use batch_size=1' if self._info is not info: self._info=info self._start=self.wm.encode_obs(self._obs(info)) self._goal=self.wm.encode_obs(self._obs(info,True)) actions=candidates[0].reshape(-1,5,5,2) actions=(actions*self.action_std+self.action_mean-self.native_action_mean)/self.native_action_std actions=actions.flatten(-2) costs=[] for chunk in actions.split(self.chunk_size): pred=self.latent_rollout(self._start,chunk) costs.append(sum((pred[k]-self._goal[k]).square().flatten(1).mean(1) for k in ['visual','proprio'])) return torch.cat(costs)[None]