BonanDing commited on
Commit
b236181
·
1 Parent(s): 927c275

Add DeMemWM geometry-aware memory attention

Browse files
algorithms/dememwm/models/dit.py CHANGED
@@ -150,11 +150,39 @@ class FrameMemoryReferenceAttention(nn.Module):
150
  self.to_q = nn.Linear(hidden_size, hidden_size, bias=False)
151
  self.to_k = nn.Linear(hidden_size, hidden_size, bias=False)
152
  self.to_v = nn.Linear(hidden_size, hidden_size, bias=False)
 
 
153
  self.out_proj = nn.Linear(hidden_size, hidden_size, bias=True)
154
 
155
  def _split_heads(self, x):
156
  return rearrange(x, "b n (h d) -> b h n d", h=self.num_heads)
157
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
158
  def forward(self, target_hidden, memory_hidden, memory_mask=None, geometry_cache=None):
159
  B, T_target, H, W, D = target_hidden.shape
160
  if memory_hidden is None or int(memory_hidden.shape[1]) == 0:
@@ -162,18 +190,34 @@ class FrameMemoryReferenceAttention(nn.Module):
162
 
163
  T_memory = memory_hidden.shape[1]
164
  P = H * W
165
- q = rearrange(target_hidden, "b t h w d -> (b t) (h w) d")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
166
  q = self._split_heads(self.to_q(q))
167
 
168
- memory_tokens = rearrange(memory_hidden, "b m h w d -> b (m h w) d")
169
- k = self._split_heads(self.to_k(memory_tokens))
170
- v = self._split_heads(self.to_v(memory_tokens))
171
-
172
- # A zero dummy key keeps all-padded samples finite; the projected delta
173
- # is masked back to zero below so padded memory cannot contribute.
174
- dummy = k.new_zeros((B, self.num_heads, 1, self.head_dim))
175
- k = torch.cat([k, dummy], dim=2)
176
- v = torch.cat([v, dummy], dim=2)
 
 
177
 
178
  if memory_mask is None:
179
  frame_valid = torch.ones((B, T_memory), device=target_hidden.device, dtype=torch.bool)
@@ -184,13 +228,22 @@ class FrameMemoryReferenceAttention(nn.Module):
184
  key_valid = torch.cat([key_valid, ~active[:, None]], dim=1)
185
 
186
  row_batch = torch.arange(B, device=target_hidden.device).repeat_interleave(T_target)
 
 
 
 
 
 
 
 
 
187
  attn_bias = torch.zeros((B * T_target, 1, 1, key_valid.shape[1]), device=q.device, dtype=q.dtype)
188
  attn_bias = attn_bias.masked_fill(~key_valid[row_batch][:, None, None], float("-inf"))
189
 
190
  x = F.scaled_dot_product_attention(
191
  query=q.contiguous(),
192
- key=k.index_select(0, row_batch).contiguous(),
193
- value=v.index_select(0, row_batch).contiguous(),
194
  attn_mask=attn_bias,
195
  )
196
  x = rearrange(x, "n h p d -> n p (h d)")
@@ -658,16 +711,18 @@ class DiT(nn.Module):
658
  else:
659
  pc = None
660
 
661
- frame_memory_geometry = self._build_frame_memory_geometry(
662
- frame_memory_segments,
663
- frame_memory_masks,
664
- frame_memory_pose,
665
- image_hw,
666
- x.shape[2],
667
- x.shape[3],
668
- x.device,
669
- x.dtype,
670
- )
 
 
671
 
672
  for i, block in enumerate(self.blocks):
