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