BonanDing commited on
Commit
1470930
·
1 Parent(s): 270eca8

Remove DeMemWM attention debugger trap

Browse files
algorithms/dememwm/models/attention.py CHANGED
@@ -94,17 +94,6 @@ class TemporalAxialAttention(nn.Module):
94
  def forward(self, x: torch.Tensor, frame_memory_segments=None, frame_memory_masks=None):
95
  B, T, H, W, D = x.shape
96
 
97
- # if T>=9:
98
- # try:
99
- # # x = torch.cat([x[:,:-1],x[:,16-T:17-T],x[:,-1:]], dim=1)
100
- # x = torch.cat([x[:,16-T:17-T],x], dim=1)
101
- # except:
102
- # import pdb;pdb.set_trace()
103
- # print("="*50)
104
- # print(x.shape)
105
-
106
- B, T, H, W, D = x.shape
107
-
108
  q, k, v = self.to_qkv(x).chunk(3, dim=-1)
109
 
110
  if self.use_domain_adapter:
@@ -145,24 +134,13 @@ class TemporalAxialAttention(nn.Module):
145
  else:
146
  attn_bias = None
147
 
148
- try:
149
- x = F.scaled_dot_product_attention(query=q, key=k, value=v, attn_mask=attn_bias)
150
- except:
151
- import pdb;pdb.set_trace()
152
 
153
  x = rearrange(x, "(B H W) h T d -> B T H W (h d)", B=B, H=H, W=W)
154
  x = x.to(q.dtype)
155
 
156
  # linear proj
157
  x = self.to_out(x)
158
-
159
- # if T>=10:
160
- # try:
161
- # # x = torch.cat([x[:,:-2],x[:,-1:]], dim=1)
162
- # x = x[:,1:]
163
- # except:
164
- # import pdb;pdb.set_trace()
165
- # print(x.shape)
166
  return x
167
 
168
  class SpatialAxialAttention(nn.Module):
 
94
  def forward(self, x: torch.Tensor, frame_memory_segments=None, frame_memory_masks=None):
95
  B, T, H, W, D = x.shape
96
 
 
 
 
 
 
 
 
 
 
 
 
97
  q, k, v = self.to_qkv(x).chunk(3, dim=-1)
98
 
99
  if self.use_domain_adapter:
 
134
  else:
135
  attn_bias = None
136
 
137
+ x = F.scaled_dot_product_attention(query=q, key=k, value=v, attn_mask=attn_bias)
 
 
 
138
 
139
  x = rearrange(x, "(B H W) h T d -> B T H W (h d)", B=B, H=H, W=W)
140
  x = x.to(q.dtype)
141
 
142
  # linear proj
143
  x = self.to_out(x)
 
 
 
 
 
 
 
 
144
  return x
145
 
146
  class SpatialAxialAttention(nn.Module):