673
  x = block(x, c, current_frame=current_frame, timestep=t, is_last_block= (i+1 == len(self.blocks)),
 
150
  self.to_q = nn.Linear(hidden_size, hidden_size, bias=False)
151
  self.to_k = nn.Linear(hidden_size, hidden_size, bias=False)
152
  self.to_v = nn.Linear(hidden_size, hidden_size, bias=False)
153
+ self.query_pose_proj = nn.Linear(6, hidden_size, bias=True)
154
+ self.key_pose_proj = nn.Linear(6, hidden_size, bias=True)
155
  self.out_proj = nn.Linear(hidden_size, hidden_size, bias=True)
156
 
157
  def _split_heads(self, x):
158
  return rearrange(x, "b n (h d) -> b h n d", h=self.num_heads)
159
 
160
+ def _geometry_rays(
161
+ self,
162
+ geometry_cache,
163
+ B,
164
+ T_target,
165
+ T_memory,
166
+ H,
167
+ W,
168
+ device,
169
+ dtype,
170
+ ):
171
+ if not isinstance(geometry_cache, dict):
172
+ return None, None
173
+ query_rays = geometry_cache.get("query_rays")
174
+ relative_rays = geometry_cache.get("relative_rays")
175
+ if not torch.is_tensor(query_rays) or not torch.is_tensor(relative_rays):
176
+ return None, None
177
+ if tuple(query_rays.shape) != (B, T_target, H, W, 6):
178
+ return None, None
179
+ if tuple(relative_rays.shape) != (B, T_target, T_memory, H, W, 6):
180
+ return None, None
181
+ return (
182
+ query_rays.to(device=device, dtype=dtype),
183
+ relative_rays.to(device=device, dtype=dtype),
184
+ )
185
+
186
  def forward(self, target_hidden, memory_hidden, memory_mask=None, geometry_cache=None):
187
  B, T_target, H, W, D = target_hidden.shape
188
  if memory_hidden is None or int(memory_hidden.shape[1]) == 0:
 
190
 
191
  T_memory = memory_hidden.shape[1]
192
  P = H * W
193
+ query_rays, relative_rays = self._geometry_rays(
194
+ geometry_cache,
195
+ B,
196
+ T_target,
197
+ T_memory,
198
+ H,
199
+ W,
200
+ target_hidden.device,
201
+ target_hidden.dtype,
202
+ )
203
+ q_input = target_hidden
204
+ if query_rays is not None:
205
+ # Values stay visual-only; Plucker features only steer Q/K matching.
206
+ q_input = q_input + self.query_pose_proj(query_rays)
207
+ q = rearrange(q_input, "b t h w d -> (b t) (h w) d")
208
  q = self._split_heads(self.to_q(q))
209
 
210
+ memory_tokens = rearrange(memory_hidden, "b m h w d -> b m (h w) d")
211
+ memory_flat = rearrange(memory_tokens, "b m p d -> b (m p) d")
212
+ if relative_rays is None:
213
+ k = self._split_heads(self.to_k(memory_flat))
214
+ else:
215
+ key_pose = self.key_pose_proj(
216
+ rearrange(relative_rays, "b t m h w r -> b t m (h w) r")
217
+ )
218
+ k = self.to_k(memory_tokens[:, None] + key_pose)
219
+ k = self._split_heads(rearrange(k, "b t m p d -> (b t) (m p) d"))
220
+ v = self._split_heads(self.to_v(memory_flat))
221
 
222
  if memory_mask is None:
223
  frame_valid = torch.ones((B, T_memory), device=target_hidden.device, dtype=torch.bool)
 
228
  key_valid = torch.cat([key_valid, ~active[:, None]], dim=1)
229
 
230
  row_batch = torch.arange(B, device=target_hidden.device).repeat_interleave(T_target)
231
+ if relative_rays is None:
232
+ k = k.index_select(0, row_batch)
233
+ # A zero dummy key keeps all-padded samples finite; the projected delta
234
+ # is masked back to zero below so padded memory cannot contribute.
235
+ dummy_k = k.new_zeros((k.shape[0], self.num_heads, 1, self.head_dim))
236
+ k = torch.cat([k, dummy_k], dim=2)
237
+ dummy_v = v.new_zeros((B, self.num_heads, 1, self.head_dim))
238
+ v = torch.cat([v, dummy_v], dim=2).index_select(0, row_batch)
239
+
240
  attn_bias = torch.zeros((B * T_target, 1, 1, key_valid.shape[1]), device=q.device, dtype=q.dtype)
241
  attn_bias = attn_bias.masked_fill(~key_valid[row_batch][:, None, None], float("-inf"))
242
 
243
  x = F.scaled_dot_product_attention(
244
  query=q.contiguous(),
245
+ key=k.contiguous(),
246
+ value=v.contiguous(),
247
  attn_mask=attn_bias,
248
  )
249
  x = rearrange(x, "n h p d -> n p (h d)")
 
711
  else:
712
  pc = None
713
 
714
+ frame_memory_geometry = None
715
+ if self.use_memory_attention and self.use_plucker:
716
+ frame_memory_geometry = self._build_frame_memory_geometry(
717
+ frame_memory_segments,
718
+ frame_memory_masks,
719
+ frame_memory_pose,
720
+ image_hw,
721
+ x.shape[2],
722
+ x.shape[3],
723
+ x.device,
724
+ x.dtype,
725
+ )
726
 
727
  for i, block in enumerate(self.blocks):
728
  x = block(x, c, current_frame=current_frame, timestep=t, is_last_block= (i+1 == len(self.blocks)),
tests/test_dememwm_temporal_attention.py CHANGED
@@ -264,6 +264,7 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
264
  action_cond_dim=3,
265
  pose_cond_dim=0,
266
  reference_length=0,
 
267
  use_memory_attention=True,
268
  )
269
  spies = []
@@ -328,5 +329,48 @@ class DeMemWMTemporalAttentionTests(unittest.TestCase):
328
  self.assertTrue(torch.allclose(second_only, torch.full_like(second_only, 10.0)))
329
  self.assertTrue(torch.equal(empty, torch.zeros_like(empty)))
330
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
331
  if __name__ == "__main__":
332
  unittest.main()
 
264
  action_cond_dim=3,
265
  pose_cond_dim=0,
266
  reference_length=0,
267
+ use_plucker=True,
268
  use_memory_attention=True,
269
  )
