Download code/dino_adapter.py from FidelityWM/planning-baselines: direct link, hf CLI and curl.
- Browser
- Download file 4.5 kB
-
https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/dino_adapter.py
- Command line
-
hf download hf://FidelityWM/planning-baselines/code/dino_adapter.py
-
curl -L -o dino_adapter.py https://huggingface.co/FidelityWM/planning-baselines/resolve/main/code/dino_adapter.py
4.5 kB
| """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] | |
| 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] | |