File size: 4,499 Bytes
9375a59
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
"""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]