270
  spies = []
 
329
  self.assertTrue(torch.allclose(second_only, torch.full_like(second_only, 10.0)))
330
  self.assertTrue(torch.equal(empty, torch.zeros_like(empty)))
331
 
332
+ def test_frame_memory_reference_attention_uses_geometry_qk(self):
333
+ attn = FrameMemoryReferenceAttention(hidden_size=2, num_heads=1)
334
+ with torch.no_grad():
335
+ attn.to_q.weight.zero_()
336
+ attn.to_q.weight[0, 0] = 1.0
337
+ attn.to_k.weight.zero_()
338
+ attn.to_k.weight[0, 0] = 1.0
339
+ attn.to_v.weight.zero_()
340
+ attn.to_v.weight[0, 1] = 1.0
341
+ attn.query_pose_proj.weight.zero_()
342
+ attn.query_pose_proj.bias.zero_()
343
+ attn.query_pose_proj.weight[0, 5] = 1.0
344
+ attn.key_pose_proj.weight.zero_()
345
+ attn.key_pose_proj.bias.zero_()
346
+ attn.key_pose_proj.weight[0, 5] = 1.0
347
+ attn.out_proj.weight.zero_()
348
+ attn.out_proj.weight[0, 0] = 1.0
349
+ attn.out_proj.bias.zero_()
350
+
351
+ target = torch.zeros((1, 1, 1, 1, 2))
352
+ memory = torch.zeros((1, 2, 1, 1, 2))
353
+ memory[0, :, 0, 0, 1] = torch.tensor([1.0, 10.0])
354
+ query_rays = torch.zeros((1, 1, 1, 1, 6))
355
+ query_rays[..., 5] = 1.0
356
+ relative_a = torch.zeros((1, 1, 2, 1, 1, 6))
357
+ relative_b = torch.zeros_like(relative_a)
358
+ relative_a[0, 0, 0, 0, 0, 5] = 2.0
359
+ relative_b[0, 0, 1, 0, 0, 5] = 2.0
360
+
361
+ out_a = attn(
362
+ target,
363
+ memory,
364
+ geometry_cache={"query_rays": query_rays, "relative_rays": relative_a},
365
+ )
366
+ out_b = attn(
367
+ target,
368
+ memory,
369
+ geometry_cache={"query_rays": query_rays, "relative_rays": relative_b},
370
+ )
371
+
372
+ self.assertLess(out_a[0, 0, 0, 0, 0].item(), out_b[0, 0, 0, 0, 0].item())
373
+ self.assertFalse(torch.allclose(out_a, out_b))
374
+
375
  if __name__ == "__main__":
376
  unittest.main()